diff --git a/nemotron-3-diarization/README.md b/nemotron-3-diarization/README.md new file mode 100644 index 00000000..2f155e75 --- /dev/null +++ b/nemotron-3-diarization/README.md @@ -0,0 +1,92 @@ +# NVIDIA Nemotron 3 Diarization on Baseten + +[NVIDIA Nemotron 3 Diarization](https://huggingface.co/nvidia/Nemotron-3-Diarization) answers +"who spoke when" in real-world audio: up to eight speakers, labels ordered by first arrival, and a +single checkpoint that runs at any of several algorithmic latencies, from a 0.32-second input +buffer for live agents to a 30-second buffer for recordings. It is the successor to Streaming +Sortformer v2.1 and is released under the [OpenMDW 1.1](https://openmdw.ai/license/1-1/) license, +which permits commercial use. + +Baseten serves the same model three ways. Each preset runs on one NVIDIA RTX PRO 6000 behind a +serving engine built around NVIDIA's NeMo streaming loop, so throughput and concurrency exceed +the reference script on the same GPU. + +| Preset | Transport | Use it for | Capacity per GPU | +|---|---|---|---| +| [`batch/`](batch/) | HTTP, one request per file | Recorded audio: a file URL or base64 in, speaker turns out | 200 six-minute files per minute, sustained | +| [`streaming/`](streaming/) | WebSocket, one connection per stream | Live audio: 100 ms PCM frames in, a live turn list on every chunk | 560 hour-long streams at the 1.04 s profile, 200 at 0.32 s | +| [`diarized-transcription-streaming/`](diarized-transcription-streaming/) | WebSocket | Live speaker-tagged words: the diarizer paired with NVIDIA's multitalker Parakeet 0.6B ASR (English) | 190 hour-long streams | + +A multilingual transcription pairing is in progress and will be added as a fourth preset. + +Model Library page: + +## Latency profiles + +One checkpoint, chosen per request (batch) or per connection (streaming). Input-buffer latency is +(chunk + right context) × 80 ms; compute runs far faster than real time, so the buffer dominates +the time from speech to label. + +| Profile | Input-buffer latency | Notes | +|---|---|---| +| `offline` | 30.4 s | Best accuracy. Default for batch. | +| `low` | 1.04 s | Real-time default. | +| `ultralow` | 0.32 s | Lowest latency; three times the step rate of `low`. | + +## Accuracy + +Diarization error rate with overlapping speech included and no collar, on identical audio and +references for every system. Lower is better. Measured on the General Access checkpoint. + +| Dataset | Nemotron 3 `offline` | `low` | `ultralow` | Streaming Sortformer v2.1 (`offline`) | +|---|---|---|---|---| +| NOTSOFAR-1 (129 meetings) | 15.8 | 17.5 | 18.8 | 23.8 | +| AMI-SDM (34 meetings) | 25.0 | 25.5 | 25.9 | 27.8 | +| CALLHOME (12 calls) | 13.0 | 13.8 | 14.0 | 16.7 | +| AISHELL-4, Mandarin (12 meetings) | 10.2 | 9.8 | 11.8 | 27.2 | + +Nemotron 3 at its fastest profile beats its predecessor at its slowest on every set. The gain is +largest on recordings with five or more speakers, where v2.1's four-speaker cap costs it. + +Real-time diarized transcription (English pairing) scores 28.6 cpWER on NOTSOFAR-1 eval-30 with the +canonical scorer. Words appear a median 0.7 s after they are spoken, speaker labels are final on +arrival, and committed words are never revised. + +## Deploying + +Deploy any preset in one click from the Baseten Model Library; each deploy link creates the model in +your workspace on an NVIDIA RTX PRO 6000: + +- Batch diarization — +- Streaming diarization — +- Real-time diarized transcription — + +The checkpoint (`nvidia/Nemotron-3-Diarization`, OpenMDW 1.1) is fetched at deploy time. + +## Calling the endpoints + +Each preset directory documents its full request and response contract and ships runnable +clients (`client.py`, and `curl.sh` for batch). All of them read two environment variables: + +```bash +export BASETEN_API_KEY=... # https://app.baseten.co/settings/account/api_keys +export MODEL_ID=... # from the model's page in the Baseten dashboard +``` + +URL forms: + +| Preset | Production URL | +|---|---| +| batch | `https://model-{MODEL_ID}.api.baseten.co/environments/production/predict` | +| streaming, diarized transcription | `wss://model-{MODEL_ID}.api.baseten.co/environments/production/websocket` | + +Replace `environments/production` with `deployment/{DEPLOYMENT_ID}` to target a specific +deployment. Authenticate with an `Authorization: Api-Key $BASETEN_API_KEY` header. + +Audio for the WebSocket presets is 16 kHz mono PCM16, base64-encoded inside JSON text frames +(100 ms frames, 3,200 bytes each, work well). The batch preset accepts any streamable format `ffmpeg` +decodes (WAV, FLAC, MP3, OGG, WebM) and resamples it server-side; see [`batch/README.md`](batch/README.md) +for the MP4/M4A and URL-fetch caveats. + +Every client and response in these directories was run against a live deployment of this code on +2026-09-22; the JSON shown is what came back, trimmed only for length. diff --git a/nemotron-3-diarization/batch/README.md b/nemotron-3-diarization/batch/README.md new file mode 100644 index 00000000..5117bfa6 --- /dev/null +++ b/nemotron-3-diarization/batch/README.md @@ -0,0 +1,139 @@ +# Nemotron 3 Diarization — batch (HTTP) + +Diarize a recording in one request: a file URL or base64 audio in, speaker turns out. The latency +profile is chosen per request. Independent requests arriving within ~100 ms are batched into one +GPU call server-side, so throughput scales with concurrency: one RTX PRO 6000 sustains 200 +six-minute files per minute. + +Deploy from the [Model Library](https://app.baseten.co/deploy/baseten/nemotron-3-diarization-batch); see the [parent README](../README.md) for +model details, accuracy and the other presets. + +## Endpoint + +``` +POST https://model-{MODEL_ID}.api.baseten.co/environments/production/predict +Authorization: Api-Key {BASETEN_API_KEY} +Content-Type: application/json +``` + +## Request + +```json +{ + "diarization_input": { + "audio": {"url": "https://github.com/ggerganov/whisper.cpp/raw/master/samples/jfk.wav"}, + "latency": "offline" + } +} +``` + +| Field | Type | Required | Description | +|---|---|---|---| +| `diarization_input` | object | yes | Wrapper object. A request without it is rejected with HTTP 400. | +| `diarization_input.audio` | object | yes | Exactly one of `url` or `audio_b64`. | +| `diarization_input.audio.url` | string | one of | Public or presigned URL. Fetched server-side with a 120 s timeout. | +| `diarization_input.audio.audio_b64` | string | one of | Base64 of the audio file bytes (the whole file, not raw PCM). Validated as base64. | +| `diarization_input.latency` | string | no | `offline` (default, 30.4 s buffer, best accuracy), `low` (1.04 s) or `ultralow` (0.32 s). | + +Audio may be any codec `ffmpeg` can decode from a stream (WAV, FLAC, MP3, OGG, WebM/Opus, …), any +sample rate, mono or stereo; it is decoded and resampled to 16 kHz mono on the server. See Errors +for the MP4/M4A caveat. Use FLAC or a +compressed format for `audio_b64` to stay under Baseten's request-size limit on long files. + +## Response + +Real response for the 11-second `jfk.wav` sample above, `offline` profile (server time 936 ms, of +which 527 ms was fetching the URL): + +```json +{ + "turns": [ + { + "start": 0.29, + "end": 2.31, + "speaker": "speaker_0" + }, + { + "start": 3.26, + "end": 4.56, + "speaker": "speaker_0" + }, + { + "start": 5.37, + "end": 10.63, + "speaker": "speaker_0" + } + ], + "segments": [ + "0.290 2.310 speaker_0", + "3.260 4.560 speaker_0", + "5.370 10.630 speaker_0" + ], + "num_speakers": 1, + "latency": "offline", + "batch_n": 1, + "timing": { + "fwd_ms": 7, + "h2d_ms": 0, + "mel_ms": 1, + "steps_ms": 5, + "graph_bs": 4, + "nonfinite_preds": 0, + "nonfinite_state": 0, + "pred_max": 1.0, + "pred_min": 0.0, + "core_dtype": "bfloat16", + "batch_ms": 7, + "batch_n": 1, + "batch_audio_s": 11.0, + "gap_ms": 1477, + "queue_ms": 401, + "decode_ms": 125, + "audio_s": 11.0, + "post_ms": 1, + "fetch_ms": 527, + "total_ms": 936 + } +} +``` + +| Field | Type | Description | +|---|---|---| +| `turns` | array | Speaker turns, ordered by start time. `start`/`end` in seconds from the beginning of the file; `speaker` is `speaker_0` … `speaker_7`, numbered in order of first appearance. Turns of different speakers may overlap. Zero-length turns are dropped. | +| `turns[].start`, `turns[].end` | number | Seconds, 10 ms resolution. | +| `turns[].speaker` | string | Session-local label, not an identity across files. | +| `segments` | array of strings | The same turns in NeMo's `"start end speaker"` text form. | +| `num_speakers` | integer | Number of distinct speakers found (1–8). | +| `latency` | string | The profile that ran. | +| `batch_n` | integer | How many concurrent requests shared this GPU call (1 when alone). | +| `timing` | object | Server-side profile of this request, milliseconds. The useful ones: `fetch_ms` (URL download or base64 decode), `decode_ms` (ffmpeg), `queue_ms` (wait for the batch window), `fwd_ms` (model), `post_ms` (turn extraction), `total_ms`; `audio_s` is the decoded duration. The remaining keys (`h2d_ms`, `mel_ms`, `steps_ms`, `graph_bs`, `gap_ms`, `batch_audio_s`, `core_dtype`, `pred_*`, `nonfinite_*`) are engine diagnostics and may change between releases. | + +## Errors + +Client errors return HTTP 400 with a JSON body `{"error": ""}`: + +| `error` | Cause | +|---|---| +| `request must be {'diarization_input': {'audio': {...}}}` | Missing wrapper or audio object. | +| `audio requires 'url' or 'audio_b64'` | Neither given. | +| `audio_b64 is not valid base64: …` | Bad encoding. | +| `could not fetch audio.url: HTTP Error 406: Not Acceptable` (or 403/404/timeout) | The host refused the server's plain `GET` (some hosts gate on `User-Agent`), the object does not exist, or the download exceeded 120 s. Use a presigned object-store URL, or send the file as `audio_b64`. | +| `audio could not be decoded: …` | `ffmpeg` rejected the bytes. Audio is decoded from a pipe, so **MP4/M4A/MOV containers whose index (`moov`) sits at the end of the file cannot be read** — remux them (`ffmpeg -movflags +faststart`) or send WAV/FLAC/MP3/OGG. | +| `latency must be one of ['offline', 'low', 'ultralow']` | Unknown profile. | + +Server faults return HTTP 500 and are logged with a traceback. + +## Limits + +- Up to 8 speakers per file; more speakers are merged into the nearest existing label. +- A request waits at most 30 minutes for a GPU slot and 10 minutes for post-processing before + failing; in practice a one-hour file returns in a few seconds. +- The server batches at most 32 files per GPU call and waits up to 1.5 s to fill a batch, so the + first request in a quiet period pays up to ~100 ms of batch-window latency. + +## Clients + +- [`client.py`](client.py) — `python client.py --url https://github.com/ggerganov/whisper.cpp/raw/master/samples/jfk.wav` or `python client.py --file local.flac` (base64), with `--latency`. +- [`curl.sh`](curl.sh) — `./curl.sh https://…/audio.wav [offline|low|ultralow]`. + +Both read `BASETEN_API_KEY` and `MODEL_ID` from the environment. diff --git a/nemotron-3-diarization/batch/client.py b/nemotron-3-diarization/batch/client.py new file mode 100644 index 00000000..21b5b9c8 --- /dev/null +++ b/nemotron-3-diarization/batch/client.py @@ -0,0 +1,60 @@ +"""Diarize a recording with the Nemotron 3 Diarization batch endpoint. + + export BASETEN_API_KEY=... MODEL_ID=... + python client.py --url https://example.com/meeting.wav + python client.py --file meeting.flac --latency low # base64 upload + +Prints one line per speaker turn: start, end, speaker. +""" + +import argparse +import base64 +import os +import sys + +import requests + + +def diarize(audio: dict, latency: str = "offline") -> dict: + url = f"https://model-{os.environ['MODEL_ID']}.api.baseten.co/environments/production/predict" + resp = requests.post( + url, + headers={"Authorization": f"Api-Key {os.environ['BASETEN_API_KEY']}"}, + json={"diarization_input": {"audio": audio, "latency": latency}}, + timeout=600, + ) + resp.raise_for_status() + return resp.json() + + +def main() -> None: + ap = argparse.ArgumentParser() + src = ap.add_mutually_exclusive_group(required=True) + src.add_argument("--url", help="public or presigned URL of the audio file") + src.add_argument( + "--file", help="local audio file, sent as base64 (any ffmpeg-decodable format)" + ) + ap.add_argument( + "--latency", default="offline", choices=["offline", "low", "ultralow"] + ) + args = ap.parse_args() + + if args.url: + audio = {"url": args.url} + else: + with open(args.file, "rb") as f: + audio = {"audio_b64": base64.b64encode(f.read()).decode()} + + result = diarize(audio, args.latency) + for turn in result["turns"]: + print(f"{turn['start']:8.2f} {turn['end']:8.2f} {turn['speaker']}") + t = result["timing"] + print( + f"\n{result['num_speakers']} speaker(s), {t['audio_s']:.1f} s of audio, " + f"profile={result['latency']}, server {t['total_ms']} ms", + file=sys.stderr, + ) + + +if __name__ == "__main__": + main() diff --git a/nemotron-3-diarization/batch/curl.sh b/nemotron-3-diarization/batch/curl.sh new file mode 100644 index 00000000..7e55d4fa --- /dev/null +++ b/nemotron-3-diarization/batch/curl.sh @@ -0,0 +1,11 @@ +#!/usr/bin/env bash +# Diarize a recording by URL with the Nemotron 3 Diarization batch endpoint. +# export BASETEN_API_KEY=... MODEL_ID=... +# ./curl.sh https://example.com/meeting.wav [offline|low|ultralow] +set -euo pipefail +URL=${1:?audio url}; LATENCY=${2:-offline} +curl -sS -X POST "https://model-${MODEL_ID}.api.baseten.co/environments/production/predict" \ + -H "Authorization: Api-Key ${BASETEN_API_KEY}" \ + -H "Content-Type: application/json" \ + -d "{\"diarization_input\": {\"audio\": {\"url\": \"${URL}\"}, \"latency\": \"${LATENCY}\"}}" +echo diff --git a/nemotron-3-diarization/diarized-transcription-streaming/README.md b/nemotron-3-diarization/diarized-transcription-streaming/README.md new file mode 100644 index 00000000..f3204e87 --- /dev/null +++ b/nemotron-3-diarization/diarized-transcription-streaming/README.md @@ -0,0 +1,171 @@ +# Nemotron 3 Diarized Transcription — streaming (WebSocket) + +Live speaker-attributed transcription: NVIDIA Nemotron 3 Diarization decides who is speaking on +every 80 ms frame, and NVIDIA's multitalker Parakeet 0.6B ASR (English) transcribes each active +speaker from the same mixed audio, conditioned on that speaker's activity. The output is +speaker-tagged words with timings, per-speaker live partials, and overlap flags. Speaker labels are +final the moment they appear, and a committed word is never revised. Words show a median 0.7 s +after they are spoken. One RTX PRO 6000 holds 190 concurrent hour-long streams. + +Deploy from the [Model Library](https://app.baseten.co/deploy/baseten/nemotron-3-diarized-transcription-streaming); see the +[parent README](../README.md). + +## Endpoint + +``` +wss://model-{MODEL_ID}.api.baseten.co/environments/production/websocket +Authorization: Api-Key {BASETEN_API_KEY} +``` + +All messages are JSON **text** frames; audio is base64-encoded inside JSON. + +## Client → server + +```json +{"session_id": "call-42", "max_speakers": 8} // optional handshake, first frame +{"type": "input_audio_buffer.append", "audio": ""} // repeat, 100 ms frames +{"type": "input_audio_buffer.commit"} // end of audio: flush, final frame, close +``` + +| Frame | Field | Type | Default | Description | +|---|---|---|---|---| +| handshake (optional) | `session_id` | string | random | Echoed in every server frame; use it to correlate logs. | +| | `max_speakers` | integer 1–8 | 8 | Caps the per-speaker ASR instances for this connection. Fewer speakers than expected is cheaper; more than the cap are merged into existing labels. | +| | `words` | 0 or 1 | 1 | Include per-word timings in `segments`. Set 0 to cut frame size by about 3×. | +| | `overlap` | 0 or 1 | 1 | Let turns of different speakers overlap in time (see below). Set 0 for a single running tail where simultaneous speech alternates word by word. | +| | `turn_segments` | 0 or 1 | 1 | Cut `segments` into turns (below). Set 0 for NeMo's native output: one running block per speaker, `partial` and `words` absent. | +| | `partials` | 0 or 1 | 1 | Send the `partial` array. | +| append | `type`, `audio` | | | `input_audio_buffer.append`; base64 of raw little-endian 16-bit mono PCM at 16 kHz. 100 ms frames (3,200 bytes) at real-time pace. | +| commit | `type` | | | `input_audio_buffer.commit`. Flushes, sends the final frame, closes. | + +## Server → client + +Real frame from streaming a single-speaker Harvard-sentences WAV (`words` trimmed to three per +segment for space; 31 frames in total for 34.7 s of audio): + +```json +{ + "type": "transcription", + "is_final": false, + "session_id": "demo", + "processed_s": 12.32, + "num_speakers": 1, + "segments": [ + { + "speaker": "speaker_0", + "start": 1.12, + "end": 8.08, + "text": "The birch canoe slid on the smooth planks, glue the sheet to the dark blue background.", + "overlap": false, + "overlaps_with": [], + "words": [ + { + "w": "The", + "start": 1.12, + "end": 1.2 + }, + { + "w": "birch", + "start": 1.28, + "end": 1.68 + }, + { + "w": "canoe", + "start": 1.84, + "end": 2.24 + } + ] + }, + { + "speaker": "speaker_0", + "start": 8.24, + "end": 11.28, + "text": "It is easy to tell the depth over well.", + "overlap": false, + "overlaps_with": [], + "words": [ + { + "w": "It", + "start": 8.24, + "end": 8.32 + }, + { + "w": "is", + "start": 8.4, + "end": 8.48 + }, + { + "w": "easy", + "start": 8.48, + "end": 8.64 + } + ] + } + ], + "partial": [ + { + "speaker": "speaker_0", + "text": "These days a chicken leg", + "start": 11.36, + "end": 12.24 + } + ] +} +``` + +The final frame carries every closed turn and an empty `partial`: + +```json +{"type": "transcription", "is_final": true, "session_id": "demo", "processed_s": 34.72, "num_speakers": 1} +``` + +| Field | Type | Description | +|---|---|---| +| `type` | string | `transcription`, or `error` (below). | +| `is_final` | boolean | `true` only on the last frame, after `commit`. | +| `session_id` | string | From the handshake, or generated. | +| `processed_s` | number | Seconds of audio the model has stepped through. | +| `num_speakers` | integer | Distinct speakers seen so far. | +| `segments` | array | **All closed turns so far**, ordered by start (replace your view). A closed turn is immutable apart from two things: a trailing word piece may be glued on if the word straddled a chunk boundary, and `overlap` may flip to `true` when a later-closing parallel turn touches it. | +| `segments[].speaker` | string | `speaker_0`…`speaker_7`, ordered by first arrival; session-local, not an identity. | +| `segments[].start`, `.end` | number | Seconds; word boundaries are on the ASR's 80 ms frame grid. | +| `segments[].text` | string | Punctuated, cased text of the turn. | +| `segments[].overlap` | boolean | This turn overlaps another speaker's turn in time. | +| `segments[].overlaps_with` | array of strings | Which speakers. Present when `overlap` is on. | +| `segments[].words` | array | `{"w", "start", "end"}` per word. Present when `words` is on. | +| `partial` | array | One entry per speaker **currently talking**: their open tail (`speaker`, `text`, `start`, `end`), re-sent every chunk (~1.1 s) until it closes and moves into `segments`. Empty on the final frame. | + +### How turns are cut + +Each speaker has its own open tail. It closes into a `segments` entry when that speaker pauses for +more than 1.2 s, at sentence-final punctuation, or when other speakers have talked for ≥ 1 s or +≥ 3 words since their last word (a hand-over). A one- or two-word backchannel never splits the +running turn; it becomes its own short overlapping segment. Because tails are per speaker, two +people talking at once produce two parallel turns rather than an alternation of fragments. + +A frame is sent on every ASR chunk (about every 1.12 s). Frames are replace-style, so a client +that falls behind can skip frames and lose nothing. + +## Errors + +One `{"type": "error", "error": "…"}` frame, then the socket closes: + +| `error` | Cause | +|---|---| +| `max_speakers must be in [1, 8]` / `max_speakers must be an integer` | Bad handshake value. | +| `frame must be a JSON object` / base64 decode error text | Malformed client frame. | +| `internal error: ` | Server fault, logged with a traceback. | + +## Limits + +- 8 speakers per connection; English only for this pairing (a multilingual pairing is in progress). +- Algorithmic latency 1.12 s (the ASR chunk); speaker activity is known ~1.04 s after the audio. +- A connection that sends no frames for 120 s is finalized and closed by the server. +- `predict_concurrency` on the deployment caps admitted connections per replica at the measured hold + (190 hour-long streams); excess connections wait at the gateway rather than degrading live ones. + +## Clients + +- [`client.py`](client.py) — `python client.py meeting.wav [--max-speakers N] [--no-words]`. Resamples any WAV to 16 kHz mono, streams at real-time pace, prints closed turns as they land and each speaker's live partial. + +Both read `BASETEN_API_KEY` and `MODEL_ID` from the environment. diff --git a/nemotron-3-diarization/diarized-transcription-streaming/client.py b/nemotron-3-diarization/diarized-transcription-streaming/client.py new file mode 100644 index 00000000..6371416f --- /dev/null +++ b/nemotron-3-diarization/diarized-transcription-streaming/client.py @@ -0,0 +1,104 @@ +"""Stream a WAV file to the Nemotron 3 diarized-transcription endpoint at real-time pace. + + export BASETEN_API_KEY=... MODEL_ID=... + pip install websockets numpy + python client.py meeting.wav [--max-speakers 8] [--no-words] + +Prints each closed, speaker-tagged turn once as it lands, and the live partial of every speaker +currently talking. Any WAV (mono/stereo, any rate) is converted to 16 kHz mono PCM16 locally. +""" + +import argparse +import asyncio +import base64 +import json +import os +import time +import wave + +import numpy as np +import websockets + +FRAME_MS = 100 +RATE = 16_000 +FRAME_BYTES = RATE * FRAME_MS // 1000 * 2 + + +def load_pcm16(path: str) -> bytes: + with wave.open(path, "rb") as w: + n_ch, width, rate = w.getnchannels(), w.getsampwidth(), w.getframerate() + raw = w.readframes(w.getnframes()) + if width != 2: + raise SystemExit("expected 16-bit PCM WAV") + x = np.frombuffer(raw, dtype=np.int16).reshape(-1, n_ch).mean(axis=1) + if rate != RATE: + x = np.interp(np.arange(0, len(x), rate / RATE), np.arange(len(x)), x) + return x.astype(np.int16).tobytes() + + +async def transcribe(pcm16: bytes, max_speakers: int, words: bool) -> dict: + url = f"wss://model-{os.environ['MODEL_ID']}.api.baseten.co/environments/production/websocket" + headers = {"Authorization": f"Api-Key {os.environ['BASETEN_API_KEY']}"} + async with websockets.connect(url, additional_headers=headers, max_size=None) as ws: + await ws.send( + json.dumps( + { + "session_id": f"demo-{int(time.time())}", + "max_speakers": max_speakers, + "words": int(words), + } + ) + ) + + async def send_audio(): + for i in range(0, len(pcm16), FRAME_BYTES): + frame = base64.b64encode(pcm16[i : i + FRAME_BYTES]).decode() + await ws.send( + json.dumps({"type": "input_audio_buffer.append", "audio": frame}) + ) + await asyncio.sleep(FRAME_MS / 1000) + await ws.send(json.dumps({"type": "input_audio_buffer.commit"})) + + sender = asyncio.create_task(send_audio()) + printed = 0 + last = {} + async for msg in ws: + frame = json.loads(msg) + if frame.get("type") == "error": + raise SystemExit(f"server error: {frame['error']}") + last = frame + for seg in frame["segments"][ + printed: + ]: # closed turns are stable: print each once + flag = " (overlap)" if seg.get("overlap") else "" + print( + f"[{seg['speaker']} {seg['start']:6.2f}-{seg['end']:6.2f}]{flag} {seg['text']}" + ) + printed = len(frame["segments"]) + for p in frame.get("partial", []): + print(f" … {p['speaker']}: {p['text']}") + if frame.get("is_final"): + break + await sender + return last + + +def main() -> None: + ap = argparse.ArgumentParser() + ap.add_argument("wav") + ap.add_argument("--max-speakers", type=int, default=8) + ap.add_argument( + "--no-words", action="store_true", help="omit per-word timings (smaller frames)" + ) + args = ap.parse_args() + final = asyncio.run( + transcribe(load_pcm16(args.wav), args.max_speakers, not args.no_words) + ) + print( + f"\nfinal: {final['num_speakers']} speaker(s), {len(final['segments'])} turns, " + f"{final['processed_s']:.1f} s processed" + ) + + +if __name__ == "__main__": + main() diff --git a/nemotron-3-diarization/streaming/README.md b/nemotron-3-diarization/streaming/README.md new file mode 100644 index 00000000..3787050a --- /dev/null +++ b/nemotron-3-diarization/streaming/README.md @@ -0,0 +1,90 @@ +# Nemotron 3 Diarization — streaming (WebSocket) + +Live speaker diarization over a WebSocket. Stream 16 kHz mono PCM16 audio in 100 ms frames; the +server drives NVIDIA's streaming Sortformer loop with per-connection state and sends back the +current speaker-turn list on every processed chunk. Speaker identity is carried forward across the +connection with no re-clustering, so a label never changes once it appears. One RTX PRO 6000 holds +560 concurrent hour-long streams at the `low` profile, or 200 at `ultralow`. + +Deploy from the [Model Library](https://app.baseten.co/deploy/baseten/nemotron-3-diarization-streaming); see the [parent README](../README.md). + +## Endpoint + +``` +wss://model-{MODEL_ID}.api.baseten.co/environments/production/websocket +Authorization: Api-Key {BASETEN_API_KEY} +``` + +All messages are JSON **text** frames. Audio travels base64-encoded inside JSON, not as binary +frames. + +## Client → server + +```json +{"latency": "low"} // optional handshake, first frame +{"type": "input_audio_buffer.append", "audio": ""} // repeat +{"type": "input_audio_buffer.commit"} // end of audio: flush, final frame, close +``` + +| Frame | Field | Type | Description | +|---|---|---|---| +| handshake (optional, must be first) | `latency` | string | `low` (default, 1.04 s buffer), `ultralow` (0.32 s) or `offline` (30.4 s). Fixed for the connection. | +| | `threshold` | number | Per-speaker activity threshold in (0, 1); default 0.5. Lower finds more speech and more overlap, higher is stricter. | +| append | `type` | string | `input_audio_buffer.append` | +| | `audio` | string | Base64 of raw little-endian 16-bit mono PCM at 16 kHz. Any frame size works; 100 ms (3,200 bytes) is a good default. Send at real-time pace for a live source. | +| commit | `type` | string | `input_audio_buffer.commit`. The server processes the remaining audio, sends the final frame and closes the socket. | + +The handshake may be omitted; the first `append` then starts a `low` session with the default +threshold. The handshake keys may also ride on the first `append` frame. + +## Server → client + +Real frames from streaming a 33.6 s single-speaker WAV at the `low` profile (47 frames in total, +one every 0.72 s of audio): + +```json +{"type": "diarization", "is_final": false, "processed_s": 8.64, "num_speakers": 1, "turns": [{"start": 0.45, "end": 3.19, "speaker": "speaker_0"}, {"start": 4.16, "end": 6.44, "speaker": "speaker_0"}, {"start": 7.81, "end": 8.64, "speaker": "speaker_0"}]} +``` +```json +{"type": "diarization", "is_final": true, "processed_s": 33.6, "num_speakers": 1, "turns": [{"start": 0.45, "end": 3.19, "speaker": "speaker_0"}, {"start": 4.16, "end": 6.44, "speaker": "speaker_0"}, …], "is_end_of_audio_flush": true} +``` + +| Field | Type | Description | +|---|---|---| +| `type` | string | `diarization`, or `error` (below). | +| `is_final` | boolean | `true` only on the last frame, sent after `commit`. | +| `is_end_of_audio_flush` | boolean | Present and `true` on the final frame: the trailing buffer was flushed. | +| `processed_s` | number | Seconds of audio the model has stepped through so far. | +| `num_speakers` | integer | Distinct speakers seen so far (1–8). | +| `turns` | array | **The complete current turn list** — replace your view with it, do not append. Each turn: `start`, `end` (seconds, 10 ms resolution), `speaker` (`speaker_0`…`speaker_7`, ordered by first arrival). Closed turns are stable; the last turn of a speaker who is still talking grows on each frame. Turns of different speakers may overlap. | + +A frame is sent for every chunk the model steps: every 0.72 s of audio at `low` (9 frames of 80 ms), +0.24 s at `ultralow`, 27 s at `offline`. Frames are replace-style, so a client that falls behind can drop intermediate +frames and lose nothing. + +## Errors + +The server sends one `{"type": "error", "error": "…"}` frame and closes the socket for: + +| `error` | Cause | +|---|---| +| `latency must be one of ['offline', 'low', 'ultralow']` | Unknown profile in the handshake (verified: `{"type": "error", "error": "latency must be one of ['offline', 'low', 'ultralow']"}` then close). | +| `threshold must be in (0, 1)` | Out-of-range threshold. | +| `frame must be a JSON object` / base64 decode error text | Malformed client frame. | +| `capacity: this replica is at its GPU budget for low sessions (…); retry on another replica` | Admission control: the replica is full. Reconnect; the platform routes to another replica when one exists. | +| `internal error: ` | Server fault, logged with a traceback. | + +## Limits + +- 8 speakers per connection. +- One profile per connection; open a new connection to change it. +- Admission: a replica admits sessions while `low + 2.8 × ultralow + 0.2 × offline ≤ 560`; beyond + that it returns the capacity error instead of degrading every stream. +- Idle connections (no frames) are closed by the platform gateway after its ping timeout; send + `commit` when you are done rather than abandoning the socket. + +## Clients + +- [`client.py`](client.py) — `python client.py meeting.wav --latency low`. Resamples any WAV to 16 kHz mono, streams it at real-time pace, prints the turn list as it grows. + +Both read `BASETEN_API_KEY` and `MODEL_ID` from the environment. diff --git a/nemotron-3-diarization/streaming/client.py b/nemotron-3-diarization/streaming/client.py new file mode 100644 index 00000000..5746eb1e --- /dev/null +++ b/nemotron-3-diarization/streaming/client.py @@ -0,0 +1,88 @@ +"""Stream a WAV file to the Nemotron 3 Diarization streaming endpoint at real-time pace. + + export BASETEN_API_KEY=... MODEL_ID=... + pip install websockets numpy + python client.py meeting.wav --latency low + +Any WAV (mono/stereo, any rate) is converted to 16 kHz mono PCM16 locally. Each server frame +carries the full current turn list; this client prints the newest turns as they change. +""" + +import argparse +import asyncio +import base64 +import json +import os +import wave + +import numpy as np +import websockets + +FRAME_MS = 100 +RATE = 16_000 +FRAME_BYTES = RATE * FRAME_MS // 1000 * 2 # PCM16 mono + + +def load_pcm16(path: str) -> bytes: + with wave.open(path, "rb") as w: + n_ch, width, rate = w.getnchannels(), w.getsampwidth(), w.getframerate() + raw = w.readframes(w.getnframes()) + if width != 2: + raise SystemExit("expected 16-bit PCM WAV") + x = np.frombuffer(raw, dtype=np.int16).reshape(-1, n_ch).mean(axis=1) + if ( + rate != RATE + ): # linear resample; good enough for a demo, use soxr/ffmpeg in production + t_new = np.arange(0, len(x), rate / RATE) + x = np.interp(t_new, np.arange(len(x)), x) + return x.astype(np.int16).tobytes() + + +async def stream(pcm16: bytes, latency: str) -> dict: + url = f"wss://model-{os.environ['MODEL_ID']}.api.baseten.co/environments/production/websocket" + headers = {"Authorization": f"Api-Key {os.environ['BASETEN_API_KEY']}"} + async with websockets.connect(url, additional_headers=headers, max_size=None) as ws: + await ws.send(json.dumps({"latency": latency})) + + async def send_audio(): + for i in range(0, len(pcm16), FRAME_BYTES): + frame = base64.b64encode(pcm16[i : i + FRAME_BYTES]).decode() + await ws.send( + json.dumps({"type": "input_audio_buffer.append", "audio": frame}) + ) + await asyncio.sleep(FRAME_MS / 1000) # real-time pacing + await ws.send(json.dumps({"type": "input_audio_buffer.commit"})) + + sender = asyncio.create_task(send_audio()) + last = {} + async for msg in ws: + frame = json.loads(msg) + if frame.get("type") == "error": + raise SystemExit(f"server error: {frame['error']}") + last = frame + tail = frame["turns"][-2:] + print( + f"t={frame['processed_s']:6.1f}s speakers={frame['num_speakers']} " + + " ".join( + f"{t['speaker']}[{t['start']:.1f}-{t['end']:.1f}]" for t in tail + ) + ) + if frame.get("is_final"): + break + await sender + return last + + +def main() -> None: + ap = argparse.ArgumentParser() + ap.add_argument("wav") + ap.add_argument("--latency", default="low", choices=["low", "ultralow", "offline"]) + args = ap.parse_args() + final = asyncio.run(stream(load_pcm16(args.wav), args.latency)) + print(f"\nfinal: {final['num_speakers']} speaker(s), {len(final['turns'])} turns") + for t in final["turns"]: + print(f"{t['start']:8.2f} {t['end']:8.2f} {t['speaker']}") + + +if __name__ == "__main__": + main()