litert-community/Nemotron-3-Diarization-LiteRT

Model

Nemotron-3-Diarization — LiteRT (CompiledModel GPU)

4

3 commits

1 linked in READMEs

updated Sep 24, 2026

See the code

README

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

filerolesize
nemotron3_diar_frontend.tflitegraph A: 8-frame mel stacking + projection (float32)2.1 MB
nemotron3_diar_encoder_low_latency_fp16.tflitegraph B, streaming: 31-layer encoder + output head, fixed T = 541 rows (float16 weights, float compute)198.7 MB
nemotron3_diar_encoder_offline_fp16.tflitegraph B, whole file: same network, fixed T = 684 rows (float16 weights)198.7 MB
assets/frontend_mel128_257.binslaney mel filter bank [128, 257], float32 little-endian131.6 KB
assets/hann400.binHann window, 400 samples, symmetric, float321.6 KB
assets/silence_embeds.binlearned silence embedding [512] that fills empty speaker-cache slots, float322.0 KB
conversion/reference host nemotron3_diar_litert.py, conversion and verification scripts (conversion/README.md)
android/Android sample app source (Kotlin)
LICENSE, NOTICE, SHA256SUMSlicense 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 + 4340 + 40
audio before the first decision1.04 s (16,680 samples)the whole file
step0.72 s27.2 s
speaker cache / FIFO (encoder frames)264 / 264264 / 40
cache update period222300
graph B rows T541 = 264 + 264 + 13684 = 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.

graphinputshapeoutputshape
A nemotron3_diar_frontendmel: log-mel of one chunk, zero rows after the last frame[1, 104, 128]chunk_embeds[1, 13, 512]
B …_low_latency_fp16packed_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_fp16the same four inputs with T = 684logits[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 precisionlatency per step, median / p95RTFagreement @ 0.5segments vs referencecache selectionscompile A / B
FP32170.7 / 175.5 ms0.238100 % (0 flips)37 / 37 identical4 / 4 identical72 / 1,477 ms
FP16 (GPU default)115.1 / 120.0 ms0.16099.973 % (21 flips)40: 2 added, 1 split, 8 boundaries moved by 10 ms4 / 4 differ (3–17 of 264 frames)73 / 1,434 ms
FP16, FP32 accumulation140.3 / 144.7 ms0.19699.997 % (2 flips)38: 1 split, 1 boundary moved by 10 ms2 / 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):
modeclipmax |Δlogit|max |Δp|flipssegmentscache state and selections
low_latency97.6 s6.9e-56.1e-60 / 78,07237 / 37 identicalequal at all 136 steps
low_latency21.5 s6.5e-52.8e-60 / 17,19210 / 10 identicalequal at all 30 steps
offline97.6 s4.6e-54.6e-60 / 78,08829 / 29 identicalequal in all 4 chunks
offline21.5 s4.6e-52.3e-60 / 17,2089 / 9 identicalequal
  • 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

android
litert
on-device
speaker-diarization
streaming
tflite
voice-activity-detection

litert-community/Nemotron-3-Diarization-LiteRT

Model

Nemotron-3-Diarization — LiteRT (CompiledModel GPU)

4

3 commits

1 linked in READMEs

updated Sep 24, 2026

See the code

README

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

filerolesize
nemotron3_diar_frontend.tflitegraph A: 8-frame mel stacking + projection (float32)2.1 MB
nemotron3_diar_encoder_low_latency_fp16.tflitegraph B, streaming: 31-layer encoder + output head, fixed T = 541 rows (float16 weights, float compute)198.7 MB
nemotron3_diar_encoder_offline_fp16.tflitegraph B, whole file: same network, fixed T = 684 rows (float16 weights)198.7 MB
assets/frontend_mel128_257.binslaney mel filter bank [128, 257], float32 little-endian131.6 KB
assets/hann400.binHann window, 400 samples, symmetric, float321.6 KB
assets/silence_embeds.binlearned silence embedding [512] that fills empty speaker-cache slots, float322.0 KB
conversion/reference host nemotron3_diar_litert.py, conversion and verification scripts (conversion/README.md)
android/Android sample app source (Kotlin)
LICENSE, NOTICE, SHA256SUMSlicense 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 + 4340 + 40
audio before the first decision1.04 s (16,680 samples)the whole file
step0.72 s27.2 s
speaker cache / FIFO (encoder frames)264 / 264264 / 40
cache update period222300
graph B rows T541 = 264 + 264 + 13684 = 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.

graphinputshapeoutputshape
A nemotron3_diar_frontendmel: log-mel of one chunk, zero rows after the last frame[1, 104, 128]chunk_embeds[1, 13, 512]
B …_low_latency_fp16packed_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_fp16the same four inputs with T = 684logits[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 precisionlatency per step, median / p95RTFagreement @ 0.5segments vs referencecache selectionscompile A / B
FP32170.7 / 175.5 ms0.238100 % (0 flips)37 / 37 identical4 / 4 identical72 / 1,477 ms
FP16 (GPU default)115.1 / 120.0 ms0.16099.973 % (21 flips)40: 2 added, 1 split, 8 boundaries moved by 10 ms4 / 4 differ (3–17 of 264 frames)73 / 1,434 ms
FP16, FP32 accumulation140.3 / 144.7 ms0.19699.997 % (2 flips)38: 1 split, 1 boundary moved by 10 ms2 / 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):
modeclipmax |Δlogit|max |Δp|flipssegmentscache state and selections
low_latency97.6 s6.9e-56.1e-60 / 78,07237 / 37 identicalequal at all 136 steps
low_latency21.5 s6.5e-52.8e-60 / 17,19210 / 10 identicalequal at all 30 steps
offline97.6 s4.6e-54.6e-60 / 78,08829 / 29 identicalequal in all 4 chunks
offline21.5 s4.6e-52.3e-60 / 17,2089 / 9 identicalequal
  • 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

android
litert
on-device
speaker-diarization
streaming
tflite
voice-activity-detection