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.

Who spoke when: waveform and per-speaker timeline of an 87.8 s dialogue, 5 speakers

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.

  1. 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.
  2. 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.
  3. Graph A turns the chunk's 104 mel frames into 13 rows (9 chunk + 4 look-ahead).
  4. Pack [cache rows | FIFO rows | 13 new rows] (L ≤ 541) plus attn_bias, run graph B.
  5. Emit the chunk's logit rows 8·(cache + FIFO) … 8·(cache + FIFO + 9).
  6. 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 + 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).

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) 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.
  • 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

Downloads last month
156
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for litert-community/Nemotron-3-Diarization-LiteRT

Quantized
(27)
this model

Papers for litert-community/Nemotron-3-Diarization-LiteRT