Download reference_runtime.py from spybyscript/Nemotron-3-Diarization-litert: direct link, hf CLI and curl.
- Browser
- Download file 19 kB
-
https://huggingface.co/spybyscript/Nemotron-3-Diarization-litert/resolve/main/reference_runtime.py
- Command line
-
hf download hf://spybyscript/Nemotron-3-Diarization-litert/reference_runtime.py
-
curl -L -o reference_runtime.py https://huggingface.co/spybyscript/Nemotron-3-Diarization-litert/resolve/main/reference_runtime.py
19 kB
| """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, ...] | |
| 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) | |
| 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:, :] | |
| 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 | |
| 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()) | |