diff --git a/README.md b/README.md index 9e983d8..e00589a 100644 --- a/README.md +++ b/README.md @@ -103,15 +103,22 @@ stream_from_file(streamer, "sample.wav") ### Speech Diarization (Who Spoke When) -Detect speaker turns and timestamps using ONNX: +Detect speaker turns and timestamps using ONNX (defaults to NVIDIA Nemotron-3 Diarization, also supports Pyannote Segmentation 3.0): ```python -from pythaiasr import diarize +from pythaiasr import diarize, segments_to_rttm -# Identify who spoke when +# NVIDIA Nemotron-3 Diarization (default: fast INT8 ONNX, up to 8 speakers) segments = diarize("meeting.wav") for seg in segments: print(f"[{seg['start']:.2f}s - {seg['end']:.2f}s] {seg['speaker']}") + +# Pyannote Segmentation 3.0 (optional) +segments_pyannote = diarize("meeting.wav", model="pyannote_segmentation") + +# Export to standard NIST RTTM format +rttm_str = segments_to_rttm(segments, uri="meeting") +print(rttm_str) ``` ### Speech Diarization + ASR (`asr_diarize`) @@ -121,7 +128,7 @@ Detect speakers and transcribe each speaker turn with ASR: ```python from pythaiasr import asr_diarize -# Attributed transcription per speaker +# Attributed transcription per speaker with Typhoon ASR and Nemotron Diarization (default) turns = asr_diarize("meeting.wav", asr_model="typhoon_asr") for turn in turns: print(f"[{turn['start']:.2f}s - {turn['end']:.2f}s] {turn['speaker']}: {turn['text']}") @@ -134,19 +141,19 @@ See examples/diarize_example.py 1. Speech Diarization (Who Spoke When) ============================================================ Processing: examples/../tests/test-diarize.wav ... -[ 0.17s -> 1.87s] SPEAKER_02 -[ 1.97s -> 4.30s] SPEAKER_01 -[ 4.83s -> 6.49s] SPEAKER_02 -[ 6.78s -> 8.36s] SPEAKER_01 +[ 0.15s -> 1.82s] SPEAKER_00 +[ 1.95s -> 4.34s] SPEAKER_01 +[ 4.81s -> 6.45s] SPEAKER_00 +[ 6.81s -> 8.39s] SPEAKER_01 ============================================================ 2. Diarization + Speech Recognition (ASR Diarize) ============================================================ Transcribing turns with Typhoon ASR: examples/../tests/test-diarize.wav ... -[ 0.17s -> 1.87s] SPEAKER_02: สวัสดีชาวโลกทุกท่าน -[ 1.97s -> 4.30s] SPEAKER_01: แล้วระบบนี้ทํางานอย่างไร -[ 4.83s -> 6.49s] SPEAKER_02: ใช้ปัญญาประดิษฐ์ในการทดสอบ -[ 6.78s -> 8.36s] SPEAKER_01: ใช้งานได้ดีทีเดียว +[ 0.15s -> 1.82s] SPEAKER_00: สวัสดีชาวโลกทุกท่าน +[ 1.95s -> 4.34s] SPEAKER_01: แล้วระบบนี้ทํางานอย่างไร +[ 4.81s -> 6.45s] SPEAKER_00: ใช้ปัญญาประดิษฐ์ในการทดสอบ +[ 6.81s -> 8.39s] SPEAKER_01: ใช้งานได้ดีทีเดียว ``` ### API @@ -217,14 +224,16 @@ You can read about models from the list: - [*biodatlab/whisper-th-medium-combined* - Thai Whisper medium model](https://huggingface.co/biodatlab/whisper-th-medium-combined) - [*biodatlab/whisper-th-large-combined* - Thai Whisper large model](https://huggingface.co/biodatlab/whisper-th-large-combined) - [*biodatlab/whisper-th-medium-timestamp* - Thai Whisper medium model with timestamp support](https://huggingface.co/biodatlab/whisper-th-medium-timestamp) +- [*joosthel/Nemotron-3-Diarization-ONNX* - NVIDIA Nemotron-3 Diarization ONNX model](https://huggingface.co/joosthel/Nemotron-3-Diarization-ONNX) #### diarize ```python diarize( data: Union[str, Path, np.ndarray], - model: str = "pyannote_segmentation", + model: str = "nemotron-3-diarization", device: Optional[str] = None, + precision: str = "int8", sampling_rate: int = 16_000, num_speakers: Optional[int] = None, min_speakers: Optional[int] = None, @@ -238,8 +247,11 @@ diarize( ``` - `data`: Audio file path or 1D numpy array of audio waveform. -- `model`: Diarization model identifier (default: `"pyannote_segmentation"`). +- `model`: Diarization model identifier: + - `"nemotron-3-diarization"` (default) / `"joosthel/Nemotron-3-Diarization-ONNX"` - NVIDIA Nemotron-3 Diarization (streaming cache, up to 8 speakers) + - `"pyannote_segmentation"` - Pyannote Segmentation 3.0 ONNX - `device`: Device to run inference on (`"auto"`, `"cpu"`, `"cuda"`). +- `precision`: Model precision for Nemotron (`"int8"` default, or `"fp32"`). - `sampling_rate`: Audio sampling rate (default: 16000). - `num_speakers`: Exact number of speakers if known. - `onset`: Speech onset probability threshold (default: 0.5). @@ -255,8 +267,9 @@ diarize( asr_diarize( data: Union[str, Path, np.ndarray], asr_model: str = "typhoon_asr", - diarize_model: str = "pyannote_segmentation", + diarize_model: str = "nemotron-3-diarization", device: Optional[str] = None, + precision: str = "int8", sampling_rate: int = 16_000, lm: bool = False, num_speakers: Optional[int] = None, @@ -269,7 +282,8 @@ asr_diarize( - `data`: Audio file path or 1D numpy array of audio waveform. - `asr_model`: The ASR model name (default: `"typhoon_asr"`). -- `diarize_model`: Diarization model name (default: `"pyannote_segmentation"`). +- `diarize_model`: Diarization model name (default: `"nemotron-3-diarization"`, or `"pyannote_segmentation"`). +- `precision`: Model precision for Nemotron (`"int8"` default, or `"fp32"`). - `merge_same_speaker`: Whether to merge adjacent speech turns from the same speaker (default: `True`). - `max_merge_gap`: Maximum gap in seconds between same-speaker segments to merge (default: 0.5). - `backend`: Diarization engine (`"onnx"` or `"sherpa-onnx"`, default: `"onnx"`). diff --git a/examples/diarize_example.py b/examples/diarize_example.py index 9f4ecd3..35f800d 100644 --- a/examples/diarize_example.py +++ b/examples/diarize_example.py @@ -13,7 +13,7 @@ # Ensure pythaiasr package in current repo can be imported directly sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) -from pythaiasr import diarize, asr_diarize +from pythaiasr import diarize, asr_diarize, segments_to_rttm def main(): @@ -27,7 +27,7 @@ def main(): sys.exit(1) print("=" * 60) - print("1. Speech Diarization (Who Spoke When)") + print("1. Speech Diarization (defaults to Nemotron-3 INT8 ONNX)") print("=" * 60) print(f"Processing: {test_audio} ...") segments = diarize(test_audio, device="auto") @@ -38,11 +38,18 @@ def main(): for seg in segments: print(f"[{seg['start']:6.2f}s -> {seg['end']:6.2f}s] {seg['speaker']}") - print("\n" + "=" * 60) + print("\nNIST RTTM Format:") + print(segments_to_rttm(segments, uri=os.path.basename(test_audio))) + + print("=" * 60) print("2. Diarization + Speech Recognition (ASR Diarize)") print("=" * 60) print(f"Transcribing turns with Typhoon ASR: {test_audio} ...") - turns = asr_diarize(test_audio, asr_model="typhoon_asr", device="auto") + turns = asr_diarize( + test_audio, + asr_model="typhoon_asr", + device="auto", + ) if not turns: print("No speech turns found.") @@ -54,3 +61,4 @@ def main(): if __name__ == "__main__": main() + diff --git a/pythaiasr/__init__.py b/pythaiasr/__init__.py index fcab21a..1cde44b 100644 --- a/pythaiasr/__init__.py +++ b/pythaiasr/__init__.py @@ -26,6 +26,7 @@ get_pythaiasr_path, get_typhoon_model_files, get_diarization_model_files, + get_nemotron_diarization_model_files, download_file, ) from pythaiasr.typhoon import ( @@ -40,9 +41,13 @@ ) from pythaiasr.diarization import ( Diarization, + NemotronDiarization, + Nemotron3Diarization, diarize, asr_diarize, merge_same_speaker_segments, + segments_to_rttm, + extract_speaker_dict, ) # Friendly alias @@ -495,9 +500,13 @@ def stream_asr( "asr", "stream_asr", "Diarization", + "NemotronDiarization", + "Nemotron3Diarization", "diarize", "asr_diarize", "merge_same_speaker_segments", + "segments_to_rttm", + "extract_speaker_dict", "FastConformerRNNT", "TyphoonASR", "RealtimeStreamASR", @@ -510,5 +519,6 @@ def stream_asr( "get_pythaiasr_path", "get_typhoon_model_files", "get_diarization_model_files", + "get_nemotron_diarization_model_files", "download_file", ] diff --git a/pythaiasr/diarization.py b/pythaiasr/diarization.py index 49129f3..5eb5a96 100644 --- a/pythaiasr/diarization.py +++ b/pythaiasr/diarization.py @@ -9,6 +9,7 @@ from __future__ import annotations +import math import os from itertools import permutations from pathlib import Path @@ -24,7 +25,10 @@ except ImportError: ort = None -from pythaiasr.download import get_diarization_model_files +from pythaiasr.download import ( + get_diarization_model_files, + get_nemotron_diarization_model_files, +) from pythaiasr.typhoon import load_audio, resample_audio @@ -141,21 +145,534 @@ def _binarize_timeline( return segments +# Constants for Nemotron-3 Diarization (joosthel/Nemotron-3-Diarization-ONNX) +NEMOTRON_SAMPLE_RATE = 16000 +NEMOTRON_HOP_LENGTH = 160 +NEMOTRON_N_FFT = 512 +NEMOTRON_PREEMPHASIS = 0.97 +NEMOTRON_FRAME_DURATION = NEMOTRON_HOP_LENGTH / NEMOTRON_SAMPLE_RATE # 0.01s (10ms) +LOG_MEL_WINDOW_FRAMES = 6000 + +NEMOTRON_DIARIZATION_MODELS = [ + "nemotron", + "nemotron_diarization", + "nemotron-diarization", + "nemotron-3-diarization", + "nemotron-3-diarization-onnx", + "joosthel/Nemotron-3-Diarization-ONNX", + "joosthel/nemotron-3-diarization-onnx", +] + + +def _sigmoid(x: np.ndarray) -> np.ndarray: + """Numerically stable logistic sigmoid.""" + z = np.exp(-np.abs(x)) + return np.where(x >= 0, 1.0 / (1.0 + z), z / (1.0 + z)) + + +def _stable_topk_indices(scores: np.ndarray, k: int, axis: int) -> np.ndarray: + """Indices of the k largest entries along axis, ties broken toward the lower index.""" + order = np.argsort(-scores, axis=axis, kind="stable") + return np.take(order, np.arange(k), axis=axis) + + +class NumpySpeakerCache: + """ + Line-by-line numpy port of Nemotron3DiarizationSpeakerCache. + Eagerly allocates fixed-size embeds/probs/fifo buffers. + """ + + def __init__( + self, + constants: dict, + fifo_length: int, + speaker_cache_update_period: int, + topk_fn=_stable_topk_indices, + ): + self._topk_fn = topk_fn + self.hidden_size = int(constants["hidden_size"]) + self.num_speakers = int(constants["num_speakers"]) + self.subsampling_factor = int(constants["subsampling_factor"]) + self.speaker_cache_length = int(constants["speaker_cache_length"]) + self.num_silence_frames = int(constants["speaker_cache_silence_frames_per_speaker"]) + self.prediction_score_threshold = float(constants["prediction_score_threshold"]) + self.latest_frames_score_boost = float(constants["latest_frames_score_boost"]) + self.silence_embeds = constants["silence_embeds"].astype(np.float32) + + self.fifo_length = fifo_length + self.speaker_cache_update_period = speaker_cache_update_period + + budget = self.speaker_cache_length // self.num_speakers - self.num_silence_frames + self.min_positive_scores = math.floor(budget * float(constants["min_positive_scores_rate"])) + self.num_strong_boosted_frames = math.floor(budget * float(constants["strong_boost_rate"])) + self.num_weak_boosted_frames = math.floor(budget * float(constants["weak_boost_rate"])) + + self.embeds = np.zeros((1, self.speaker_cache_length, self.hidden_size), dtype=np.float32) + self.probs = np.zeros((1, self.speaker_cache_length, self.num_speakers), dtype=np.float32) + self.fifo = np.zeros((1, self.fifo_length, self.hidden_size), dtype=np.float32) + self.num_cache_frames = 0 + self.num_fifo_frames = 0 + self.is_compressed = False + + def get_embeds(self) -> np.ndarray: + return np.concatenate( + [self.embeds[:, : self.num_cache_frames], self.fifo[:, : self.num_fifo_frames]], + axis=1, + ) + + def _pool_probs(self, logits: np.ndarray) -> np.ndarray: + """Speaker probabilities at encoder frame rate (avg_pool1d(kernel=stride=8)).""" + factor = self.subsampling_factor + probs = _sigmoid(logits) + batch, num_frames, num_speakers = probs.shape + pooled = probs.reshape(batch, num_frames // factor, factor, num_speakers).mean(axis=2) + return pooled.astype(self.probs.dtype) + + def _num_popped_frames(self, num_fifo_frames: int) -> int: + if num_fifo_frames <= self.fifo_length: + return 0 + num_popped = max(self.speaker_cache_update_period, num_fifo_frames - self.fifo_length) + return min(num_popped, num_fifo_frames) + + def update(self, chunk_input_embeds: np.ndarray, chunk_logits: np.ndarray, num_chunk_frames: int): + num_cache_frames, num_fifo_frames = self.num_cache_frames, self.num_fifo_frames + probs = self._pool_probs(chunk_logits) + + chunk_start = num_cache_frames + num_fifo_frames + chunk_embeds = chunk_input_embeds[:, chunk_start : chunk_start + num_chunk_frames] + fifo_embeds = np.concatenate([self.fifo[:, :num_fifo_frames], chunk_embeds], axis=1) + + num_popped = self._num_popped_frames(fifo_embeds.shape[1]) + if num_popped: + fifo_probs = probs[:, num_cache_frames : num_cache_frames + fifo_embeds.shape[1]] + stored_probs = self.probs[:, :num_cache_frames] if self.is_compressed else probs[:, :num_cache_frames] + cache_embeds = np.concatenate([self.embeds[:, :num_cache_frames], fifo_embeds[:, :num_popped]], axis=1) + cache_probs = np.concatenate([stored_probs, fifo_probs[:, :num_popped]], axis=1) + fifo_embeds = fifo_embeds[:, num_popped:] + + if cache_embeds.shape[1] > self.speaker_cache_length: + cache_embeds, cache_probs = self._compress(cache_embeds, cache_probs) + self.is_compressed = True + self.num_cache_frames = cache_embeds.shape[1] + + self.embeds[:, : self.num_cache_frames] = cache_embeds + self.probs[:, : self.num_cache_frames] = cache_probs + + self.num_fifo_frames = fifo_embeds.shape[1] + self.fifo[:, : self.num_fifo_frames] = fifo_embeds + + def _get_frame_scores(self, probs: np.ndarray) -> np.ndarray: + threshold = self.prediction_score_threshold + log_probs = np.log(np.clip(probs, threshold, None)) + log_complements = np.log(np.clip(1.0 - probs, threshold, None)) + + scores = log_probs - log_complements + log_complements.sum(axis=-1, keepdims=True) - math.log(0.5) + + is_speech = probs > 0.5 + scores = np.where(is_speech, scores, -np.inf) + is_positive = scores > 0 + has_enough_positive = is_positive.sum(axis=1, keepdims=True) >= self.min_positive_scores + scores = np.where(~is_positive & is_speech & has_enough_positive, -np.inf, scores) + return scores + + def _boost_scores(self, scores: np.ndarray, num_boosted: int, boost: float) -> np.ndarray: + if num_boosted <= 0: + return scores + topk_idx = self._topk_fn(scores, num_boosted, axis=1) + scores = scores.copy() + boosted = np.take_along_axis(scores, topk_idx, axis=1) + boost + np.put_along_axis(scores, topk_idx, boosted, axis=1) + return scores + + def _compress(self, embeds: np.ndarray, probs: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: + batch_size, num_frames, num_speakers = probs.shape + + scores = self._get_frame_scores(probs) + scores = scores.copy() + scores[:, self.speaker_cache_length :] += self.latest_frames_score_boost + + scores = self._boost_scores(scores, self.num_strong_boosted_frames, boost=-2.0 * math.log(0.5)) + scores = self._boost_scores(scores, self.num_weak_boosted_frames, boost=-math.log(0.5)) + scores = np.pad(scores, ((0, 0), (0, self.num_silence_frames), (0, 0)), constant_values=np.inf) + embeds = np.concatenate( + [embeds, np.broadcast_to(self.silence_embeds.reshape(1, 1, -1), (batch_size, 1, embeds.shape[-1]))], + axis=1, + ) + probs = np.pad(probs, ((0, 0), (0, 1), (0, 0))) + + num_scored_frames = num_frames + self.num_silence_frames + sentinel = num_scored_frames * num_speakers + flat_scores = scores.transpose(0, 2, 1).reshape(batch_size, -1) + + topk_idx = self._topk_fn(flat_scores, self.speaker_cache_length, axis=1) + topk_scores = np.take_along_axis(flat_scores, topk_idx, axis=1) + topk_idx = np.where(topk_scores == -np.inf, sentinel, topk_idx) + topk_idx = np.sort(topk_idx, axis=1) + frame_indices = np.where(topk_idx == sentinel, num_frames, np.minimum(topk_idx % num_scored_frames, num_frames)) + + batch_indices = np.arange(batch_size)[:, None] + return embeds[batch_indices, frame_indices], probs[batch_indices, frame_indices] + + +class NemotronDiarization: + """ + ONNX-based Speaker Diarization using NVIDIA Nemotron-3 Diarization (joosthel/Nemotron-3-Diarization-ONNX). + + Pure-numpy + onnxruntime offline speaker diarization engine, supporting up to 8 + concurrently tracked speakers with 10ms frame resolution. + """ + + def __init__( + self, + model_dir: Optional[Union[str, Path]] = None, + precision: str = "int8", + device: Optional[str] = None, + threads: Optional[int] = None, + providers: Optional[List[str]] = None, + **kwargs, + ) -> None: + if ort is None: + raise ImportError( + "onnxruntime is required for Nemotron Diarization. " + "Install it with: pip install onnxruntime" + ) + + self.device = device or "auto" + self.precision = precision + + prep_path, model_path, const_path = get_nemotron_diarization_model_files( + model_dir=model_dir, precision=precision + ) + + sess_options = ort.SessionOptions() + sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL + sess_options.log_severity_level = 3 + if threads: + sess_options.intra_op_num_threads = threads + + if providers is None: + providers = self._resolve_providers(self.device, ort) + + self.preprocessor_core = ort.InferenceSession(prep_path, sess_options=sess_options, providers=providers) + self.model = ort.InferenceSession(model_path, sess_options=sess_options, providers=providers) + self.constants = dict(np.load(const_path)) + + self.subsampling_factor = int(self.constants["subsampling_factor"]) + self.chunk_length = int(self.constants["chunk_length"]) + self.chunk_right_context = int(self.constants["chunk_right_context"]) + self.fifo_length = int(self.constants["fifo_length"]) + self.speaker_cache_update_period = int(self.constants["speaker_cache_update_period"]) + + @staticmethod + def _resolve_providers(device: str, ort_mod) -> List[str]: + avail = ort_mod.get_available_providers() + if device in ("cuda", "gpu"): + if "CUDAExecutionProvider" in avail: + return ["CUDAExecutionProvider", "CPUExecutionProvider"] + raise RuntimeError("CUDA requested, but 'CUDAExecutionProvider' is not available in onnxruntime.") + elif device == "cpu": + return ["CPUExecutionProvider"] + else: # auto + if "CUDAExecutionProvider" in avail: + return ["CUDAExecutionProvider", "CPUExecutionProvider"] + return ["CPUExecutionProvider"] + + def log_mel(self, waveform: np.ndarray) -> np.ndarray: + """Bounded-memory windowed log-mel feature extraction.""" + waveform = np.asarray(waveform, dtype=np.float32) + num_samples = waveform.shape[0] + num_frames_total = 1 + num_samples // NEMOTRON_HOP_LENGTH + mel = np.empty((1, num_frames_total, 128), dtype=np.float32) + + start_frame = 0 + while start_frame < num_frames_total: + end_frame = min(start_frame + LOG_MEL_WINDOW_FRAMES, num_frames_total) + raw_start = start_frame * NEMOTRON_HOP_LENGTH - NEMOTRON_N_FFT // 2 + raw_end = (end_frame - 1) * NEMOTRON_HOP_LENGTH - NEMOTRON_N_FFT // 2 + NEMOTRON_N_FFT + lookback = raw_start - 1 + real_start = max(0, lookback) + real_end = max(real_start, min(num_samples, raw_end)) + real_segment = waveform[real_start:real_end] + + if lookback >= 0: + preemph_real = real_segment[1:] - NEMOTRON_PREEMPHASIS * real_segment[:-1] + content_start = raw_start + else: + preemph_real = real_segment.copy() + preemph_real[1:] -= NEMOTRON_PREEMPHASIS * real_segment[:-1] + content_start = real_start + + left_pad = content_start - raw_start + right_pad = raw_end - real_end + preemphasized = ( + np.pad(preemph_real, (left_pad, right_pad)) + if (left_pad or right_pad) + else preemph_real + ) + + (mel_chunk,) = self.preprocessor_core.run(None, {"preemphasized": preemphasized[None, :]}) + mel[:, start_frame:end_frame] = mel_chunk[:, : end_frame - start_frame] + start_frame = end_frame + + if num_frames_total: + mel[:, -1, :] = 0.0 + return mel + + def run_chunk( + self, chunk_mel: np.ndarray, chunk_mel_length: int, context_embeds: np.ndarray + ) -> Tuple[np.ndarray, np.ndarray]: + logits, embeds = self.model.run( + None, + { + "chunk_mel": chunk_mel, + "chunk_mel_length": np.array(chunk_mel_length, dtype=np.int64), + "context_embeds": context_embeds, + "context_length": np.array(context_embeds.shape[1], dtype=np.int64), + }, + ) + return logits, embeds + + def predict_proba( + self, + data: Union[str, Path, np.ndarray], + sampling_rate: int = NEMOTRON_SAMPLE_RATE, + ) -> np.ndarray: + """ + Get frame-level speaker probabilities of shape (1, num_frames, 8). + """ + if isinstance(data, (str, Path)): + audio = load_audio(data, target_sr=NEMOTRON_SAMPLE_RATE) + elif isinstance(data, np.ndarray): + audio = np.asarray(data, dtype=np.float32) + if audio.ndim > 1: + audio = np.mean(audio, axis=1) + if sampling_rate != NEMOTRON_SAMPLE_RATE: + audio = resample_audio(audio, orig_sr=sampling_rate, target_sr=NEMOTRON_SAMPLE_RATE) + else: + raise TypeError(f"Unsupported data type for diarization: {type(data)}") + + if len(audio) == 0: + return np.zeros((1, 0, 8), dtype=np.float32) + + mel = self.log_mel(audio) + num_mel_frames = mel.shape[1] + num_embeds = -(-num_mel_frames // self.subsampling_factor) + valid_mel_frames = max(num_mel_frames - 1, 0) + + cache = NumpySpeakerCache( + self.constants, + fifo_length=self.fifo_length, + speaker_cache_update_period=self.speaker_cache_update_period, + topk_fn=_stable_topk_indices, + ) + + chunk_logits = [] + for start_idx in range(0, num_embeds, self.chunk_length): + end_idx = min(start_idx + self.chunk_length, num_embeds) + num_chunk_frames = end_idx - start_idx + + mel_start = start_idx * self.subsampling_factor + mel_end = min((end_idx + self.chunk_right_context) * self.subsampling_factor, num_mel_frames) + chunk_mel = mel[:, mel_start:mel_end] + chunk_mel_length = max(min(mel_end, valid_mel_frames) - mel_start, 0) + + context_embeds = cache.get_embeds() + context_length = context_embeds.shape[1] + + logits, chunk_input_embeds = self.run_chunk(chunk_mel, chunk_mel_length, context_embeds) + cache.update(chunk_input_embeds, logits, num_chunk_frames) + + start_logit_idx = context_length * self.subsampling_factor + end_logit_idx = (context_length + num_chunk_frames) * self.subsampling_factor + chunk_logits.append(logits[:, start_logit_idx:end_logit_idx]) + + if not chunk_logits: + return np.zeros((1, 0, 8), dtype=np.float32) + + logits = np.concatenate(chunk_logits, axis=1)[:, :num_mel_frames] + return _sigmoid(logits) + + def diarize( + self, + data: Union[str, Path, np.ndarray], + sampling_rate: int = NEMOTRON_SAMPLE_RATE, + threshold: float = 0.5, + onset: Optional[float] = None, + offset: Optional[float] = None, + num_speakers: Optional[int] = None, + min_speakers: Optional[int] = None, + max_speakers: Optional[int] = None, + min_duration_on: float = 0.3, + min_duration_off: float = 0.5, + **kwargs, + ) -> List[Dict[str, Union[float, str]]]: + """ + Diarize audio data and return list of speaker turns. + """ + probs = self.predict_proba(data, sampling_rate=sampling_rate) + if probs.shape[1] == 0: + return [] + + segments = [] + num_frames = probs.shape[1] + num_speakers_total = probs.shape[2] + on_th = onset if onset is not None else threshold + off_th = offset if offset is not None else on_th + + for spk_idx in range(num_speakers_total): + spk_probs = probs[0, :, spk_idx] + is_active = False + start_frame = 0 + spk_turns = [] + + for f in range(num_frames): + p = spk_probs[f] + if not is_active: + if p >= on_th: + is_active = True + start_frame = f + else: + if p < off_th: + is_active = False + end_frame = f + start_t = round(start_frame * NEMOTRON_FRAME_DURATION, 3) + end_t = round(end_frame * NEMOTRON_FRAME_DURATION, 3) + if end_t > start_t: + spk_turns.append((start_t, end_t)) + + if is_active: + start_t = round(start_frame * NEMOTRON_FRAME_DURATION, 3) + end_t = round(num_frames * NEMOTRON_FRAME_DURATION, 3) + if end_t > start_t: + spk_turns.append((start_t, end_t)) + + # Filter out turns shorter than min_duration_on + if min_duration_on > 0: + spk_turns = [seg for seg in spk_turns if (seg[1] - seg[0]) >= min_duration_on] + + if not spk_turns: + continue + + # Merge gaps smaller than min_duration_off + if min_duration_off > 0 and len(spk_turns) > 1: + merged = [spk_turns[0]] + for curr_s, curr_e in spk_turns[1:]: + prev_s, prev_e = merged[-1] + if (curr_s - prev_e) < min_duration_off: + merged[-1] = (prev_s, max(prev_e, curr_e)) + else: + merged.append((curr_s, curr_e)) + spk_turns = merged + + for s, e in spk_turns: + segments.append({ + "start": round(s, 3), + "end": round(e, 3), + "speaker": f"SPEAKER_{spk_idx:02d}", + }) + + segments.sort(key=lambda x: (x["start"], x["end"])) + + # Apply speaker filtering if requested + if num_speakers is not None or max_speakers is not None: + target_speakers = num_speakers if num_speakers is not None else max_speakers + durations: Dict[str, float] = {} + for seg in segments: + spk = str(seg["speaker"]) + durations[spk] = durations.get(spk, 0.0) + (float(seg["end"]) - float(seg["start"])) + + sorted_spks = sorted(durations.keys(), key=lambda k: durations[k], reverse=True) + keep_spks = set(sorted_spks[:target_speakers]) + segments = [seg for seg in segments if seg["speaker"] in keep_spks] + + return segments + + def to_rttm(self, segments: List[Dict[str, Union[float, str]]], uri: str = "audio") -> str: + """Convert speaker turns to NIST RTTM format.""" + return segments_to_rttm(segments, uri=uri) + + +# Friendly alias +Nemotron3Diarization = NemotronDiarization + + +def extract_speaker_dict(probs: np.ndarray, threshold: float = 0.5) -> List[Dict[str, Union[float, int]]]: + """ + Extract speaker turns dictionary from frame-level probabilities. + + :param np.ndarray probs: Array of shape (1, num_frames, num_speakers) + :param float threshold: Activation threshold (default: 0.5) + :return: List of dicts with 'Start', 'End', and 'Speaker' (int) keys. + """ + active = (probs > threshold).astype(np.int32) + boundary = np.zeros((active.shape[0], 1, active.shape[2]), dtype=np.int32) + changes = np.diff(np.concatenate([boundary, active, boundary], axis=1), axis=1) + + segments = [] + for speaker in range(changes.shape[2]): + starts = np.nonzero(changes[0, :, speaker] == 1)[0] + ends = np.nonzero(changes[0, :, speaker] == -1)[0] + for start, end in zip(starts, ends): + segments.append({ + "Start": round(start * NEMOTRON_FRAME_DURATION, 2), + "End": round(end * NEMOTRON_FRAME_DURATION, 2), + "Speaker": int(speaker), + }) + segments.sort(key=lambda seg: (seg["Start"], seg["Speaker"])) + return segments + + +def segments_to_rttm( + segments: List[Dict[str, Union[float, str]]], + uri: str = "audio", +) -> str: + """ + Convert diarization segments to standard NIST RTTM string format. + + :param list segments: List of segment dicts (with 'start'/'end'/'speaker' or 'Start'/'End'/'Speaker') + :param str uri: Recording URI / identifier + :return: Formatted RTTM string + """ + lines = [] + for seg in segments: + s = float(seg.get("start", seg.get("Start", 0.0))) + e = float(seg.get("end", seg.get("End", 0.0))) + dur = round(e - s, 3) + if dur <= 0: + continue + spk = seg.get("speaker", seg.get("Speaker", "speaker_00")) + if isinstance(spk, int): + spk_label = f"speaker_{spk:02d}" + else: + spk_label = str(spk).lower() + lines.append(f"SPEAKER {uri} 1 {s:.3f} {dur:.3f} {spk_label} ") + return "\n".join(lines) + ("\n" if lines else "") + + class Diarization: """ - ONNX-based Speaker Diarization using Pyannote Segmentation. + ONNX-based Speaker Diarization supporting: + - NVIDIA Nemotron-3 Diarization ("nemotron-3-diarization" / "joosthel/Nemotron-3-Diarization-ONNX") (default) + - Pyannote Segmentation 3.0 ("pyannote_segmentation") """ def __init__( self, - model: str = "pyannote_segmentation", + model: str = "nemotron-3-diarization", model_path: Optional[str] = None, device: Optional[str] = None, + precision: str = "int8", + threads: Optional[int] = None, + **kwargs, ) -> None: """ - :param str model: Diarization model identifier (default: "pyannote_segmentation") - :param Optional[str] model_path: Explicit path to ONNX model file + :param str model: Diarization model identifier: + * "nemotron-3-diarization" (default) / "nemotron_diarization" / "joosthel/Nemotron-3-Diarization-ONNX" + * "pyannote_segmentation" - Pyannote Segmentation 3.0 + :param Optional[str] model_path: Explicit path to ONNX model file or directory :param Optional[str] device: Inference device ("cpu", "cuda", "auto") + :param str precision: Model precision for Nemotron ('int8' or 'fp32', default: 'int8') + :param Optional[int] threads: Intra-op num threads for onnxruntime """ if ort is None: raise ImportError( @@ -165,14 +682,26 @@ def __init__( self.model_name = model self.device = device or "auto" - - if model_path is not None and os.path.exists(model_path): - self.model_path = model_path + self.precision = precision + + if self.model_name.lower() in [m.lower() for m in NEMOTRON_DIARIZATION_MODELS]: + self.is_nemotron = True + self.engine = NemotronDiarization( + model_dir=model_path, + precision=precision, + device=self.device, + threads=threads, + **kwargs, + ) else: - self.model_path = get_diarization_model_files() + self.is_nemotron = False + if model_path is not None and os.path.exists(model_path): + self.model_path = model_path + else: + self.model_path = get_diarization_model_files() - self.session = self._init_session(self.model_path, self.device) - self.input_name = self.session.get_inputs()[0].name + self.session = self._init_session(self.model_path, self.device) + self.input_name = self.session.get_inputs()[0].name def _init_session(self, path: str, device: str) -> ort.InferenceSession: sess_options = ort.SessionOptions() @@ -226,14 +755,11 @@ def _sliding_window_inference( if total_samples <= _WINDOW_SAMPLES: chunk_probs = self._infer_chunk(audio) - # Calculate number of valid frames corresponding to actual audio valid_frames = max(1, min(len(chunk_probs), int((total_samples - _OFFSET_SAMPLES) / _STEP_SAMPLES) + 1)) return chunk_probs[:valid_frames] - # First chunk accumulated_probs = self._infer_chunk(audio[:_WINDOW_SAMPLES]) overlap_samples = _WINDOW_SAMPLES - step_samples - # Overlap in output frames overlap_frames = int(overlap_samples / _STEP_SAMPLES) pos = step_samples @@ -241,13 +767,11 @@ def _sliding_window_inference( chunk = audio[pos : pos + _WINDOW_SAMPLES] curr_probs = self._infer_chunk(chunk) - # Match local speakers of curr_probs with accumulated_probs on overlap region overlap_m = min(overlap_frames, len(accumulated_probs), len(curr_probs)) if overlap_m > 0: prev_overlap = accumulated_probs[-overlap_m:, :3] curr_overlap = curr_probs[:overlap_m, :3] - # Find permutation that minimizes mean absolute difference best_diff = float("inf") best_perm = (0, 1, 2) for perm in permutations(range(3)): @@ -256,16 +780,13 @@ def _sliding_window_inference( best_diff = diff best_perm = perm - # Apply best permutation to current chunk curr_probs = curr_probs[:, best_perm] - # Blend overlap frames linearly alpha = np.linspace(0.0, 1.0, overlap_m)[:, None] accumulated_probs[-overlap_m:, :3] = ( (1.0 - alpha) * accumulated_probs[-overlap_m:, :3] + alpha * curr_probs[:overlap_m, :3] ) - # Append new non-overlapping frames new_frames = curr_probs[overlap_m:] if len(new_frames) > 0: accumulated_probs = np.vstack([accumulated_probs, new_frames]) @@ -274,10 +795,34 @@ def _sliding_window_inference( pos += step_samples - # Trim to valid frames for total audio length valid_frames = max(1, min(len(accumulated_probs), int((total_samples - _OFFSET_SAMPLES) / _STEP_SAMPLES) + 1)) return accumulated_probs[:valid_frames] + def predict_proba( + self, + data: Union[str, Path, np.ndarray], + sampling_rate: int = _SAMPLE_RATE, + ) -> np.ndarray: + """ + Get frame-level speaker probabilities. + """ + if getattr(self, "is_nemotron", False): + return self.engine.predict_proba(data, sampling_rate=sampling_rate) + + if isinstance(data, (str, Path)): + audio = load_audio(data, target_sr=_SAMPLE_RATE) + elif isinstance(data, np.ndarray): + audio = data.flatten().astype(np.float32) + if sampling_rate != _SAMPLE_RATE: + audio = resample_audio(audio, orig_sr=sampling_rate, target_sr=_SAMPLE_RATE) + else: + raise TypeError(f"Unsupported data type for diarization: {type(data)}") + + if len(audio) == 0: + return np.zeros((0, 3), dtype=np.float32) + + return self._sliding_window_inference(audio) + def diarize( self, data: Union[str, Path, np.ndarray], @@ -289,10 +834,25 @@ def diarize( offset: float = 0.5, min_duration_on: float = 0.3, min_duration_off: float = 0.5, + **kwargs, ) -> List[Dict[str, Union[float, str]]]: """ Diarize audio data and return list of speaker turns. """ + if getattr(self, "is_nemotron", False): + return self.engine.diarize( + data=data, + sampling_rate=sampling_rate, + num_speakers=num_speakers, + min_speakers=min_speakers, + max_speakers=max_speakers, + onset=onset, + offset=offset, + min_duration_on=min_duration_on, + min_duration_off=min_duration_off, + **kwargs, + ) + if isinstance(data, (str, Path)): audio = load_audio(data, target_sr=_SAMPLE_RATE) elif isinstance(data, np.ndarray): @@ -316,10 +876,8 @@ def diarize( min_duration_off=min_duration_off, ) - # Apply speaker filtering if requested if num_speakers is not None or max_speakers is not None: target_speakers = num_speakers if num_speakers is not None else max_speakers - # Rank speakers by total speech duration durations: Dict[str, float] = {} for seg in segments: spk = str(seg["speaker"]) @@ -334,8 +892,10 @@ def diarize( # Module-level cache _diarizer_instance: Optional[Diarization] = None -_diarizer_model: str = "pyannote_segmentation" +_diarizer_model: str = "nemotron-3-diarization" _diarizer_device: Optional[str] = None +_diarizer_precision: str = "int8" + def merge_same_speaker_segments( @@ -410,8 +970,9 @@ def _run_sherpa_onnx_diarize( def diarize( data: Union[str, Path, np.ndarray], - model: str = "pyannote_segmentation", + model: str = "nemotron-3-diarization", device: Optional[str] = None, + precision: str = "int8", sampling_rate: int = _SAMPLE_RATE, num_speakers: Optional[int] = None, min_speakers: Optional[int] = None, @@ -427,8 +988,10 @@ def diarize( Perform speaker diarization on audio data using ONNX. :param Union[str, Path, np.ndarray] data: Path to sound file or numpy array of audio. - :param str model: Diarization model name (default: "pyannote_segmentation"). + :param str model: Diarization model name (default: "nemotron-3-diarization"). + Options: "nemotron-3-diarization", "pyannote_segmentation", etc. :param Optional[str] device: Inference device ("cpu", "cuda", "auto"). + :param str precision: Model precision for Nemotron ('int8' or 'fp32', default: 'int8'). :param int sampling_rate: Audio sample rate (default: 16000). :param Optional[int] num_speakers: Exact number of speakers if known. :param Optional[int] min_speakers: Minimum number of speakers. @@ -461,11 +1024,17 @@ def diarize( if backend != "onnx": raise ValueError(f"Unknown backend '{backend}'. Supported backends are 'onnx' and 'sherpa-onnx'.") - global _diarizer_instance, _diarizer_model, _diarizer_device - if _diarizer_instance is None or _diarizer_model != model or _diarizer_device != device: - _diarizer_instance = Diarization(model=model, device=device) + global _diarizer_instance, _diarizer_model, _diarizer_device, _diarizer_precision + if ( + _diarizer_instance is None + or _diarizer_model != model + or _diarizer_device != device + or _diarizer_precision != precision + ): + _diarizer_instance = Diarization(model=model, device=device, precision=precision, **kwargs) _diarizer_model = model _diarizer_device = device + _diarizer_precision = precision return _diarizer_instance.diarize( data=data, @@ -483,8 +1052,9 @@ def diarize( def asr_diarize( data: Union[str, Path, np.ndarray], asr_model: str = "typhoon_asr", - diarize_model: str = "pyannote_segmentation", + diarize_model: str = "nemotron-3-diarization", device: Optional[str] = None, + precision: str = "int8", sampling_rate: int = _SAMPLE_RATE, lm: bool = False, num_speakers: Optional[int] = None, @@ -505,8 +1075,9 @@ def asr_diarize( :param Union[str, Path, np.ndarray] data: Path to sound file or numpy array of audio. :param str asr_model: The ASR model name (default: "typhoon_asr"). - :param str diarize_model: Diarization model name (default: "pyannote_segmentation"). + :param str diarize_model: Diarization model name (default: "nemotron-3-diarization"). :param Optional[str] device: Inference device ("cpu", "cuda", "auto"). + :param str precision: Model precision for Nemotron ('int8' or 'fp32', default: 'int8'). :param int sampling_rate: Audio sample rate (default: 16000). :param bool lm: Use language model for ASR if supported. :param Optional[int] num_speakers: Exact number of speakers if known. @@ -527,7 +1098,7 @@ def asr_diarize( from pythaiasr import asr_diarize - turns = asr_diarize("meeting.wav", asr_model="typhoon_asr") + turns = asr_diarize("meeting.wav", asr_model="typhoon_asr", diarize_model="nemotron-3-diarization") for turn in turns: print(f"[{turn['start']:.2f}s - {turn['end']:.2f}s] {turn['speaker']}: {turn['text']}") """ @@ -549,6 +1120,7 @@ def asr_diarize( data=audio, model=diarize_model, device=device, + precision=precision, sampling_rate=sampling_rate, num_speakers=num_speakers, min_speakers=min_speakers, @@ -599,3 +1171,16 @@ def asr_diarize( return results + +__all__ = [ + "Diarization", + "NemotronDiarization", + "Nemotron3Diarization", + "diarize", + "asr_diarize", + "merge_same_speaker_segments", + "segments_to_rttm", + "extract_speaker_dict", +] + + diff --git a/pythaiasr/download.py b/pythaiasr/download.py index 3b4a019..7ab5ee5 100644 --- a/pythaiasr/download.py +++ b/pythaiasr/download.py @@ -28,6 +28,20 @@ "segmentation": "segmentation-3.0.onnx", } +DEFAULT_NEMOTRON_DIARIZATION_URLS = { + "preprocessor": "https://huggingface.co/joosthel/Nemotron-3-Diarization-ONNX/resolve/main/preprocessor_core.onnx", + "model_int8": "https://huggingface.co/joosthel/Nemotron-3-Diarization-ONNX/resolve/main/model.int8.onnx", + "model_fp32": "https://huggingface.co/joosthel/Nemotron-3-Diarization-ONNX/resolve/main/model.onnx", + "constants": "https://huggingface.co/joosthel/Nemotron-3-Diarization-ONNX/resolve/main/constants.npz", +} + +DEFAULT_NEMOTRON_DIARIZATION_FILENAMES = { + "preprocessor": "preprocessor_core.onnx", + "model_int8": "model.int8.onnx", + "model_fp32": "model.onnx", + "constants": "constants.npz", +} + def get_pythaiasr_path() -> str: """ @@ -240,3 +254,86 @@ def get_diarization_model_files( download_file(urls["segmentation"], seg_path) return seg_path + + +def get_nemotron_diarization_model_files( + model_dir: Optional[str] = None, + urls: Optional[Dict[str, str]] = None, + precision: str = "int8", +) -> Tuple[str, str, str]: + """ + Ensure Nemotron-3 Diarization ONNX model files are present. + Downloads to `~/pythaiasr-data/nemotron-3-diarization-onnx/` if missing. + + :param model_dir: Custom directory to store or load model files. + :param urls: Custom dictionary with URLs for 'preprocessor', 'model_int8', 'model_fp32', 'constants'. + :param precision: Model precision: 'int8' (default, ~104 MB) or 'fp32' (~397 MB). + :return: Tuple of (preprocessor_path, model_path, constants_path). + """ + if model_dir is None: + root_data = get_pythaiasr_path() + model_dir = os.path.join(root_data, "nemotron-3-diarization-onnx") + else: + model_dir = os.path.abspath(os.path.expanduser(model_dir)) + + try: + os.makedirs(model_dir, exist_ok=True) + except OSError: + pass + + urls = urls or DEFAULT_NEMOTRON_DIARIZATION_URLS + model_key = "model_fp32" if precision == "fp32" else "model_int8" + model_filename = DEFAULT_NEMOTRON_DIARIZATION_FILENAMES[model_key] + prep_filename = DEFAULT_NEMOTRON_DIARIZATION_FILENAMES["preprocessor"] + const_filename = DEFAULT_NEMOTRON_DIARIZATION_FILENAMES["constants"] + + prep_path = os.path.join(model_dir, prep_filename) + model_path = os.path.join(model_dir, model_filename) + const_path = os.path.join(model_dir, const_filename) + + # Check if files exist in target model_dir + if os.path.exists(prep_path) and os.path.exists(model_path) and os.path.exists(const_path): + return prep_path, model_path, const_path + + # Check local candidate paths + local_candidates = [ + os.path.abspath("nemotron_diarization"), + os.path.abspath("nemotron-3-diarization-onnx"), + os.path.abspath("Nemotron-3-Diarization-ONNX"), + os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "nemotron_diarization")), + os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "nemotron-3-diarization-onnx")), + os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "Nemotron-3-Diarization-ONNX")), + ] + + for candidate in local_candidates: + if os.path.isdir(candidate): + cand_prep = os.path.join(candidate, prep_filename) + cand_model = os.path.join(candidate, model_filename) + cand_const = os.path.join(candidate, const_filename) + if os.path.exists(cand_prep) and os.path.exists(cand_model) and os.path.exists(cand_const): + try: + os.makedirs(model_dir, exist_ok=True) + if not os.path.exists(prep_path): + shutil.copy2(cand_prep, prep_path) + if not os.path.exists(model_path): + shutil.copy2(cand_model, model_path) + if not os.path.exists(const_path): + shutil.copy2(cand_const, const_path) + except OSError: + return cand_prep, cand_model, cand_const + return prep_path, model_path, const_path + + # Download missing files + try: + os.makedirs(model_dir, exist_ok=True) + except OSError: + pass + + if not os.path.exists(prep_path): + download_file(urls["preprocessor"], prep_path) + if not os.path.exists(model_path): + download_file(urls[model_key], model_path) + if not os.path.exists(const_path): + download_file(urls["constants"], const_path) + + return prep_path, model_path, const_path diff --git a/tests/test_asr.py b/tests/test_asr.py index 44f75c5..dfb6d8f 100644 --- a/tests/test_asr.py +++ b/tests/test_asr.py @@ -99,8 +99,7 @@ def test_stream_asr_import(self): def test_stream_asr_without_pyaudio(self): """Test that stream_asr raises ImportError when pyaudio is not available""" pyaudio_backup = sys.modules.get('pyaudio') - if 'pyaudio' in sys.modules: - del sys.modules['pyaudio'] + sys.modules['pyaudio'] = None try: gen = stream_asr(device="cpu") @@ -112,6 +111,8 @@ def test_stream_asr_without_pyaudio(self): finally: if pyaudio_backup is not None: sys.modules['pyaudio'] = pyaudio_backup + else: + sys.modules.pop('pyaudio', None) def test_typhoon_path_resolution(self): """Test root user path defaults to ~/pythaiasr-data and respects env var.""" diff --git a/tests/test_diarize.py b/tests/test_diarize.py index 050a7a1..7ec92f6 100644 --- a/tests/test_diarize.py +++ b/tests/test_diarize.py @@ -8,20 +8,30 @@ from pythaiasr import ( Diarization, + NemotronDiarization, + Nemotron3Diarization, diarize, asr_diarize, merge_same_speaker_segments, + segments_to_rttm, + extract_speaker_dict, get_diarization_model_files, + get_nemotron_diarization_model_files, ) from pythaiasr.diarization import ( _powerset_to_multilabel, _binarize_timeline, + _sigmoid, + _stable_topk_indices, + NumpySpeakerCache, _SAMPLE_RATE, _STEP_SAMPLES, _OFFSET_SAMPLES, + NEMOTRON_FRAME_DURATION, ) TEST_WAV_FILE = os.path.join(".", "tests", "test.wav") +TEST_DIARIZE_FILE = os.path.join(".", "tests", "test-diarize.wav") COMMON_VOICE_FILE = os.path.join(".", "tests", "common_voice_th_25686161.wav") @@ -144,7 +154,7 @@ def mock_run(output_names, input_feed): with patch.object(Diarization, "_init_session", return_value=mock_session): with patch("pythaiasr.diarization.get_diarization_model_files", return_value="dummy_path.onnx"): - diarizer = Diarization() + diarizer = Diarization(model="pyannote_segmentation") # 1. Short audio (3 seconds = 48,000 samples) short_audio = np.random.randn(48000).astype(np.float32) @@ -225,11 +235,11 @@ def test_empty_audio_handling(self): mock_session = MagicMock() with patch.object(Diarization, "_init_session", return_value=mock_session): with patch("pythaiasr.diarization.get_diarization_model_files", return_value="dummy.onnx"): - diarizer = Diarization() + diarizer = Diarization(model="pyannote_segmentation") res = diarizer.diarize(empty_audio) self.assertEqual(res, []) - res_asr = asr_diarize(empty_audio) + res_asr = asr_diarize(empty_audio, diarize_model="pyannote_segmentation") self.assertEqual(res_asr, []) @patch("pythaiasr.diarization.diarize") @@ -279,7 +289,7 @@ def mock_run(output_names, input_feed): with patch.object(Diarization, "_init_session", return_value=mock_session): with patch("pythaiasr.diarization.get_diarization_model_files", return_value="dummy.onnx"): - diarizer = Diarization() + diarizer = Diarization(model="pyannote_segmentation") segments = diarizer.diarize(audio_24k, sampling_rate=24000) self.assertIsInstance(segments, list) self.assertGreater(len(segments), 0) @@ -295,6 +305,173 @@ def test_get_diarization_model_files_resolution(self): path = get_diarization_model_files(model_dir=tmpdir) self.assertEqual(path, dummy_seg) + def test_nemotron_imports_and_aliases(self): + """Verify Nemotron diarization classes and functions are available.""" + self.assertTrue(callable(NemotronDiarization)) + self.assertTrue(callable(Nemotron3Diarization)) + self.assertIs(Nemotron3Diarization, NemotronDiarization) + self.assertTrue(callable(segments_to_rttm)) + self.assertTrue(callable(extract_speaker_dict)) + self.assertTrue(callable(get_nemotron_diarization_model_files)) + + def test_extract_speaker_dict(self): + """Test extract_speaker_dict converts frame probabilities to turns.""" + probs = np.zeros((1, 200, 8), dtype=np.float32) + # Speaker 0 active from frame 10 to 30 (0.10s to 0.30s) + probs[0, 10:30, 0] = 0.95 + # Speaker 1 active from frame 50 to 90 (0.50s to 0.90s) + probs[0, 50:90, 1] = 0.95 + + segments = extract_speaker_dict(probs, threshold=0.5) + self.assertEqual(len(segments), 2) + self.assertEqual(segments[0]["Speaker"], 0) + self.assertAlmostEqual(segments[0]["Start"], 0.10, places=2) + self.assertAlmostEqual(segments[0]["End"], 0.30, places=2) + + self.assertEqual(segments[1]["Speaker"], 1) + self.assertAlmostEqual(segments[1]["Start"], 0.50, places=2) + self.assertAlmostEqual(segments[1]["End"], 0.90, places=2) + + def test_segments_to_rttm(self): + """Test segments_to_rttm output formatting.""" + segments = [ + {"start": 0.5, "end": 2.0, "speaker": "SPEAKER_00"}, + {"Start": 2.5, "End": 4.0, "Speaker": 1}, + ] + rttm = segments_to_rttm(segments, uri="recording_01") + lines = rttm.strip().split("\n") + self.assertEqual(len(lines), 2) + self.assertTrue(lines[0].startswith("SPEAKER recording_01 1 0.500 1.500 speaker_00")) + self.assertTrue(lines[1].startswith("SPEAKER recording_01 1 2.500 1.500 speaker_01")) + + def test_numpy_speaker_cache(self): + """Test NumpySpeakerCache initialization and buffer updates.""" + constants = { + "hidden_size": np.array(192), + "num_speakers": np.array(8), + "subsampling_factor": np.array(8), + "speaker_cache_length": np.array(64), + "speaker_cache_silence_frames_per_speaker": np.array(2), + "prediction_score_threshold": np.array(0.5), + "latest_frames_score_boost": np.array(1.0), + "silence_embeds": np.zeros((1, 16, 192), dtype=np.float32), + "min_positive_scores_rate": np.array(0.1), + "strong_boost_rate": np.array(0.1), + "weak_boost_rate": np.array(0.1), + } + cache = NumpySpeakerCache( + constants=constants, + fifo_length=8, + speaker_cache_update_period=4, + ) + # Initially empty + init_embeds = cache.get_embeds() + self.assertEqual(init_embeds.shape, (1, 0, 192)) + + # Push a chunk: 16 frames subsampled, 128 frames logits + chunk_embeds = np.random.randn(1, 16, 192).astype(np.float32) + chunk_logits = np.random.randn(1, 128, 8).astype(np.float32) + cache.update(chunk_embeds, chunk_logits, num_chunk_frames=16) + + new_embeds = cache.get_embeds() + self.assertGreater(new_embeds.shape[1], 0) + self.assertEqual(new_embeds.shape[2], 192) + + @patch("onnxruntime.InferenceSession") + def test_mock_nemotron_diarization(self, mock_ort_session): + """Test NemotronDiarization inference flow with mocked ONNX runtime.""" + mock_prep = MagicMock() + mock_model = MagicMock() + mock_ort_session.side_effect = [mock_prep, mock_model] + + def mock_prep_run(output_names, input_feed): + sig_len = input_feed["preemphasized"].shape[1] + mel_frames = 1 + sig_len // 160 + 10 + return [np.zeros((1, mel_frames, 128), dtype=np.float32)] + + def mock_model_run(output_names, input_feed): + ctx_len = int(input_feed["context_length"]) + mel_len = int(input_feed["chunk_mel_length"]) + num_frames = -(-mel_len // 8) + total_embeds = ctx_len + num_frames + logits = np.zeros((1, total_embeds * 8, 8), dtype=np.float32) + embeds = np.zeros((1, total_embeds, 192), dtype=np.float32) + return logits, embeds + + mock_prep.run.side_effect = mock_prep_run + mock_model.run.side_effect = mock_model_run + + dummy_constants = { + "hidden_size": np.array(192), + "num_speakers": np.array(8), + "subsampling_factor": np.array(8), + "speaker_cache_length": np.array(64), + "speaker_cache_silence_frames_per_speaker": np.array(2), + "prediction_score_threshold": np.array(0.5), + "latest_frames_score_boost": np.array(1.0), + "silence_embeds": np.zeros((1, 16, 192), dtype=np.float32), + "min_positive_scores_rate": np.array(0.1), + "strong_boost_rate": np.array(0.1), + "weak_boost_rate": np.array(0.1), + "chunk_length": np.array(64), + "chunk_right_context": np.array(8), + "fifo_length": np.array(8), + "speaker_cache_update_period": np.array(4), + } + + with patch("pythaiasr.diarization.get_nemotron_diarization_model_files", return_value=("prep.onnx", "const.npz", "model.onnx")): + with patch("numpy.load", return_value=dummy_constants): + engine = NemotronDiarization(device="cpu") + self.assertIsNotNone(engine) + + # Test predict_proba + audio = np.zeros(16000, dtype=np.float32) + probs = engine.predict_proba(audio) + self.assertEqual(probs.ndim, 3) + self.assertEqual(probs.shape[-1], 8) + + # Test diarize + segments = engine.diarize(audio) + self.assertIsInstance(segments, list) + + def test_diarization_class_nemotron_dispatch(self): + """Test Diarization wrapper dispatches default model and nemotron aliases to NemotronDiarization.""" + with patch("pythaiasr.diarization.NemotronDiarization") as mock_engine_class: + mock_inst = MagicMock() + mock_engine_class.return_value = mock_inst + + # Default model should be nemotron + diarizer_default = Diarization() + self.assertTrue(diarizer_default.is_nemotron) + self.assertEqual(diarizer_default.model_name, "nemotron-3-diarization") + + # Explicit nemotron model name + diarizer = Diarization(model="nemotron-3-diarization", precision="int8") + self.assertTrue(diarizer.is_nemotron) + + diarizer.diarize(np.zeros(16000, dtype=np.float32)) + mock_inst.diarize.assert_called_once() + + def test_real_nemotron_diarize_and_rttm(self): + """Test real Nemotron diarization on tests/test-diarize.wav if model files are cached.""" + home = os.path.expanduser("~") + cache_dir = os.path.join(home, "pythaiasr-data", "nemotron-3-diarization-onnx") + int8_model = os.path.join(cache_dir, "model.int8.onnx") + + if not (os.path.exists(TEST_DIARIZE_FILE) and os.path.exists(int8_model)): + self.skipTest("Nemotron model cache or test-diarize.wav not present; skipping live inference test.") + + segments = diarize(TEST_DIARIZE_FILE, model="nemotron-3-diarization") + self.assertIsInstance(segments, list) + self.assertGreaterEqual(len(segments), 2) + speakers = {s["speaker"] for s in segments} + self.assertIn("SPEAKER_00", speakers) + self.assertIn("SPEAKER_01", speakers) + + # Convert to RTTM + rttm = segments_to_rttm(segments, uri="test_audio") + self.assertIn("SPEAKER test_audio 1", rttm) + if __name__ == "__main__": unittest.main()