Nemotron-3-Diarization — LiteRT (CompiledModel GPU)
NVIDIA's Nemotron-3-Diarization (streaming Sortformer, 100M parameters, up to 8 speakers) as LiteRT graphs: who spoke when, one decision per 10 ms frame. On a Galaxy S26 the whole network runs on the LiteRT CompiledModel GPU (2,915 of 2,915 nodes, one partition) and handles audio arriving in real time in 171 ms per 0.72 s step, with FP32 speaker decisions identical to the reference. Two graphs do the neural work; the streaming state (speaker cache + FIFO) is host code, included here in Python and Kotlin.
Input: "The Mad Tea-Party", chapter 7 of Alice's Adventures in Wonderland*, LibriVox dramatic reading
(archive.org, public domain, Public Domain Mark 1.0),
an 87.75 s excerpt starting at 1:39.70 of the chapter file (the audio is not part of this repo). Output of
conversion/nemotron3_diar_litert.py with the files in this repo: streaming low_latency mode, CPU, FP32; 5
speakers, 42 segments; colors as in the Android sample.*
assets/demo.mp4 (78 s): the Android sample's Play live mode on a Galaxy S26, graph B at FP32. The audio is a synthetic
eight-person meeting (Kokoro-82M voices, fictional names and a fictional company; the audio file is not part of this
repo). The timeline is drawn as each 0.72 s step completes, 1.04 s of audio plus about 0.16 s of compute behind the
sound.
Contents
| file | role | size |
|---|---|---|
nemotron3_diar_frontend.tflite |
graph A: 8-frame mel stacking + projection (float32) | 2.1 MB |
nemotron3_diar_encoder_low_latency_fp16.tflite |
graph B, streaming: 31-layer encoder + output head, fixed T = 541 rows (float16 weights, float compute) | 198.7 MB |
nemotron3_diar_encoder_offline_fp16.tflite |
graph B, whole file: same network, fixed T = 684 rows (float16 weights) | 198.7 MB |
assets/frontend_mel128_257.bin |
slaney mel filter bank [128, 257], float32 little-endian | 131.6 KB |
assets/hann400.bin |
Hann window, 400 samples, symmetric, float32 | 1.6 KB |
assets/silence_embeds.bin |
learned silence embedding [512] that fills empty speaker-cache slots, float32 | 2.0 KB |
conversion/ |
reference host nemotron3_diar_litert.py, conversion and verification scripts (conversion/README.md) |
|
android/ |
Android sample app source (Kotlin) | |
LICENSE, NOTICE, SHA256SUMS |
license agreement, origin and changes, checksums |
All of the checkpoint's weights are bfloat16 values; storing them as float16 changes 0.0055 % of them, by at most 3.0e-8.
Streaming profile
One encoder frame = 8 mel frames = 80 ms. Graph B sees [speaker cache | FIFO | chunk + look-ahead] and scores
every row; the host keeps the chunk's own rows.
low_latency (streaming) |
offline (whole file) |
|
|---|---|---|
| chunk + look-ahead (encoder frames) | 9 + 4 | 340 + 40 |
| audio before the first decision | 1.04 s (16,680 samples) | the whole file |
| step | 0.72 s | 27.2 s |
| speaker cache / FIFO (encoder frames) | 264 / 264 | 264 / 40 |
| cache update period | 222 | 300 |
| graph B rows T | 541 = 264 + 264 + 13 | 684 = 264 + 40 + 380 |
Speaker-cache constants (both modes, from the upstream config): 1 silence slot per speaker, score threshold 0.25, minimum positive-score rate 0.5, strong / weak boost rates 0.75 / 1.5, newest-frame boost 0.05.
Graph interfaces
All tensors float32. Signature names as below; read the output by name.
| graph | input | shape | output | shape |
|---|---|---|---|---|
A nemotron3_diar_frontend |
mel: log-mel of one chunk, zero rows after the last frame |
[1, 104, 128] | chunk_embeds |
[1, 13, 512] |
B …_low_latency_fp16 |
packed_embeds: L real rows, then zero rows |
[1, 541, 512] | logits (first 8·L rows real) |
[1, 4328, 8] |
attn_bias: added to every attention score |
[1, 1, 1, 541] | |||
rope_cos, rope_sin: rotary tables for positions 0 … T−1 |
[1, 1, 541, 64] | |||
B …_offline_fp16 |
the same four inputs with T = 684 | logits |
[1, 5472, 8] |
attn_bias, low_latency: 0 on the L real rows, −30000 on the zero rows. Offline: 0 real, −16384 on a real row whose key is masked (the frame after the last full hop of a file), −32768 on the zero rows. The graph derives its row mask from these values, so the zero rows never reach the output convolution.- The rotary tables are inputs, computed once by the host (they are the same every step): the graph holds no large position table as a constant, a pattern that has miscomputed on the ML Drift GPU path in other models.
- Speaker probabilities are
sigmoid(logits); the host does it.
Streaming loop (host)
conversion/nemotron3_diar_litert.py (numpy) and android/…/Nemotron3Diarizer.kt + SpeakerCache.kt +
MelFrontend.kt implement this loop; both follow transformers' Nemotron3DiarizationProcessor and
Nemotron3DiarizationSpeakerCache.
- Log-mel on the continuous 16 kHz stream: pre-emphasis 0.97 (first sample kept), frame i = samples [160i − 256, 160i + 256) (zero outside the stream), the 400-sample Hann window centered in 512, float32 real FFT, |X|², the mel bank, log(x + 2⁻²⁴). No normalization.
- Schedule (low_latency): the first chunk is mel frames [0, 104) once 16,680 samples have arrived; chunk k ≥ 1 is frames [72k, 72k + 104) once the stream reaches 72k·160 − 256 + 17,040 samples; when the audio ends, the remaining frames form the last chunk, with no look-ahead.
- Graph A turns the chunk's 104 mel frames into 13 rows (9 chunk + 4 look-ahead).
- Pack
[cache rows | FIFO rows | 13 new rows](L ≤ 541) plusattn_bias, run graph B. - Emit the chunk's logit rows
8·(cache + FIFO) … 8·(cache + FIFO + 9). - Update the state: sigmoid, then the mean of every 8 rows (one probability row per encoder frame). The 9 chunk
rows join the FIFO; when the FIFO exceeds 264 rows, max(222, overflow) of its oldest rows move to the cache.
When the cache exceeds 264 rows it is compressed: each frame gets a per-speaker score (log-likelihood of that
speaker alone; −∞ for frames without that speaker's speech, and for weak frames of speakers with ≥ 16 positive
ones), the candidates after the first 264 (the newest) get +0.05, each speaker's top 24 frames get +2·ln 2 and
top 48 +ln 2, and the 264 best (speaker, frame) pairs, with one silence slot per speaker, are kept in speaker
order; silence slots take
silence_embeds.
Float summation order matters only where two frames' scores tie exactly: both hosts sum the 8 per-speaker terms in
the reference's order ((s, s+4) pairs) and pick the lower frame index on equal scores; on the test clips every cache
selection equals the reference's (next sections). The offline mode (run_file / runFile) computes the whole
file's mel (centered frames 0 … N/160, the last one masked), runs graph A over 104-frame blocks and graph B over
340-frame chunks with the same state update (FIFO 40, period 300).
On-device results
Galaxy S26 (SM-S942Q, Snapdragon SM8850, Adreno GPU), Android 16, LiteRT 2.2.0, CompiledModel GPU: graph A 2 / 2
nodes, graph B 2,915 / 2,915 (offline 2,919 / 2,919) on LITERT_CL, one partition each, no CPU fallback. Graph A
always runs with GPU precision FP32; graph B with the precision in the first column. The 97.6 s example clip of the
upstream card, pushed 0.1 s at a time at the audio rate like a microphone; one step = log-mel + graph A + graph B
(write → run → read) + state update. Latency = from the push that completes a chunk to its logits. Reference =
transformers FP32 streaming on the same clip (78,072 frame × speaker cells, 37 segments, 4 cache compressions).
| graph B precision | latency per step, median / p95 | RTF | agreement @ 0.5 | segments vs reference | cache selections | compile A / B |
|---|---|---|---|---|---|---|
| FP32 | 170.7 / 175.5 ms | 0.238 | 100 % (0 flips) | 37 / 37 identical | 4 / 4 identical | 72 / 1,477 ms |
| FP16 (GPU default) | 115.1 / 120.0 ms | 0.160 | 99.973 % (21 flips) | 40: 2 added, 1 split, 8 boundaries moved by 10 ms | 4 / 4 differ (3–17 of 264 frames) | 73 / 1,434 ms |
| FP16, FP32 accumulation | 140.3 / 144.7 ms | 0.196 | 99.997 % (2 flips) | 38: 1 split, 1 boundary moved by 10 ms | 2 / 4 differ (1 frame each) | 71 / 1,382 ms |
- The first decision needs 1.04 s of audio plus 191 ms (FP32) / 133 ms (FP16) / 154 ms (FP16 + FP32 accumulation) of compute, measured after one warm-up inference at start-up. Compile = load + GPU compile at start-up.
- Step breakdown, FP32, median: log-mel 15.4 ms, graph A 11.3 ms, graph B 136.2 ms, state update 7.5 ms.
- Offline file mode, the same 97.6 s clip (graph B T = 684, 4 chunks), second and third of three passes in one launch: FP32 0.99–1.00 s (RTF 0.010), segments identical to transformers' offline forward (29 / 29, 0 flips); FP16 0.64–0.65 s (RTF 0.0066), 7 flips, 7 boundaries moved by 10 ms. Compile A / B: 71 / 1,444 ms (FP32).
- Conditions: screen on, unlocked, USB power, battery 96–98 %, 35.8–37.9 °C, Android thermal status 0 at the start of every run; the streaming runs in one session with 60 s between them, the file-mode runs 30 s apart.
- Back to back (feeding a file through the streaming path as fast as it runs) the GPU heats up: after about 8 s its thermal governor lowers the clock step by step (1300 → 500 MHz) and graph B goes from 135 ms (first 10 steps) to 288 ms (last 10 steps) (FP32, RTF 0.316). At real-time pace graph B stayed at 136 ms. For files, use the offline mode.
FP32 is the recommended setting: it is the only one whose segments and cache contents equal the reference, and at 171 ms per 720 ms step it runs at 0.24× real time. FP16 is faster (115 ms) but moves segment boundaries: its per-step differences are small (7 single steps fed the reference's inputs: max |Δp| 0.035, 1 flip in 155,072 cells), yet they change which frames the speaker cache keeps, and later decisions follow.
Validation
Reference: transformers Nemotron3DiarizationForAudioFrameClassification, FP32, commit
4b28d51d0d5f17ec20c23a187d0475a8e68810c8; transformers' integration test checks that model against NeMo's
probabilities within atol 1e-3 (the last encoder frame's 8 output frames excepted). Clips: the upstream example
(97.6 s) and a 21.5 s multi-speaker clip (not distributed).
- This repo's files through
conversion/nemotron3_diar_litert.py, Mac CPU (XNNPACK, 4 threads):
| mode | clip | max |Δlogit| | max |Δp| | flips | segments | cache state and selections |
|---|---|---|---|---|---|---|
| low_latency | 97.6 s | 6.9e-5 | 6.1e-6 | 0 / 78,072 | 37 / 37 identical | equal at all 136 steps |
| low_latency | 21.5 s | 6.5e-5 | 2.8e-6 | 0 / 17,192 | 10 / 10 identical | equal at all 30 steps |
| offline | 97.6 s | 4.6e-5 | 4.6e-6 | 0 / 78,088 | 29 / 29 identical | equal in all 4 chunks |
| offline | 21.5 s | 4.6e-5 | 2.3e-6 | 0 / 17,208 | 9 / 9 identical | equal |
- Same files on the Mac GPU (Metal,
--accelerator gpu): FP32 identical segments in all four runs (max |Δlogit| ≤ 2.3e-4); FP16 moves them (97.6 s streaming: 149 flips, 44 segments instead of 37). - Galaxy S26, closed loop (the table above): FP32 max |Δlogit| 1.8e-4, max |Δp| 1.5e-5, cache contents equal to the reference at all 136 steps; the 21.5 s clip: 0 flips, 10 / 10 segments identical.
- Galaxy S26, one step (7 steps of the 97.6 s clip fed with the reference's own inputs): FP32 max |Δlogit| 5.3e-5; FP16 0.343 (max |Δp| 0.035, 1 flip in 155,072 cells, the chunk's output rows 100 %); FP16 + FP32 accumulation 0.093 (2 flips).
- Kotlin host on the JVM, replaying the reference's graph outputs: log-mel max |Δ| 1.9e-6 (log domain), packed inputs bit-identical, cache state equal at all steps, emitted logits bit-identical to the reference.
- Graph checks: no GATHER / TOPK / CAST / WHERE-type ops and no tensor above rank 4 (2,915 ops; offline 2,919); native GELU (erf), attention as rank-4 matmul + softmax.
Minimal usage
Python
# pip install ai-edge-litert==2.2.0 numpy soundfile
# hf download litert-community/Nemotron-3-Diarization-LiteRT --local-dir n3d
# ffmpeg -i aliceinwonderland_07_carroll_64kb.mp3 -ss 99.70 -t 87.75 -ac 1 -ar 16000 tea_party_16k.wav
import sys
import numpy as np
sys.path.insert(0, "n3d/conversion")
from nemotron3_diar_litert import Nemotron3Diarizer, load_wav, speaker_segments
diarizer = Nemotron3Diarizer("n3d", mode="low_latency") # graph A + graph B on the LiteRT CompiledModel (CPU)
audio = load_wav("tea_party_16k.wav") # 16 kHz mono float32
logits = []
for i in range(0, len(audio), 1600): # 0.1 s pushes, as a microphone delivers them
for step in diarizer.push(audio[i : i + 1600]): # one step per 0.72 s of audio
logits.append(step.logits) # [frames, 8]: one row per 10 ms, one column per speaker
logits += [step.logits for step in diarizer.finish()] # the rest of the audio, no look-ahead
segments = speaker_segments(np.concatenate(logits)) # sigmoid > 0.5, speakers in order of arrival
for s in segments[:6]:
print(f"speaker_{s['Speaker']}: {s['Start']:.2f}s - {s['End']:.2f}s")
print(len(segments), "segments,", len({s["Speaker"] for s in segments}), "speakers")
Output (the hero clip, Mac CPU):
speaker_0: 0.17s - 3.97s
speaker_1: 4.38s - 7.20s
speaker_2: 7.57s - 9.84s
speaker_0: 10.33s - 11.15s
speaker_2: 11.59s - 12.96s
speaker_1: 13.71s - 13.73s
42 segments, 5 speakers
Whole file at once: Nemotron3Diarizer("n3d", mode="offline").run_file(audio). Command line:
python n3d/conversion/nemotron3_diar_litert.py audio_16k.wav [--mode offline] [--accelerator gpu]. Inside,
each graph is a CompiledModel.from_file(...) with buffers by signature name (Graph in the same file).
Kotlin (Android)
The classes are in android/app/src/main/java/com/nemotron3diar/; the models go to filesDir/models
(conversion/install_to_device.sh), the three tables ship as APK assets.
// implementation("com.google.ai.edge.litert:litert:2.2.0")
val env = Environment.create() // one Environment for both graphs
// One graph on the GPU, as LiteRtEngine does it: graph B at precision FP32 (graph A is always FP32)
val options = CompiledModel.Options(Accelerator.GPU)
options.gpuOptions = CompiledModel.GpuOptions(precision = CompiledModel.GpuOptions.Precision.FP32)
val file = File(filesDir, "models/nemotron3_diar_encoder_low_latency_fp16.tflite")
val model = CompiledModel.create(file.absolutePath, options, env)
val ins = listOf("packed_embeds", "attn_bias", "rope_cos", "rope_sin").associateWith { model.createInputBuffer(it) }
val outs = listOf("logits").associateWith { model.createOutputBuffer(it) }
val (cos, sin) = Nemotron3Diarizer.ropeTables(541)
ins.getValue("packed_embeds").writeFloat(packed) // [541 x 512]: L rows, then zero rows
ins.getValue("attn_bias").writeFloat(bias) // [541]: 0 real rows, -3e4 zero rows
ins.getValue("rope_cos").writeFloat(cos)
ins.getValue("rope_sin").writeFloat(sin)
model.run(ins, outs)
val logits = outs.getValue("logits").readFloat() // [4328 x 8]
// The streaming loop, as the sample's MainActivity runs it (LiteRtEngine = graph A + graph B as above)
val engine = LiteRtEngine(env, File(filesDir, "models"), precision = LiteRtEngine.Precision.FP32)
val d = Nemotron3Diarizer(
engine,
MelFrontend(StepLog.floatsFromStream(assets.open("frontend_mel128_257.bin")),
StepLog.floatsFromStream(assets.open("hann400.bin"))),
StepLog.floatsFromStream(assets.open("silence_embeds.bin")),
StreamConfig.LOW_LATENCY,
)
val rec = AudioRecord(MediaRecorder.AudioSource.MIC, 16000, AudioFormat.CHANNEL_IN_MONO,
AudioFormat.ENCODING_PCM_FLOAT, 16000 * 4 * 4)
val buf = FloatArray(1600) // 0.1 s pushes
rec.startRecording()
while (recording) {
val r = rec.read(buf, 0, buf.size, AudioRecord.READ_BLOCKING)
for (step in d.push(buf, 0, r)) timeline.append(step.logits, step.numFrames) // [numFrames x 8] (UI thread)
}
for (step in d.finish()) timeline.append(step.logits, step.numFrames)
sigmoid(logit) > 0.5 marks a speaker as active in that 10 ms frame (TimelineView.append); speaker k is the k-th
voice to appear.
Android sample
android/ is the sample app: Record (microphone, up to 5 min), Pick clip, or Play live (plays a WAV through the
speaker while diarizing it at audio rate, the mode in the video above), a per-speaker timeline that grows every
0.72 s step, graph B precision switch (FP32 default / FP16 / FP16 + FP32 accumulation), per-step ms and real-time
factor. ClosedLoopTest.kt and SelfTest.kt are the device checks behind the numbers above (started through
files/selftest.json; see conversion/README.md). LiteRT 2.2.0, AGP 8.9.1, Kotlin 2.2.21, minSdk 26, arm64-v8a.
cd android && gradle wrapper --gradle-version 8.11.1 && ./gradlew :app:installDebug # or open android/ in Android Studio
cd .. && conversion/install_to_device.sh . nemotron3_diar_frontend.tflite nemotron3_diar_encoder_low_latency_fp16.tflite
adb shell am start -n com.nemotron3diar/.MainActivity
Conversion notes
The network was re-authored in plain PyTorch (conversion/nemotron3diar_model.py) with the checkpoint loaded
unchanged, exported with litert-torch 0.9.4, and the weights cast to float16 with ai-edge-quantizer.
conversion/README.md has the full procedure.
LayerNorm in float16. The LayerNorm inputs reach |x| ≈ 956 (final norm) and 804 (layer 30), so the plain (x − μ)² reaches about 9·10⁵, beyond float16's 65,504. At the S26 GPU's default precision a plain-LayerNorm graph B ran without errors but returned wrong values (7 test steps: max |Δlogit| 38.2, logit correlation down to 0.0006). All 64 LayerNorms use a scaled form that keeps every intermediate within O(max |x|):
def safe_layer_norm(x, weight, bias, eps=1e-5): s = (x.abs().amax(-1, keepdim=True) * 0.125).clamp(min=1.0) # per row xs = x / s d = xs - xs.mean(-1, keepdim=True) var = (d * d).mean(-1, keepdim=True) # down-scaled variance, never multiplied back by s^2 return d * torch.rsqrt(var + eps / (s * s)) * weight + biasThe eps is divided by s²: adding eps in the down-scaled domain equals eps·s² in the original units, which moved the FP32 logits by up to 3.0e-3 here; with eps / s² the FP32 export stays within 7.1e-5 of the reference.
Row mask. The output head starts with a k = 3 convolution over time; the reference convolves exactly L rows. In the fixed-T graph the row after the last real one would leak into it, so the graph multiplies the projection by
relu(attn_bias + 1)(1 on real rows, 0 on zero rows). Without it the last real row's logits moved by up to 13.2.Three bias levels offline. The offline pass masks the key of the frame after the last full hop but still runs that frame through the head. The offline graph reads 0 / −16384 / −32768 and builds the mask as relu(y) − relu(y − 1) with y = bias·2⁻¹⁴ + 2, exact in float16 and without RELU_0_TO_1. Feeding that frame as a zero row instead moved the logits by 11.4.
The FFT of the host mel. The reference's torch.stft rounds quiet mel bins (energy near the 2⁻²⁴ guard) differently from a textbook FFT: a float32 radix-2 FFT was 2.7e-4 (log domain) from the processor, even a float64 FFT 1.8e-4. The Kotlin host ports pocketfft's real FFT (factors 2·4·4·4·4, twiddle products with a single rounding) and matches
torch.fft.rffton all 65,792 values tested; the Python host uses numpy's float32 path (rfft(norm="forward") * 512): 1.9e-6 from the processor.Graph A in FP32. Its rows stay in the speaker cache and FIFO for the whole session. At the GPU's default precision they differ from the reference by up to 0.38 (|x| ≤ 141); FP32 precision: 2.0e-4, for 0.2 ms.
float16 storage. For the plain-LayerNorm graph the converter folded LayerNorm γ into the next Linear, which made the weights non-bfloat16 and float16 storage lossy (max |Δlogit| 8.3e-3); the scaled LayerNorm is not folded, and the float16 files equal the float32 exports within 1e-7 per weight.
Native GELU (erf) is kept; attention is written as rank-4 matmul + softmax (no SDPA op, no KV cache).
Limitations
- FP16 GPU precision changes which frames the speaker cache keeps and moves segment boundaries (numbers above). FP32 matches the reference on the test clips.
- Two latency modes are exported:
low_latency(1.04 s) andoffline. The upstreamvery_low_latency(0.64 s) andultra_low_latency(0.32 s) modes need graph B builds with their own T; not included. - Checked on two clips (97.6 s, 21.5 s) against transformers; diarization error rate was not measured here (see the upstream card).
- Measured on the Galaxy S26 GPU (Adreno), Mac CPU and Mac GPU (Metal); other phones and GPUs are untested.
- Continuous back-to-back streaming heats the S26 GPU until its clock is capped (above); real-time use and the offline mode stayed at full speed in these runs.
- Audio shorter than the first chunk (1.04 s) runs as a single chunk; that path was not compared with the reference.
License
OpenMDW-1.1, as the original model. NOTICE records the origin (repository, revision, weight checksum) and what was changed.
References
- Model: nvidia/Nemotron-3-Diarization (the upstream card has the training data, evaluation and intended use).
- Streaming Sortformer: Speaker Cache-Based Online Speaker Diarization with Arrival-Time Ordering
- Sortformer: A Novel Approach for Permutation-Resolved Speaker Supervision in Speech-to-Text Systems
- transformers implementation: Nemotron3Diarization docs
- LiteRT (CompiledModel API); sample apps for other models: litert-samples.
- Downloads last month
- 156
Model tree for litert-community/Nemotron-3-Diarization-LiteRT
Base model
nvidia/Nemotron-3-Diarization