The gen-small EngineCore OOM (04:21 PT): our parked 3,582 MiB window cache left no room for vLLM's runtime workspace. Seat-side fix, three controls: - windowed path wraps every window in torch.cuda.empty_cache(), so the seat returns to ~2,108 MiB rest after a 12-min file instead of parking at the peak (measured: peak 3,028 MiB during, rest after, restarts=0); - MEM_CAP_MIB=3840 hard set_per_process_memory_fraction: over-cap requests answer 503 with the seat alive (proved at cap=2000), so the failure lands on us, never on a neighbour; - CUDA_GRAPHS=0: the graph decoder pins cache blocks that empty_cache must free (illegal-memory-access wedge when both were on first try). Cost: 12-min file 3.0 s vs 1.2 s, short bins 35-62 ms vs 33-42 ms -- still 4-15x under the sherpa seat. Measured, not computed: gen-small moved ZERO from 36,116 MiB across three realistic requests (1,351 in / ~180 out) -- its workspace lands at engine init; the growth window is restart-relative, matching infra-ops's observation. WINDOW_S is now a real compose tunable. README memory section rewritten.
175 lines
7.7 KiB
Python
175 lines
7.7 KiB
Python
"""Parakeet ASR seat: nvidia/parakeet-unified-en-0.6b under NeMo torch (bf16 weights).
|
||
|
||
Fork of the A/B winner (services/parakeet-ab-2026-09-30 arm B-bf16w), with the three shipping
|
||
changes that arm's doc called for and a longer warm-up. Same HTTP shape as the sherpa seat it
|
||
replaces: model loaded at import, warm-up before traffic, async handlers over a serialised
|
||
blocking decode, {"text": ...} responses on /transcribe and /v1/audio/transcriptions.
|
||
|
||
Hard-wired, because every one of these is load-bearing on GPU 0:
|
||
bf16 weights — the encoder/decoder/joint are cast to bfloat16 (mel front end stays fp32).
|
||
WER is identical to fp32 (399/400 utterances byte-equal in the A/B).
|
||
CPU-then-cast — the .nemo restores on CPU, weights are cast to bf16 there, and only then move
|
||
to the GPU. Restoring to CUDA spikes the load by ~1.5 GB; GPU 0 cannot absorb it.
|
||
local attn — rel_pos_local_attn ±128 (NeMo's documented long-audio mode). This is what fixes
|
||
the old seat's >400 s HTTP 500 and the long-form dropouts; memory grows linearly
|
||
in file length. A 30-min file transcribes in ~2.6 s in one request (A/B, 2026-09-30).
|
||
Warm-up runs ascending silent clips (1 s, 8 s, 60 s): the CUDA-graph greedy decoder and the
|
||
encoder kernels cost ~330 ms extra on the first call at a new maximum length.
|
||
"""
|
||
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"]
|
||
WARMUP_SECONDS = [int(x) for x in os.environ.get("WARMUP_SECONDS", "1,8,60").split(",")]
|
||
# Hard ceiling for the whole process, MiB. The measured window peak is 3,582 (audit 2026-10-01);
|
||
# 3,840 gives a little headroom and NO more. GPU 0 is shared with vLLM seats that grow at RUNTIME
|
||
# (~0.8 GB for gen-small), and an unbounded window cache OOMed gen-small's EngineCore on its first
|
||
# request after the switch — a cached peak is an unpaid debt to the neighbours.
|
||
MEM_CAP_MIB = int(os.environ.get("MEM_CAP_MIB", "3840"))
|
||
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()
|
||
m = nemo_asr.models.ASRModel.restore_from(MODEL_PATH, map_location="cpu")
|
||
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"] = os.environ.get("CUDA_GRAPHS", "1") == "1"
|
||
m.change_decoding_strategy(d, verbose=False)
|
||
# transcribe() sets these on entry; the direct path must match, and must not dither (dither is
|
||
# a training-time augmentation and makes the same file decode differently on each call).
|
||
m.preprocessor.featurizer.dither = 0.0
|
||
m.preprocessor.featurizer.pad_to = 0
|
||
m.change_attention_model("rel_pos_local_attn", [128, 128])
|
||
# bf16 the serving modules BEFORE the H2D copy: halves the transfer and skips the GPU-side
|
||
# fp32->bf16 transient entirely (the measured load spike goes from 3,194 MiB to under 2,600).
|
||
# ⚠ Must run AFTER change_attention_model: that call rebuilds the attention modules in fp32,
|
||
# and casting first leaves fp32 islands behind (RuntimeError: mat1 and mat2 ... BFloat16/Float,
|
||
# hit live at the first boot of this image, 2026-10-01).
|
||
for mod in (m.encoder, m.decoder, m.joint):
|
||
mod.to(torch.bfloat16)
|
||
m = m.to("cuda")
|
||
logger.info("loaded %s (%s) bf16w local_att=128,128 in %.1fs", os.path.basename(MODEL_PATH),
|
||
type(m).__name__, time.monotonic() - t0)
|
||
return m
|
||
|
||
|
||
model = _load()
|
||
# Hard per-process ceiling on torch allocations (see MEM_CAP_MIB): an over-size request must
|
||
# fail HERE, at us, instead of stealing runtime room from a neighbour's process. torch counts
|
||
# RESERVED bytes against this, which is exactly the ledger we want capped.
|
||
torch.cuda.set_per_process_memory_fraction(MEM_CAP_MIB / (torch.cuda.get_device_properties(0).total_memory / 2**20))
|
||
|
||
|
||
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:
|
||
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(torch.bfloat16)
|
||
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:
|
||
for secs in WARMUP_SECONDS:
|
||
t0 = time.monotonic()
|
||
_infer(np.zeros(SR * secs, dtype=np.float32))
|
||
torch.cuda.synchronize()
|
||
logger.info("warmup %ss decode complete in %.1fs", secs, time.monotonic() - t0)
|
||
|
||
|
||
_warm()
|
||
app = FastAPI(title="Parakeet ASR (NeMo torch, bf16w)")
|
||
|
||
|
||
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()
|
||
samples = np.ascontiguousarray(samples, dtype=np.float32)
|
||
# Windowed long-form (> WINDOW_S): NeMo's attention mask is materialised T×T even under
|
||
# rel_pos_local_attn, so one whole 12-min request wanted +1.1 GiB of scratch and OOMed on
|
||
# GPU 0 (measured at acceptance, 2026-10-01). The A/B's own long-form arm used ~6-min
|
||
# windows and lost zero clean speech on unified-en in 4/4 placements (doc § 5.4), so the
|
||
# seat chunks at the same size: bounded memory, any length, no API change for callers.
|
||
win = int(os.environ.get("WINDOW_S", "360")) * SR
|
||
if len(samples) <= win:
|
||
return _infer(samples)
|
||
# Windowed path: empty BEFORE each window too — torch's cap counts RESERVED bytes, and cached
|
||
# blocks from the previous window would otherwise count against it and bite spuriously.
|
||
torch.cuda.empty_cache()
|
||
try:
|
||
parts = []
|
||
for i in range(0, len(samples), win):
|
||
parts.append(_infer(samples[i:i + win]))
|
||
torch.cuda.empty_cache()
|
||
except torch.OutOfMemoryError as exc:
|
||
# Our cap (or a neighbour's pressure) bit: answer 503, give the cache back either way.
|
||
torch.cuda.empty_cache()
|
||
raise HTTPException(503, f"GPU memory limit reached for this file: {exc}") from exc
|
||
return " ".join(p for p in parts if p)
|
||
|
||
|
||
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-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())
|