- 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)
81 lines
2.8 KiB
Python
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})
|