"""A/B arm B: parakeet-unified-en-0.6b under NeMo torch, behind the SAME thin HTTP shape as the seat. Mirrors stacks/parakeet/app.py: model loaded at import, one warm-up decode before traffic, and `async def` handlers that call a blocking decode (so, like the seat, requests are serialised on the event loop). Same endpoints, same {"text": ...} body, plus the harness's x-ab-decode-ms header. Decode path (PATH_MODE): direct preprocessor -> encoder -> RNNT greedy_batch, no DataLoader (what a latency seat would run) transcribe ASRModel.transcribe([samples]) (NeMo's public API; used to check `direct` gives the same text) DTYPE: fp32 | bf16 (torch.autocast over encoder+decoding, fp32 weights) | bf16w (encoder, decoder and joint WEIGHTS cast to bf16, no autocast; the mel front end stays fp32). LOAD_CPU=1: restore the .nemo on the CPU, then move it to the GPU (avoids the ~1.8 GB load-time transient). LOCAL_ATT=L,R: switch the encoder to rel_pos_local_attn with that context (NeMo's documented long-audio mode; memory linear in length instead of quadratic). Unset = full attention, as the seat runs. CUDA_GRAPHS: 1 | 0 for NeMo's label-looping CUDA-graph greedy decoder. """ from __future__ import annotations import io import logging import os import time import numpy as np import soundfile as sf import torch from fastapi import FastAPI, File, HTTPException, UploadFile from fastapi.responses import JSONResponse MODEL_PATH = os.environ["MODEL_PATH"] DTYPE = os.environ.get("DTYPE", "fp32") CUDA_GRAPHS = os.environ.get("CUDA_GRAPHS", "1") == "1" PATH_MODE = os.environ.get("PATH_MODE", "direct") LOAD_CPU = os.environ.get("LOAD_CPU", "0") == "1" LOCAL_ATT = os.environ.get("LOCAL_ATT", "") SR = 16000 logger = logging.getLogger("parakeet-nemo") logging.basicConfig(level=os.environ.get("LOG_LEVEL", "INFO")) def _load(): import nemo.collections.asr as nemo_asr from omegaconf import open_dict t0 = time.monotonic() if LOAD_CPU: m = nemo_asr.models.ASRModel.restore_from(MODEL_PATH, map_location="cpu").to("cuda") else: m = nemo_asr.models.ASRModel.restore_from(MODEL_PATH, map_location="cuda") m.eval() if m.cfg.get("validation_ds") is None: # the unified .nemo ships without it; transcribe() reads it with open_dict(m.cfg): m.cfg.validation_ds = {} d = m.cfg.decoding with open_dict(d): d.strategy = "greedy_batch" d.greedy["use_cuda_graph_decoder"] = CUDA_GRAPHS m.change_decoding_strategy(d, verbose=False) # What transcribe() does on entry: no dither, no pad_to. The direct path must match it. m.preprocessor.featurizer.dither = 0.0 m.preprocessor.featurizer.pad_to = 0 if LOCAL_ATT: m.change_attention_model("rel_pos_local_attn", [int(x) for x in LOCAL_ATT.split(",")]) if DTYPE == "bf16w": for mod in (m.encoder, m.decoder, m.joint): mod.to(torch.bfloat16) logger.info("loaded %s (%s) dtype=%s cuda_graphs=%s path=%s local_att=%s in %.1fs", os.path.basename(MODEL_PATH), type(m).__name__, DTYPE, CUDA_GRAPHS, PATH_MODE, LOCAL_ATT or "full", time.monotonic() - t0) return m model = _load() _AMP = dict(device_type="cuda", dtype=torch.bfloat16, enabled=(DTYPE == "bf16")) _FEAT_DTYPE = torch.bfloat16 if DTYPE == "bf16w" else torch.float32 def _hyp_text(h) -> str: if isinstance(h, str): return h t = getattr(h, "text", None) if isinstance(t, str): return t return model.tokenizer.ids_to_text([int(i) for i in h.y_sequence]) @torch.inference_mode() def _infer(samples: np.ndarray) -> str: if PATH_MODE == "transcribe": with torch.autocast(**_AMP): out = model.transcribe([samples], batch_size=1, verbose=False) if isinstance(out, tuple): out = out[0] return _hyp_text(out[0]) x = torch.from_numpy(samples).to("cuda", non_blocking=True).unsqueeze(0) xl = torch.tensor([x.shape[1]], device="cuda", dtype=torch.long) feats, fl = model.preprocessor(input_signal=x, length=xl) feats = feats.to(_FEAT_DTYPE) with torch.autocast(**_AMP): enc, el = model.encoder(audio_signal=feats, length=fl) hyps = model.decoding.rnnt_decoder_predictions_tensor(encoder_output=enc, encoded_lengths=el, return_hypotheses=False) if isinstance(hyps, tuple): hyps = hyps[0] return _hyp_text(hyps[0]) def _warm() -> None: """Same as the seat: one throwaway 1 s of silence before traffic.""" t0 = time.monotonic() _infer(np.zeros(SR, dtype=np.float32)) torch.cuda.synchronize() logger.info("warmup decode complete in %.1fs", time.monotonic() - t0) _warm() app = FastAPI(title="Parakeet ASR (NeMo torch, A/B arm)") def _decode(raw: bytes) -> str: try: samples, sample_rate = sf.read(io.BytesIO(raw), dtype="float32") except Exception as exc: raise HTTPException(400, f"Could not decode audio: {exc}") from exc if samples.ndim > 1: samples = samples.mean(axis=1).astype(np.float32) if sample_rate != SR: import torchaudio.functional as AF samples = AF.resample(torch.from_numpy(samples), sample_rate, SR).numpy() return _infer(np.ascontiguousarray(samples, dtype=np.float32)) def _timed(raw: bytes) -> JSONResponse: t0 = time.perf_counter() text = _decode(raw) torch.cuda.synchronize() ms = (time.perf_counter() - t0) * 1000.0 return JSONResponse({"text": text}, headers={"x-ab-decode-ms": f"{ms:.3f}"}) @app.get("/healthz") def healthz() -> dict[str, str]: return {"status": "ok"} @app.post("/transcribe") async def transcribe(file: UploadFile = File(...)): return _timed(await file.read()) @app.post("/v1/audio/transcriptions") async def openai_transcriptions(file: UploadFile = File(...)): return _timed(await file.read())