Files
esh-pfi-infrastructure/services/parakeet-ab-2026-09-30/code/serve_nemo.py
T
vh a6c1d3c454 docs(parakeet): seat A/B vs parakeet-unified-en-0.6b - latency is the int8-on-CPU runtime; unified wins WER
A/B of the live STT seat (fv-ml1 GPU 0, sherpa-onnx int8 v3) against
nvidia/parakeet-unified-en-0.6b, measured on GPU 3 with the seat's own image,
k2-fsa's published unified int8 export, fp32/fp16 exports made with k2-fsa's
recipe, v2 int8, and NeMo 3.0.0 (fp32, bf16 autocast, bf16 weights).

- Seat int8 graph runs on one CPU thread (cpu/wall 1.00, GPU 2-9%).
- unified-en under NeMo: -121/-234/-530 ms vs the seat at 1-3/3-8/8-20 s
  (paired, n=120/bin; floor <=6 ms; +50 ms positive control reads +52-54).
- unified-en WER lower in every runtime: -0.7 pp clean, -1.5 pp other,
  -3.2 to -4.4 pp AMI (paired CIs exclude 0).
- Seat defects found: hard 400 s input ceiling (HTTP 500), truncation after
  a quiet 1.5 s pause, and severe long-window dropouts (int8 v3 only).
- B-bf16w needs +0.8 to +1.5 GB over the seat's 1,690 MiB on GPU 0.

Raw requests, hypotheses, manifests and the full harness under
services/parakeet-ab-2026-09-30/. No deploy; live seat untouched apart
from 240 light test requests.
2026-09-30 18:51:44 -07:00

154 lines
5.8 KiB
Python

"""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())