Nemotron-3-Diarization-litert / reference_runtime.py
spybyscript's picture
Publish fused LiteRT INT8 v2
47382f6 verified
Raw History Blame Contribute Delete
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, ...]
@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())