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})