"""Portable low-latency runtime for the Nemotron-3-Diarization LiteRT bundle.""" from __future__ import annotations import argparse import hashlib import json import math import time from pathlib import Path from typing import Any import librosa import numpy as np SHARDED_ADAPTER_ID = "nemotron3-diarization-low-latency-v1" FUSED_ADAPTER_ID = "nemotron3-diarization-low-latency-v2" SUPPORTED_ADAPTER_IDS = {SHARDED_ADAPTER_ID, FUSED_ADAPTER_ID} SAMPLE_RATE = 16_000 HOP_LENGTH = 160 N_FFT = 512 WIN_LENGTH = 400 MEL_BINS = 128 MEL_FRAMES = 104 SUBSAMPLING_FACTOR = 8 CHUNK_ENCODER_FRAMES = 9 LOOKAHEAD_ENCODER_FRAMES = 4 SPEAKER_CACHE_FRAMES = 264 FIFO_FRAMES = 264 SEQUENCE_FRAMES = 541 HIDDEN_SIZE = 512 NUM_SPEAKERS = 8 NUM_LAYERS = 31 class Graph: """One resident LiteRT graph with a numerically ordered signature.""" def __init__(self, path: Path, threads: int): from ai_edge_litert.interpreter import Interpreter self.path = path self.interpreter = Interpreter(model_path=str(path), num_threads=threads) self.interpreter.allocate_tensors() signatures = self.interpreter.get_signature_list() if set(signatures) != {"serving_default"}: raise ValueError(f"Unexpected signatures in {path.name}: {list(signatures)}") self.runner = self.interpreter.get_signature_runner("serving_default") def __call__(self, *arguments: np.ndarray) -> np.ndarray: """Invoke the graph and return its one finite output.""" result = self.runner(**{f"args_{index}": value for index, value in enumerate(arguments)}) ordered = [result[key] for key in sorted(result, key=lambda key: int(key.rsplit("_", 1)[1]))] if len(ordered) != 1: raise ValueError(f"Expected one output from {self.path.name}, got {len(ordered)}") output = ordered[0] if not np.isfinite(output).all(): raise ValueError(f"Non-finite output from {self.path.name}") return output class SpeakerCache: """NumPy implementation of the upstream AOSC and FIFO streaming policy.""" def __init__(self, silence_embedding: np.ndarray): self.silence_embedding = np.asarray(silence_embedding, dtype=np.float32).reshape(HIDDEN_SIZE) self.embeddings = np.zeros((SPEAKER_CACHE_FRAMES, HIDDEN_SIZE), dtype=np.float32) self.probabilities = np.zeros((SPEAKER_CACHE_FRAMES, NUM_SPEAKERS), dtype=np.float32) self.fifo = np.zeros((FIFO_FRAMES, HIDDEN_SIZE), dtype=np.float32) self.cache_count = 0 self.fifo_count = 0 self.compressed = False def get_embeddings(self) -> np.ndarray: """Return the currently populated AOSC followed by FIFO frames.""" return np.concatenate( [self.embeddings[: self.cache_count], self.fifo[: self.fifo_count]], axis=0 )[None, ...] @staticmethod def _pool_probabilities(logits: np.ndarray) -> np.ndarray: """Pool 10 ms logits to the 80 ms encoder frame rate.""" probabilities = 1.0 / (1.0 + np.exp(-logits.astype(np.float32))) frame_count = probabilities.shape[1] // SUBSAMPLING_FACTOR pooled = probabilities[:, : frame_count * SUBSAMPLING_FACTOR] pooled = pooled.reshape(1, frame_count, SUBSAMPLING_FACTOR, NUM_SPEAKERS).mean(axis=2) return pooled.astype(np.float32) @staticmethod def _top_indices(values: np.ndarray, count: int) -> np.ndarray: """Return unsorted indices for the largest values along the frame axis.""" if count >= values.shape[1]: return np.broadcast_to(np.arange(values.shape[1]), (values.shape[0], values.shape[1], values.shape[2])) partition = np.argpartition(values, values.shape[1] - count, axis=1) return partition[:, -count:, :] @classmethod def _boost_scores(cls, scores: np.ndarray, count: int, boost: float) -> np.ndarray: """Boost each speaker's highest-scoring frames in place.""" indices = cls._top_indices(scores, count) batches = np.arange(scores.shape[0])[:, None, None] speakers = np.arange(scores.shape[2])[None, None, :] scores[batches, indices, speakers] += boost return scores @classmethod def _compress( cls, embeddings: np.ndarray, probabilities: np.ndarray, silence_embedding: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: """Select and arrival-order the 264 AOSC frames used by the source runtime.""" threshold = 0.25 log_probabilities = np.log(np.maximum(probabilities, threshold)) log_complements = np.log(np.maximum(1.0 - probabilities, threshold)) scores = ( log_probabilities - log_complements + log_complements.sum(axis=-1, keepdims=True) - math.log(0.5) ) speech = probabilities > 0.5 scores[~speech] = -np.inf positive = scores > 0.0 enough_positive = positive.sum(axis=1, keepdims=True) >= 16 scores[(~positive) & speech & enough_positive] = -np.inf scores[:, SPEAKER_CACHE_FRAMES:] += 0.05 scores = cls._boost_scores(scores, 24, -2.0 * math.log(0.5)) scores = cls._boost_scores(scores, 48, -math.log(0.5)) frame_count = embeddings.shape[1] scores = np.pad(scores, ((0, 0), (0, 1), (0, 0)), constant_values=np.inf) silence = np.broadcast_to(silence_embedding.reshape(1, 1, -1), (1, 1, HIDDEN_SIZE)) embeddings = np.concatenate([embeddings, silence], axis=1) probabilities = np.pad(probabilities, ((0, 0), (0, 1), (0, 0))) scored_frames = frame_count + 1 flat_scores = scores.transpose(0, 2, 1).reshape(1, -1) selected = np.argpartition(flat_scores, -SPEAKER_CACHE_FRAMES, axis=1)[ :, -SPEAKER_CACHE_FRAMES: ] selected_scores = np.take_along_axis(flat_scores, selected, axis=1) sentinel = scored_frames * NUM_SPEAKERS selected[selected_scores == -np.inf] = sentinel selected.sort(axis=1) frame_indices = np.where(selected == sentinel, frame_count, selected % scored_frames) return embeddings[:, frame_indices[0]], probabilities[:, frame_indices[0]] def update( self, step_embeddings: np.ndarray, step_logits: np.ndarray, chunk_frame_count: int, ) -> None: """Push one processed chunk and update AOSC/FIFO state.""" probabilities = self._pool_probabilities(step_logits) chunk_start = self.cache_count + self.fifo_count chunk = step_embeddings[:, chunk_start : chunk_start + chunk_frame_count] fifo = np.concatenate([self.fifo[None, : self.fifo_count], chunk], axis=1) popped = 0 if fifo.shape[1] > FIFO_FRAMES: popped = min(max(222, fifo.shape[1] - FIFO_FRAMES), fifo.shape[1]) if popped: fifo_probabilities = probabilities[ :, self.cache_count : self.cache_count + fifo.shape[1] ] stored_probabilities = ( self.probabilities[None, : self.cache_count] if self.compressed else probabilities[:, : self.cache_count] ) cache_embeddings = np.concatenate( [self.embeddings[None, : self.cache_count], fifo[:, :popped]], axis=1 ) cache_probabilities = np.concatenate( [stored_probabilities, fifo_probabilities[:, :popped]], axis=1 ) fifo = fifo[:, popped:] if cache_embeddings.shape[1] > SPEAKER_CACHE_FRAMES: cache_embeddings, cache_probabilities = self._compress( cache_embeddings, cache_probabilities, self.silence_embedding ) self.compressed = True self.cache_count = cache_embeddings.shape[1] self.embeddings[: self.cache_count] = cache_embeddings[0] self.probabilities[: self.cache_count] = cache_probabilities[0] self.fifo_count = fifo.shape[1] self.fifo[: self.fifo_count] = fifo[0] def sha256_file(path: Path) -> str: """Return the SHA-256 digest of a file.""" digest = hashlib.sha256() with path.open("rb") as handle: for block in iter(lambda: handle.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() def verify_manifest(bundle: Path) -> dict[str, Any]: """Validate the bundle identity, inventory, sizes, and hashes.""" manifest = json.loads((bundle / "nemotron3-diarization-manifest.json").read_text()) if manifest.get("adapter_id") not in SUPPORTED_ADAPTER_IDS: raise ValueError(f"Wrong adapter: {manifest.get('adapter_id')!r}") for entry in manifest["files"]: path = bundle / entry["path"] if path.parent != bundle or not path.is_file(): raise ValueError(f"Invalid or missing bundle file: {entry['path']}") if path.stat().st_size != entry["bytes"] or sha256_file(path) != entry["sha256"]: raise ValueError(f"Bundle checksum mismatch: {entry['path']}") return manifest def log_mel_features(audio: np.ndarray, *, center: bool) -> np.ndarray: """Compute the source checkpoint's unnormalized 128-bin log-mel features.""" audio = np.asarray(audio, dtype=np.float32) emphasized = np.concatenate([audio[:1], audio[1:] - 0.97 * audio[:-1]]) if center: emphasized = np.pad(emphasized, (N_FFT // 2, N_FFT // 2)) if emphasized.size < N_FFT: return np.zeros((0, MEL_BINS), dtype=np.float32) frames = np.lib.stride_tricks.sliding_window_view(emphasized, N_FFT)[::HOP_LENGTH] window = np.pad(np.hanning(WIN_LENGTH).astype(np.float32), ((N_FFT - WIN_LENGTH) // 2,) * 2) spectrum = np.fft.rfft(frames * window, axis=1) power = np.square(np.abs(spectrum).astype(np.float32)) filters = librosa.filters.mel( sr=SAMPLE_RATE, n_fft=N_FFT, n_mels=MEL_BINS, fmin=0.0, fmax=SAMPLE_RATE / 2, norm="slaney", ).astype(np.float32) return np.log(power @ filters.T + 2**-24).astype(np.float32) def rotary_embeddings() -> tuple[np.ndarray, np.ndarray]: """Build the checkpoint's fixed RoPE cosine and sine inputs.""" inverse_frequency = 1.0 / (10_000.0 ** (np.arange(0, 64, 2, dtype=np.float32) / 64.0)) frequencies = np.outer(np.arange(SEQUENCE_FRAMES, dtype=np.float32), inverse_frequency) embedding = np.concatenate([frequencies, frequencies], axis=-1)[None, ...] return np.cos(embedding).astype(np.float32), np.sin(embedding).astype(np.float32) def audio_chunks(audio: np.ndarray) -> list[tuple[np.ndarray, bool, bool]]: """Split complete audio into the exact low-latency overlapping input windows.""" first_samples = 16_680 regular_samples = 17_040 first_end = min(first_samples, audio.shape[0]) if first_end == audio.shape[0]: return [(audio, True, True)] chunks = [(audio[:first_end], True, False)] mel_frame = CHUNK_ENCODER_FRAMES * SUBSAMPLING_FACTOR start = mel_frame * HOP_LENGTH - N_FFT // 2 while start + regular_samples <= audio.shape[0]: chunks.append((audio[start : start + regular_samples], False, False)) mel_frame += CHUNK_ENCODER_FRAMES * SUBSAMPLING_FACTOR start = mel_frame * HOP_LENGTH - N_FFT // 2 if start < audio.shape[0]: chunks.append((audio[start:], False, True)) return chunks def stabilize_activity(activity: np.ndarray, frames: int = 10) -> np.ndarray: """Close sub-100 ms gaps and remove sub-100 ms speaker events.""" stable = activity.copy() for speaker in range(stable.shape[1]): values = stable[:, speaker] changes = np.diff(np.pad(values.astype(np.int8), (1, 1))) speech_starts = np.flatnonzero(changes == 1) speech_ends = np.flatnonzero(changes == -1) for gap_start, gap_end in zip(speech_ends[:-1], speech_starts[1:], strict=True): if gap_end - gap_start <= frames: values[gap_start:gap_end] = True changes = np.diff(np.pad(values.astype(np.int8), (1, 1))) starts = np.flatnonzero(changes == 1) ends = np.flatnonzero(changes == -1) for start, end in zip(starts, ends, strict=True): if end - start < frames: values[start:end] = False return stable def segments_from_activity(activity: np.ndarray) -> list[dict[str, float | int]]: """Convert 10 ms speaker activity to arrival-ordered segments.""" segments: list[dict[str, float | int]] = [] for speaker in range(activity.shape[1]): changes = np.diff(np.pad(activity[:, speaker].astype(np.int8), (1, 1))) starts = np.flatnonzero(changes == 1) ends = np.flatnonzero(changes == -1) segments.extend( { "start": round(float(start) * 0.01, 2), "end": round(float(end) * 0.01, 2), "speaker": speaker, } for start, end in zip(starts, ends, strict=True) ) return sorted(segments, key=lambda segment: (segment["start"], segment["speaker"])) class Runtime: """Complete fixed-shape low-latency LiteRT diarization runtime.""" def __init__(self, bundle: Path, threads: int = 4, validate: bool = True): self.bundle = Path(bundle) if validate: verify_manifest(self.bundle) paths = {path.stem: path for path in self.bundle.glob("*.tflite")} sharded = { "feature_stacker", "classification_head", *(f"encoder_layer_{index:02d}" for index in range(NUM_LAYERS)), } fused = {"feature_stacker", "encoder_head"} if set(paths) == fused: self.fused = True elif set(paths) == sharded: self.fused = False else: raise ValueError( "Incomplete or unexpected graph inventory: " f"{sorted(set(paths) ^ (fused if 'encoder_head' in paths else sharded))}" ) self.graphs = {name: Graph(path, threads) for name, path in paths.items()} self.silence_embedding = np.load(self.bundle / "silence_embedding.npy") self.rope_cos, self.rope_sin = rotary_embeddings() def diarize(self, audio: np.ndarray, stabilize_ms: int = 100) -> dict[str, Any]: """Diarize finite mono Float32 PCM sampled at exactly 16 kHz.""" audio = np.asarray(audio, dtype=np.float32) if audio.ndim != 1 or not np.isfinite(audio).all(): raise ValueError("Expected finite mono Float32 PCM") cache = SpeakerCache(self.silence_embedding) output_chunks: list[np.ndarray] = [] graph_seconds = 0.0 chunks = audio_chunks(audio) for chunk, first, last in chunks: features = log_mel_features(chunk, center=first) valid_mel_frames = chunk.shape[0] // HOP_LENGTH if first else max( 0, (chunk.shape[0] - N_FFT) // HOP_LENGTH + 1 ) features = features[:valid_mel_frames] padded = np.zeros((1, MEL_FRAMES, MEL_BINS), dtype=np.float32) padded[:, : features.shape[0]] = features started = time.perf_counter() current = self.graphs["feature_stacker"](padded) current_count = math.ceil(valid_mel_frames / SUBSAMPLING_FACTOR) current = current[:, :current_count] cached = cache.get_embeddings() step_embeddings = np.concatenate([cached, current], axis=1) pad_frames = SEQUENCE_FRAMES - step_embeddings.shape[1] if pad_frames < 0: raise ValueError(f"Streaming state exceeded {SEQUENCE_FRAMES} frames") hidden = np.zeros((1, SEQUENCE_FRAMES, HIDDEN_SIZE), dtype=np.float32) hidden[:, pad_frames:] = step_embeddings valid = np.zeros((1, SEQUENCE_FRAMES), dtype=np.bool_) valid[:, pad_frames:] = True attention = np.where( valid[:, None, None, :], 0.0, np.finfo(np.float32).min ).astype(np.float32) if self.fused: full_logits = self.graphs["encoder_head"]( hidden, attention, self.rope_cos, self.rope_sin, valid ) else: for index in range(NUM_LAYERS): hidden = self.graphs[f"encoder_layer_{index:02d}"]( hidden, attention, self.rope_cos, self.rope_sin ) full_logits = self.graphs["classification_head"](hidden, valid) graph_seconds += time.perf_counter() - started lookahead = 0 if last else LOOKAHEAD_ENCODER_FRAMES chunk_encoder_frames = current_count - lookahead start_encoder = pad_frames + cached.shape[1] start_logit = start_encoder * SUBSAMPLING_FACTOR output_frames = min(chunk_encoder_frames * SUBSAMPLING_FACTOR, valid_mel_frames) output_chunks.append(full_logits[:, start_logit : start_logit + output_frames]) valid_start = pad_frames * SUBSAMPLING_FACTOR valid_end = (pad_frames + step_embeddings.shape[1]) * SUBSAMPLING_FACTOR cache.update( step_embeddings, full_logits[:, valid_start:valid_end], chunk_encoder_frames, ) logits = np.concatenate(output_chunks, axis=1)[0] if output_chunks else np.zeros((0, 8)) activity = logits > 0.0 stable = stabilize_activity(activity, max(1, round(stabilize_ms / 10))) return { "segments": segments_from_activity(stable), "raw_segments": segments_from_activity(activity), "audio_seconds": audio.shape[0] / SAMPLE_RATE, "graph_seconds": graph_seconds, "graph_rtf": graph_seconds / (audio.shape[0] / SAMPLE_RATE) if audio.size else 0.0, "chunks": len(chunks), "stabilize_ms": stabilize_ms, } def main() -> int: """Run the portable reference from the command line.""" import soundfile as sf from scipy.signal import resample_poly parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("bundle", type=Path) parser.add_argument("audio", type=Path) parser.add_argument("--threads", type=int, default=4) parser.add_argument("--stabilize-ms", type=int, default=100) args = parser.parse_args() waveform, sample_rate = sf.read(args.audio, dtype="float32", always_2d=True) mono = waveform.mean(axis=1) divisor = math.gcd(sample_rate, SAMPLE_RATE) audio = resample_poly(mono, SAMPLE_RATE // divisor, sample_rate // divisor).astype(np.float32) result = Runtime(args.bundle, args.threads).diarize(audio, args.stabilize_ms) print(json.dumps(result, indent=2, sort_keys=True)) return 0 if __name__ == "__main__": raise SystemExit(main())