"""FastAPI wrapper around NVIDIA NeMo Canary ASR models. Native Canary HTTP API (POST /inference). Not a Whisper /asr clone. Inspired by EvilFreelancer/docker-canary-serve, written from scratch. """ import copy import json import logging import os import subprocess import tempfile import time import wave from typing import Optional import torch from fastapi import FastAPI, File, Form, HTTPException, UploadFile from fastapi.responses import JSONResponse, PlainTextResponse from nemo.collections.asr.models import ASRModel from nemo.collections.asr.parts.utils.streaming_utils import FrameBatchMultiTaskAED from nemo.collections.asr.parts.utils.transcribe_utils import get_buffered_pred_feat_multitaskAED from omegaconf import OmegaConf logging.basicConfig( level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s", ) log = logging.getLogger("canary") MODEL_NAME = os.getenv("CANARY_MODEL", "nvidia/canary-1b-v2") DEVICE = os.getenv("CANARY_DEVICE", "cuda") PRECISION = os.getenv("CANARY_PRECISION", "bf16").lower() MODELS_DIR = os.getenv("CANARY_MODELS_DIR", "/models") CHUNK_LEN_SECS = float(os.getenv("CANARY_CHUNK_LEN_SECS", "40.0")) CHUNKED_THRESHOLD_SECS = float(os.getenv("CANARY_CHUNKED_THRESHOLD_SECS", "25.0")) MODEL_STRIDE = int(os.getenv("CANARY_MODEL_STRIDE", "8")) # FastConformer = 8 DTYPE_MAP = { "fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16, } if PRECISION not in DTYPE_MAP: raise ValueError(f"Unsupported CANARY_PRECISION={PRECISION}; use fp32, fp16, or bf16") os.environ.setdefault("HF_HOME", MODELS_DIR) log.info("Loading model=%s device=%s precision=%s", MODEL_NAME, DEVICE, PRECISION) t_load = time.monotonic() model = ASRModel.from_pretrained(MODEL_NAME) model = model.to(DEVICE) if DTYPE_MAP[PRECISION] != torch.float32: model = model.to(DTYPE_MAP[PRECISION]) model.eval() log.info("Model loaded in %.1fs", time.monotonic() - t_load) # Prepare config for the buffered/chunked inference path (long audio). # See CHUNKING.md and NVIDIA-NeMo's speech_to_text_aed_chunked_infer.py. _model_cfg = copy.deepcopy(model._cfg) OmegaConf.set_struct(_model_cfg.preprocessor, False) _model_cfg.preprocessor.dither = 0.0 _model_cfg.preprocessor.pad_to = 0 OmegaConf.set_struct(_model_cfg.preprocessor, True) if _model_cfg.preprocessor.normalize != "per_feature": log.warning( "Preprocessor normalize=%s; chunked inference assumes per_feature", _model_cfg.preprocessor.normalize, ) MODEL_STRIDE_SECS = float(_model_cfg.preprocessor["window_stride"]) * MODEL_STRIDE log.info( "Long-audio config: chunk_len=%.1fs threshold=%.1fs model_stride=%.4fs", CHUNK_LEN_SECS, CHUNKED_THRESHOLD_SECS, MODEL_STRIDE_SECS, ) app = FastAPI(title="doctate-canary", version="0.1.0") def _transcode_to_wav16k(src_path: str) -> str: """Convert any audio file to 16 kHz mono WAV. Returns the new path.""" dst = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) dst.close() cmd = [ "ffmpeg", "-y", "-loglevel", "error", "-i", src_path, "-ac", "1", "-ar", "16000", "-f", "wav", dst.name, ] proc = subprocess.run(cmd, capture_output=True, text=True) if proc.returncode != 0: os.unlink(dst.name) raise HTTPException(status_code=400, detail=f"ffmpeg failed: {proc.stderr.strip()}") return dst.name def _wav_duration(path: str) -> float: with wave.open(path, "rb") as w: return w.getnframes() / float(w.getframerate()) def _transcribe_buffered( wav_path: str, duration: float, language: str, pnc: str, timestamps: bool, ): """Long-audio path: explicit chunked inference via FrameBatchMultiTaskAED. Writes a temporary single-line JSONL manifest because language and pnc must be passed via manifest fields when using the buffered API. """ manifest_tmp = tempfile.NamedTemporaryFile(suffix=".json", mode="w", delete=False) try: json.dump( { "audio_filepath": wav_path, "duration": duration, "taskname": "asr", "source_lang": language, "target_lang": language, "pnc": pnc, }, manifest_tmp, ) manifest_tmp.write("\n") manifest_tmp.close() frame_asr = FrameBatchMultiTaskAED( asr_model=model, frame_len=CHUNK_LEN_SECS, total_buffer=CHUNK_LEN_SECS, batch_size=8, ) amp_dtype = DTYPE_MAP[PRECISION] use_amp = amp_dtype != torch.float32 with torch.amp.autocast(model.device.type, enabled=use_amp, dtype=amp_dtype): with torch.no_grad(): hyps = get_buffered_pred_feat_multitaskAED( frame_asr, _model_cfg.preprocessor, MODEL_STRIDE_SECS, model.device, manifest=manifest_tmp.name, filepaths=None, timestamps=timestamps, ) return hyps finally: try: os.unlink(manifest_tmp.name) except OSError: pass @app.get("/health") def health(): return {"status": "ok", "model": MODEL_NAME, "device": DEVICE, "precision": PRECISION} @app.get("/info") def info(): return { "model": MODEL_NAME, "device": DEVICE, "precision": PRECISION, "models_dir": MODELS_DIR, } @app.post("/inference") async def inference( file: UploadFile = File(...), language: str = Form("de"), pnc: str = Form("yes"), timestamps: str = Form("no"), response_format: str = Form("text"), ): if language not in {"en", "de", "fr", "es"}: raise HTTPException(status_code=400, detail=f"Unsupported language: {language}") if pnc not in {"yes", "no"}: raise HTTPException(status_code=400, detail="pnc must be yes or no") if timestamps not in {"yes", "no"}: raise HTTPException(status_code=400, detail="timestamps must be yes or no") if response_format not in {"text", "json"}: raise HTTPException(status_code=400, detail="response_format must be text or json") upload_suffix = os.path.splitext(file.filename or "")[1] or ".bin" upload_tmp = tempfile.NamedTemporaryFile(suffix=upload_suffix, delete=False) wav_path: Optional[str] = None try: upload_tmp.write(await file.read()) upload_tmp.flush() upload_tmp.close() wav_path = _transcode_to_wav16k(upload_tmp.name) duration = _wav_duration(wav_path) ts_flag = timestamps == "yes" t0 = time.monotonic() chunked = duration > CHUNKED_THRESHOLD_SECS if chunked: hyps = _transcribe_buffered(wav_path, duration, language, pnc, ts_flag) else: hyps = model.transcribe( [wav_path], source_lang=language, target_lang=language, pnc=pnc, timestamps=ts_flag, ) infer_secs = time.monotonic() - t0 if not hyps: raise HTTPException(status_code=500, detail="Empty transcription result") hyp = hyps[0] text = (getattr(hyp, "text", None) or str(hyp)).strip() log.info(json.dumps({ "event": "transcribe", "duration_audio": round(duration, 2), "duration_infer": round(infer_secs, 2), "language": language, "chars": len(text), "chunked": chunked, })) if response_format == "json": payload = { "text": text, "language": language, "duration": duration, } if timestamps == "yes": ts = getattr(hyp, "timestamp", None) if ts: payload["segments"] = ts.get("segment", []) return JSONResponse(payload) return PlainTextResponse(text) finally: for p in (upload_tmp.name, wav_path): if p and os.path.exists(p): try: os.unlink(p) except OSError: pass