Nemotron-3-Diarization — LiteRT (CompiledModel GPU)
4
3 commits
1 linked in READMEs
updated Sep 24, 2026
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.
| 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.
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.
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.sigmoid(logits); the host does it.conversion/nemotron3_diar_litert.py (numpy) and android/…/Nemotron3Diarizer.kt + SpeakerCache.kt +
MelFrontend.kt implement this loop; both follow transformers' Nemotron3DiarizationProcessor and
Nemotron3DiarizationSpeakerCache.
[cache rows | FIFO rows | 13 new rows] (L ≤ 541) plus attn_bias, run graph B.8·(cache + FIFO) … 8·(cache + FIFO + 9).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).
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 |
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.
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).
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 |
--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).# 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).
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/ 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
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 + bias
The 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.rfft on 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).
low_latency (1.04 s) and offline. The upstream very_low_latency (0.64 s) and
ultra_low_latency (0.32 s) modes need graph B builds with their own T; not included.OpenMDW-1.1, as the original model. NOTICE records the origin (repository, revision, weight checksum) and what was changed.
Nemotron-3-Diarization — LiteRT (CompiledModel GPU)
4
3 commits
1 linked in READMEs
updated Sep 24, 2026
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.
| 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.
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.
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.sigmoid(logits); the host does it.conversion/nemotron3_diar_litert.py (numpy) and android/…/Nemotron3Diarizer.kt + SpeakerCache.kt +
MelFrontend.kt implement this loop; both follow transformers' Nemotron3DiarizationProcessor and
Nemotron3DiarizationSpeakerCache.
[cache rows | FIFO rows | 13 new rows] (L ≤ 541) plus attn_bias, run graph B.8·(cache + FIFO) … 8·(cache + FIFO + 9).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).
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 |
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.
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).
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 |
--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).# 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).
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/ 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
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 + bias
The 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.rfft on 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).
low_latency (1.04 s) and offline. The upstream very_low_latency (0.64 s) and
ultra_low_latency (0.32 s) modes need graph B builds with their own T; not included.OpenMDW-1.1, as the original model. NOTICE records the origin (repository, revision, weight checksum) and what was changed.