Files
voice-assistant-stack/asr-server/app.py
T
bhetherman c2b92d96cb Initial voice-assistant-stack: open-webui + spotify-voice-assistant submodules, asr/tts sidecars
- open-webui: submodule, LAN OIDC/Authentik-fronted, GPU-passthrough ollama,
  ollama-auth proxy, asr/tts wired via docker-compose.audio.yaml
- spotify-voice-assistant: submodule, Home Assistant custom integration
- tts-server / asr-server: plain directories, built as sidecars by
  open-webui/docker-stack.sh (no separate git history)
2026-08-31 00:47:35 -04:00

81 lines
2.8 KiB
Python

import hmac
import logging
import os
import tempfile
import torch
from fastapi import Depends, FastAPI, File, Form, HTTPException, UploadFile
from fastapi.responses import JSONResponse, PlainTextResponse
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from transformers import AutoModelForRNNT, AutoProcessor
from transformers.audio_utils import load_audio
logging.basicConfig(level=logging.INFO)
log = logging.getLogger("asr-server")
AUTH_TOKEN = os.environ["ASR_AUTH_TOKEN"]
bearer_scheme = HTTPBearer()
def require_auth(creds: HTTPAuthorizationCredentials = Depends(bearer_scheme)) -> None:
if not hmac.compare_digest(creds.credentials, AUTH_TOKEN):
raise HTTPException(status_code=401, detail="invalid token")
MODEL_ID = os.environ.get("ASR_MODEL_ID", "nvidia/nemotron-3.5-asr-streaming-0.6b")
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32
# The RNN-T decoder's LSTM layers hit cuDNN's fused RNN kernel on `.to(device)`,
# which requires SM >= 7.5 in recent cuDNN builds. This GPU (Pascal, SM 6.1) is
# below that floor, so disable cuDNN and fall back to PyTorch's generic CUDA
# RNN kernels instead.
torch.backends.cudnn.enabled = False
app = FastAPI()
log.info("Loading %s on %s (%s)", MODEL_ID, DEVICE, DTYPE)
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = AutoModelForRNNT.from_pretrained(MODEL_ID, dtype=DTYPE).to(DEVICE).eval()
SAMPLING_RATE = processor.feature_extractor.sampling_rate
log.info("Model ready (sampling_rate=%s)", SAMPLING_RATE)
@app.get("/health")
def health():
return {"status": "ok", "device": DEVICE}
@app.post("/v1/audio/transcriptions", dependencies=[Depends(require_auth)])
async def transcribe(
file: UploadFile = File(...),
model_name: str = Form("nemotron-3.5-asr-streaming-0.6b", alias="model"),
language: str | None = Form(None),
response_format: str = Form("json"),
):
suffix = os.path.splitext(file.filename or "")[1] or ".wav"
data = await file.read()
with tempfile.NamedTemporaryFile(suffix=suffix) as tmp:
tmp.write(data)
tmp.flush()
try:
audio = load_audio(tmp.name, sampling_rate=SAMPLING_RATE)
except Exception as exc:
raise HTTPException(status_code=400, detail=f"could not decode audio: {exc}") from exc
lang = language or "auto"
inputs = processor(audio, sampling_rate=SAMPLING_RATE, language=lang)
inputs = inputs.to(DEVICE, dtype=DTYPE)
with torch.inference_mode():
output = model.generate(**inputs, return_dict_in_generate=True)
text = processor.decode(output.sequences, skip_special_tokens=True)
if isinstance(text, list):
text = text[0] if text else ""
if response_format == "text":
return PlainTextResponse(text)
return JSONResponse({"text": text})