diff --git a/scripts/generate_golden.py b/scripts/generate_golden.py index 50bbdf10d..f3724bf86 100644 --- a/scripts/generate_golden.py +++ b/scripts/generate_golden.py @@ -53,6 +53,30 @@ _MISSING_ATTRIBUTE = object() +def _best_effort_package_commit(package_name: str) -> str | None: + """Best-effort immutable VCS commit for a pip-installed package. + + When a dependency is installed from a git URL (``pip install + git+https://...@``) -- the only way to get pre-release/unreleased + HuggingFace model code such as Nemotron3Diarization's -- pip records the + resolved commit SHA in the package's ``direct_url.json`` metadata. This + lets golden-data provenance capture an immutable source identity instead + of just a mutable, potentially ambiguous ``X.Y.Z.devN`` version string. + Returns ``None`` (rather than raising) for a normal PyPI install, or if + the metadata is missing/malformed, since this is diagnostic-only. + """ + try: + import json + from importlib.metadata import distribution + + direct_url = json.loads( + distribution(package_name).read_text("direct_url.json") or "{}" + ) + return direct_url.get("vcs_info", {}).get("commit_id") + except Exception: + return None + + @contextlib.contextmanager def _temporary_processor_max_pixels(processor: object, max_pixels: int | None): """Temporarily override processor pixel limits and restore their exact state.""" @@ -1907,6 +1931,207 @@ def _hook(_module, _args, kwargs, output): save_drafter_inputs(drafter_inputs_path_for_case(case), arrays) +def _generate_diarization_offline(case: TestCase, json_path: Path, device: str) -> None: + """Generate golden data for offline speaker diarization. + + Diarization models emit continuous per-frame per-speaker sigmoid + probabilities, not a vocab logit vector, so the reference is stored via + ``save_diarization_golden()`` (full-precision ``probs`` array in a + companion ``.npz``) and compared with ``compare_diarization_golden()`` + instead of the argmax-gated ``save_golden_ref()``/``compare_golden()`` + used elsewhere in this file. + + Two reference implementations are dispatched by ``case.model_type``: + + * ``sortformer`` (NeMo Sortformer): built from a ``.nemo`` archive via + ``nemo_toolkit`` (not a mobius runtime dependency) -- + ``frontend_encoder`` (mel -> embeddings) then ``forward_infer`` + (embeddings -> per-frame speaker sigmoids). ``mel`` is generated + already channel-first ``[batch, feat, frames]`` -- NeMo's own input + contract -- matching the mobius ``DiarizationTask`` graph directly. + * otherwise (e.g. ``nemotron3_diarization``): a HuggingFace + ``AutoModelForAudioFrameClassification`` checkpoint. HuggingFace's + ``input_features`` are channel-last ``[batch, frames, feat]``; they + are transposed to channel-first before saving so the golden matches + the mobius task's input contract with no runtime transpose needed. + """ + import numpy as np + import torch + + from mobius._testing.golden import save_diarization_golden + + seed = int(case.generation_params.get("seed", 0)) + torch.manual_seed(seed) + + if case.model_type == "sortformer": + import nemo # type: ignore[import-not-found] + from huggingface_hub import hf_hub_download + from nemo.collections.asr.models import ( # type: ignore[import-not-found] + SortformerEncLabelModel, + ) + + filename = case.generation_params.get( + "nemo_filename", "diar_streaming_sortformer_4spk-v2.1.nemo" + ) + nemo_path = hf_hub_download( + repo_id=case.model_id, filename=filename, revision=case.revision + ) + model = SortformerEncLabelModel.restore_from(nemo_path, map_location=device) + model.eval() + model.streaming_mode = False # offline: full-context attention. + + num_frames = int(case.generation_params.get("num_frames", 400)) + feat_dim = int(model.cfg.encoder.feat_in) + mel = torch.randn(1, feat_dim, num_frames, device=device) + mel_len = torch.tensor([num_frames], dtype=torch.long, device=device) + with torch.no_grad(): + emb_seq, emb_len = model.frontend_encoder( + processed_signal=mel, processed_signal_length=mel_len + ) + preds = model.forward_infer(emb_seq, emb_len) + + arrays = { + "mel": mel.cpu().numpy().astype(np.float32), + "emb_seq": emb_seq.cpu().numpy().astype(np.float32), + "emb_len": emb_len.cpu().numpy().astype(np.int64), + "probs": preds.cpu().numpy().astype(np.float32), + } + provenance = { + "model_id": case.model_id, + "revision": case.revision, + "nemo_version": nemo.__version__, + "nemo_commit": _best_effort_package_commit("nemo_toolkit"), + "seed": seed, + "feat_dim": feat_dim, + "num_frames": num_frames, + "num_speakers": int(preds.shape[-1]), + } + else: + import transformers + from transformers import AutoModelForAudioFrameClassification + + model = AutoModelForAudioFrameClassification.from_pretrained( + case.model_id, revision=case.revision, dtype=torch.float32 + ) + model = model.to(device) + model.eval() + + mel_dim = model.config.audio_config.num_mel_bins + # Below chunk_length * subsampling_factor so HuggingFace's own + # offline forward does not internally re-chunk. + num_frames = int(case.generation_params.get("num_frames", 2000)) + mel = torch.randn(1, num_frames, mel_dim, device=device) + with torch.no_grad(): + out = model(input_features=mel) + + arrays = { + "mel": np.transpose(mel.cpu().numpy(), (0, 2, 1)).astype(np.float32), + "probs": out.logits.sigmoid().cpu().numpy().astype(np.float32), + } + provenance = { + "model_id": case.model_id, + "revision": case.revision, + "transformers_version": transformers.__version__, + "transformers_commit": _best_effort_package_commit("transformers"), + "torch_version": torch.__version__, + "seed": seed, + "num_speakers": int(model.config.head_config.num_speakers), + "mel_dim": mel_dim, + } + + save_diarization_golden(json_path, arrays=arrays, provenance=provenance) + + +def _generate_diarization_streaming(case: TestCase, json_path: Path, device: str) -> None: + """Generate golden data for a multi-chunk streaming diarization session. + + Runs the real HuggingFace ``Nemotron3DiarizationForAudioFrameClassification`` + checkpoint (the ground-truth reference implementation) chunk by chunk, + threading the HuggingFace ``Nemotron3DiarizationSpeakerCache`` (Arrival- + Order Speaker Cache + FIFO queue) across chunks. Each chunk is a full + ``chunk_length + chunk_right_context`` window (matching the mobius + ``DiarizationStreamingTask`` ONNX graph's fixed input shape). + + NOTE: the HF cache's ``.embeds``/``.probs``/``.fifo`` buffers are mutable + and updated in place via ``index_copy_`` across chunks -- ``.numpy()`` + returns a *view*, so every per-chunk snapshot below is explicitly copied. + Without this, all per-chunk arrays would silently alias the final + chunk's values once ``np.savez`` serializes them (a real bug caught while + first authoring this golden -- see git history for + ``scripts/generate_nemotron3_diarization_golden.py``). + """ + import numpy as np + import torch + import transformers + from transformers import AutoModelForAudioFrameClassification + + from mobius._testing.golden import save_diarization_golden + + seed = int(case.generation_params.get("seed", 0)) + num_chunks = int(case.generation_params.get("num_chunks", 3)) + + model = AutoModelForAudioFrameClassification.from_pretrained( + case.model_id, revision=case.revision, dtype=torch.float32 + ) + model = model.to(device) + model.eval() + + mel_dim = model.config.audio_config.num_mel_bins + chunk_length = model.config.chunk_length + chunk_right_context = model.config.chunk_right_context + subsampling = model.config.audio_config.subsampling_factor + raw_window = (chunk_length + chunk_right_context) * subsampling + + arrays: dict[str, np.ndarray] = {"num_chunks": np.array(num_chunks, dtype=np.int64)} + cache = None + for i in range(num_chunks): + torch.manual_seed(seed + 100 + i) + mel = torch.randn(1, raw_window, mel_dim, device=device) + is_last = i == num_chunks - 1 + lookahead = 0 if is_last else chunk_right_context + with torch.no_grad(): + out = model( + input_features=mel, speaker_cache=cache, num_lookahead_frames=lookahead + ) + cache = out.speaker_cache + # HF's input_features are channel-last; transpose to the mobius + # DiarizationStreamingTask's channel-first contract. + arrays[f"chunk{i}_mel"] = ( + np.transpose(mel.cpu().numpy(), (0, 2, 1)).astype(np.float32).copy() + ) + arrays[f"chunk{i}_lookahead"] = np.array(lookahead, dtype=np.int64) + arrays[f"chunk{i}_probs"] = out.logits.sigmoid().cpu().numpy().copy() + arrays[f"chunk{i}_cache_embeds"] = cache.embeds.cpu().numpy().copy() + arrays[f"chunk{i}_cache_probs"] = cache.probs.cpu().numpy().copy() + arrays[f"chunk{i}_fifo"] = cache.fifo.cpu().numpy().copy() + arrays[f"chunk{i}_num_cache_frames"] = np.array(cache.num_cache_frames, dtype=np.int64) + arrays[f"chunk{i}_num_fifo_frames"] = np.array(cache.num_fifo_frames, dtype=np.int64) + arrays[f"chunk{i}_is_compressed"] = np.array(cache.is_compressed, dtype=np.bool_) + + provenance = { + "model_id": case.model_id, + "revision": case.revision, + "transformers_version": transformers.__version__, + "transformers_commit": _best_effort_package_commit("transformers"), + "torch_version": torch.__version__, + "seed": seed, + "num_speakers": int(model.config.head_config.num_speakers), + "mel_dim": mel_dim, + "chunk_length": int(chunk_length), + "chunk_right_context": int(chunk_right_context), + "subsampling_factor": int(subsampling), + "streaming_fifo_length": int(model.config.streaming_config.fifo_length), + "streaming_speaker_cache_length": int( + model.config.streaming_config.speaker_cache_length + ), + "streaming_speaker_cache_update_period": int( + model.config.streaming_config.speaker_cache_update_period + ), + "num_stream_chunks": num_chunks, + } + save_diarization_golden(json_path, arrays=arrays, provenance=provenance) + + # ---- Dispatcher ---- # Map task_type strings to generator functions. @@ -1930,6 +2155,8 @@ def _hook(_module, _args, kwargs, output): "phi4mm-multimodal": _generate_phi4mm_multimodal, "gemma4-assistant": _generate_gemma4_assistant, "dflash-draft": _generate_dflash_draft, + "diarization": _generate_diarization_offline, + "diarization-streaming": _generate_diarization_streaming, } diff --git a/scripts/generate_sortformer_golden.py b/scripts/generate_sortformer_golden.py deleted file mode 100644 index bc5bde225..000000000 --- a/scripts/generate_sortformer_golden.py +++ /dev/null @@ -1,106 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT License. - -"""Regenerate the Sortformer diarization golden reference used by the L4/L5 tests. - -This produces ``testdata/golden/speech/sortformer_diarization.npz`` by running -the *real* NeMo Sortformer model through the NeMo toolkit (the ground-truth -reference implementation). It must be run inside an environment that has -``nemo_toolkit`` installed (it is **not** a mobius runtime dependency):: - - python -m venv /tmp/nemo_ref_venv - source /tmp/nemo_ref_venv/bin/activate - pip install "nemo_toolkit[asr]" - python scripts/generate_sortformer_golden.py \ - --model nvidia/diar_streaming_sortformer_4spk-v2.1 \ - --revision fafaab5faa1617a0ca52d38dd3dc4bd636800d3d \ - --out testdata/golden/speech/sortformer_diarization.npz - -The offline forward path is ``frontend_encoder`` (mel features -> embedding -sequence) followed by ``forward_infer`` (embeddings -> per-frame speaker -activity sigmoids). The committed ``.npz`` stores the mel input, the encoder -embeddings, and the speaker probabilities, plus a ``meta`` JSON blob (model id, -revision, NeMo version, dtype, seed) so the reference is self-describing and -auditable. -""" - -from __future__ import annotations - -import argparse -import json - -import numpy as np -import torch - -# Deterministic mel-feature fixture (also recorded in metadata). -_SEED = 0 -_T = 400 # mel frames; with 8x subsampling -> 50 output diarization frames. - - -def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--model", default="nvidia/diar_streaming_sortformer_4spk-v2.1") - parser.add_argument( - "--revision", - default="fafaab5faa1617a0ca52d38dd3dc4bd636800d3d", - help="HuggingFace Hub commit SHA to pin the reference model.", - ) - parser.add_argument( - "--out", - default="testdata/golden/speech/sortformer_diarization.npz", - ) - args = parser.parse_args() - - import nemo # type: ignore[import-not-found] - from huggingface_hub import hf_hub_download - from nemo.collections.asr.models import ( # type: ignore[import-not-found] - SortformerEncLabelModel, - ) - - torch.manual_seed(_SEED) - - nemo_path = hf_hub_download( - repo_id=args.model, - filename="diar_streaming_sortformer_4spk-v2.1.nemo", - revision=args.revision, - ) - model = SortformerEncLabelModel.restore_from(nemo_path, map_location="cpu") - model.eval() - # Offline (non-streaming) forward path: full-context attention. - model.streaming_mode = False - - feat_dim = int(model.cfg.encoder.feat_in) - mel = torch.randn(1, feat_dim, _T) - mel_len = torch.tensor([_T], dtype=torch.long) - - with torch.no_grad(): - emb_seq, emb_len = model.frontend_encoder( - processed_signal=mel, processed_signal_length=mel_len - ) - preds = model.forward_infer(emb_seq, emb_len) - - num_spks = int(preds.shape[-1]) - meta = { - "model_id": args.model, - "revision": args.revision, - "nemo_version": nemo.__version__, - "dtype": "float32", - "seed": _SEED, - "feat_dim": feat_dim, - "input_frames": _T, - "num_spks": num_spks, - } - - np.savez_compressed( - args.out, - mel=mel.numpy().astype(np.float32), - emb_seq=emb_seq.numpy().astype(np.float32), - emb_len=emb_len.numpy().astype(np.int64), - preds=preds.numpy().astype(np.float32), - meta=np.array(json.dumps(meta)), - ) - print(f"saved {args.out}\n{json.dumps(meta, indent=2)}") - - -if __name__ == "__main__": - main() diff --git a/src/mobius/_configs/__init__.py b/src/mobius/_configs/__init__.py index 3a59a6e94..18a3c7bcf 100644 --- a/src/mobius/_configs/__init__.py +++ b/src/mobius/_configs/__init__.py @@ -62,6 +62,7 @@ MoonshineStreamingConfig, MuseGlimmerConfig, NanoChatConfig, + Nemotron3DiarizationConfig, NemotronHConfig, NemotronParseConfig, ParakeetCTCConfig, @@ -167,6 +168,7 @@ "MoonshineStreamingConfig", "MuseGlimmerConfig", "NanoChatConfig", + "Nemotron3DiarizationConfig", "NemotronParseConfig", "NemotronHConfig", "ParakeetCTCConfig", diff --git a/src/mobius/_configs/_base.py b/src/mobius/_configs/_base.py index b34eeda16..61a7e15ee 100644 --- a/src/mobius/_configs/_base.py +++ b/src/mobius/_configs/_base.py @@ -4987,3 +4987,92 @@ def from_transformers(cls, config, parent_config=None) -> ParakeetCTCConfig: scale_input=getattr(encoder, "scale_input", True), layer_norm_eps=getattr(encoder, "layer_norm_eps", 1e-5), ) + + +@dataclasses.dataclass +class Nemotron3DiarizationConfig(ArchitectureConfig): + """Configuration for HuggingFace ``Nemotron3DiarizationForAudioFrameClassification``. + + Extracted from the nested ``audio_config`` (a bidirectional, partial-RoPE + Transformer encoder — its fields populate the base :class:`ArchitectureConfig` + fields directly), ``head_config`` (the speaker-classification head that + projects, upsamples, and classifies the encoder output), and + `` streaming_config`` (the Arrival-Order Speaker Cache / FIFO policy shared by + both the offline and streaming forwards' cache-update logic, see + ``models/nemotron3_diarization.py`` for both the offline (chunked + internally via an ONNX ``Loop``, not a single full-sequence pass) and + streaming (one chunk per call, stateful across calls) exports). + """ + + feat_in: int = 128 + subsampling_factor: int = 8 + head_hidden_size: int = 192 + num_speakers: int = 8 + + # Offline chunking (also used as the default streaming step window). + chunk_length: int = 340 + chunk_right_context: int = 40 + + # Offline-only AOSC/FIFO policy, from the top-level ``config`` (distinct + # from ``streaming_config``'s values below): the offline forward reuses + # ``streaming_speaker_cache_length`` (AOSC capacity) but has its own + # ``fifo_length``/``speaker_cache_update_period``. + offline_fifo_length: int = 40 + offline_speaker_cache_update_period: int = 300 + + # Arrival-Order Speaker Cache (AOSC) + FIFO queue policy, from + # ``config.streaming_config``. Shared by both the offline and streaming + # forwards (via ``_run_chunk_and_update_cache``'s cache-update/compression + # logic) — only ``fifo_length``/``speaker_cache_update_period`` differ + # per mode (offline uses its own ``offline_*`` pair above instead). + streaming_fifo_length: int = 264 + streaming_speaker_cache_length: int = 264 + streaming_speaker_cache_update_period: int = 222 + streaming_silence_frames_per_speaker: int = 1 + streaming_prediction_score_threshold: float = 0.25 + streaming_latest_frames_score_boost: float = 0.05 + streaming_min_positive_scores_rate: float = 0.5 + streaming_strong_boost_rate: float = 0.75 + streaming_weak_boost_rate: float = 1.5 + + @classmethod + def from_transformers(cls, config, parent_config=None) -> Nemotron3DiarizationConfig: + """Extract the nested Nemotron3Diarization audio, head, and streaming configs.""" + audio = config.audio_config + head = config.head_config + streaming = getattr(config, "streaming_config", None) + base = ArchitectureConfig.from_transformers(audio, parent_config=config) + fields = _shallow_fields(base) + fields.update(model_type=getattr(config, "model_type", fields["model_type"])) + return cls( + **fields, + feat_in=getattr(audio, "num_mel_bins", 128), + subsampling_factor=getattr(audio, "subsampling_factor", 8), + head_hidden_size=getattr(head, "hidden_size", 192), + num_speakers=getattr(head, "num_speakers", 8), + chunk_length=getattr(config, "chunk_length", 340), + chunk_right_context=getattr(config, "chunk_right_context", 40), + offline_fifo_length=getattr(config, "fifo_length", 40), + offline_speaker_cache_update_period=getattr( + config, "speaker_cache_update_period", 300 + ), + streaming_fifo_length=getattr(streaming, "fifo_length", 264), + streaming_speaker_cache_length=getattr(streaming, "speaker_cache_length", 264), + streaming_speaker_cache_update_period=getattr( + streaming, "speaker_cache_update_period", 222 + ), + streaming_silence_frames_per_speaker=getattr( + streaming, "speaker_cache_silence_frames_per_speaker", 1 + ), + streaming_prediction_score_threshold=getattr( + streaming, "prediction_score_threshold", 0.25 + ), + streaming_latest_frames_score_boost=getattr( + streaming, "latest_frames_score_boost", 0.05 + ), + streaming_min_positive_scores_rate=getattr( + streaming, "min_positive_scores_rate", 0.5 + ), + streaming_strong_boost_rate=getattr(streaming, "strong_boost_rate", 0.75), + streaming_weak_boost_rate=getattr(streaming, "weak_boost_rate", 1.5), + ) diff --git a/src/mobius/_passes/_fold_transpose.py b/src/mobius/_passes/_fold_transpose.py index e7f6cae36..957d57b86 100644 --- a/src/mobius/_passes/_fold_transpose.py +++ b/src/mobius/_passes/_fold_transpose.py @@ -156,7 +156,10 @@ def call(self, model: ir.Model) -> ir.passes.PassResult: # That orphaned initializer is then serialized into the ONNX file # and triggers an ORT warning: # "Removing initializer X. It is not used by any node" - model.graph.remove(node, safe=True) + # `model.graph.all_nodes()` recurses into subgraphs (e.g. a Loop + # body), so `node` may not belong to the top-level `model.graph` + # -- remove it from its own (sub)graph instead. + node.graph.remove(node, safe=True) folded_nodes += 1 modified = True diff --git a/src/mobius/_passes/_fold_transpose_test.py b/src/mobius/_passes/_fold_transpose_test.py index 0e3def2ad..dc64970ef 100644 --- a/src/mobius/_passes/_fold_transpose_test.py +++ b/src/mobius/_passes/_fold_transpose_test.py @@ -558,3 +558,105 @@ def test_transposed_dtype_follows_const_value_when_declared_dtype_missing(self): ) assert packed.const_value.dtype == ir.DataType.FLOAT16 assert packed.const_value.numpy().dtype == np.float16 + + +def test_transpose_inside_loop_body_folds_against_root_initializer(): + """A Transpose inside a Loop body, consuming an outer-scope initializer. + + Folds correctly and registers the new pre-transposed initializer at the + root graph — not a per-name collision risk. + + This documents (and guards) the invariant the pass's ``folded`` cache + relies on: it is keyed by initializer *name* only, and always inserts + into ``model.graph`` (the root), regardless of which (sub)graph the + consuming ``Transpose`` node actually lives in. That is only safe if two + *different* initializers can never share a name across sibling + subgraphs. + + onnxscript's ``OpBuilder`` enforces exactly that: every initializer + (constant or parameter) is always registered on the **root** graph via + ``self._root._graph.register_initializer(...)`` (see + ``onnxscript._internal.builder``), even when the builder is currently + emitting into a Loop/If body. Subgraphs only ever reference outer-scope + initializers by name — they never declare their own local initializers. + So every mobius-built graph has a single, globally unique initializer + namespace by construction, and this pass's name-only cache can never + conflate two distinct initializers. This test builds a Loop body + manually (bypassing the builder) to pin that expectation directly. + """ + weight_data = np.arange(6, dtype=np.float32).reshape(2, 3) + weight_val = ir.Value(name="weight") + weight_val.shape = ir.Shape([2, 3]) + weight_val.dtype = ir.DataType.FLOAT + weight_val.const_value = ir.tensor(weight_data) + + # Loop body: (iter_num, cond_in) -> (cond_out, y), with a Transpose of the + # *outer-scope* root initializer plus a MatMul against the loop-carried x. + iter_num = ir.Value(name="iter_num") + iter_num.dtype = ir.DataType.INT64 + iter_num.shape = ir.Shape([]) + cond_in = ir.Value(name="cond_in") + cond_in.dtype = ir.DataType.BOOL + cond_in.shape = ir.Shape([]) + x_carried = ir.Value(name="x_carried") + x_carried.shape = ir.Shape([2, 2]) + x_carried.dtype = ir.DataType.FLOAT + + body_transpose = ir.Node( + "", + "Transpose", + inputs=[weight_val], + attributes=[ir.Attr("perm", ir.AttributeType.INTS, [1, 0])], + num_outputs=1, + ) + body_w_t = body_transpose.outputs[0] + body_matmul = ir.Node("", "MatMul", inputs=[x_carried, body_w_t], num_outputs=1) + cond_out = ir.Node("", "Identity", inputs=[cond_in], num_outputs=1).outputs[0] + + body_graph = ir.Graph( + inputs=[iter_num, cond_in, x_carried], + outputs=[cond_out, body_matmul.outputs[0]], + nodes=[body_transpose, body_matmul, cond_out.producer()], + name="loop_body", + opset_imports={"": 20}, + ) + + x_init = ir.Value(name="x_init") + x_init.shape = ir.Shape([2, 2]) + x_init.dtype = ir.DataType.FLOAT + trip_count = ir.Value(name="trip_count") + trip_count.dtype = ir.DataType.INT64 + trip_count.shape = ir.Shape([]) + cond = ir.Value(name="cond") + cond.dtype = ir.DataType.BOOL + cond.shape = ir.Shape([]) + + loop_node = ir.Node( + "", + "Loop", + inputs=[trip_count, cond, x_init], + attributes=[ir.Attr("body", ir.AttributeType.GRAPH, body_graph)], + num_outputs=1, + ) + + root_graph = ir.Graph( + inputs=[trip_count, cond, x_init], + outputs=[loop_node.outputs[0]], + nodes=[loop_node], + name="root_graph", + opset_imports={"": 20}, + ) + root_graph.register_initializer(weight_val) + model = ir.Model(root_graph, ir_version=10) + + result = FoldTransposedInitializerPass()(model) + assert result.modified + + # The pre-transposed initializer is registered at the root, and the + # subgraph's Transpose node is gone. + assert "weight_t" in model.graph.initializers + node_types = [n.op_type for n in model.graph.all_nodes()] + assert "Transpose" not in node_types + + w_t = model.graph.initializers["weight_t"] + np.testing.assert_array_equal(w_t.const_value.numpy(), weight_data.T) diff --git a/src/mobius/_registry.py b/src/mobius/_registry.py index 869518b0f..1dde0d8cb 100644 --- a/src/mobius/_registry.py +++ b/src/mobius/_registry.py @@ -122,6 +122,7 @@ MoonshineForConditionalGeneration, MoonshineStreamingForConditionalGeneration, NanoChatCausalLMModel, + Nemotron3DiarizationModel, NemotronCausalLMModel, NemotronParseForConditionalGeneration, OLMo2CausalLMModel, @@ -637,6 +638,7 @@ def _detect_fallback_registration(hf_config) -> ModelRegistration | None: "mpt": ModelRegistration(MPTCausalLMModel), "nanochat": ModelRegistration(NanoChatCausalLMModel), "nemotron": ModelRegistration(NemotronCausalLMModel), + "nemotron3_diarization": ModelRegistration(Nemotron3DiarizationModel, task="diarization"), "olmo": ModelRegistration(OLMoCausalLMModel), "olmo2": ModelRegistration(OLMo2CausalLMModel), "olmo3": ModelRegistration(OLMo2CausalLMModel), @@ -1283,6 +1285,7 @@ def _create_default_registry() -> ModelRegistry: "internlm2": "internlm/internlm2_5-7b-chat", "llama_embed_gguf": "mradermacher/llama-embed-nemotron-8b-GGUF", "nemotron": "nvidia/Nemotron-Mini-4B-Instruct", + "nemotron3_diarization": "nvidia/Nemotron-3-Diarization", "olmo": "allenai/OLMo-1B-hf", "olmo2": "allenai/OLMo-2-1124-7B", "phi": "microsoft/phi-1_5", diff --git a/src/mobius/_testing/golden.py b/src/mobius/_testing/golden.py index bdca71cad..2b6bbf8ad 100644 --- a/src/mobius/_testing/golden.py +++ b/src/mobius/_testing/golden.py @@ -516,6 +516,90 @@ def load_drafter_inputs( return {k: data[k] for k in data.files} +def diarization_probs_path_for_case( + case: GoldenTestCase, + golden_dir: Path = GOLDEN_DIR, +) -> Path: + """Return the companion ``*_probs.npz`` path for a diarization test case. + + Diarization models (``diarization`` / ``diarization-streaming``) emit + continuous per-frame per-speaker sigmoid probabilities rather than a + vocab logit vector, so ``compare_golden()``'s argmax/top-K gate does not + apply and the raw arrays cannot be compactly summarized the way + ``GoldenRef`` summarizes a logit vector. Streaming cases additionally + need the per-chunk Arrival-Order Speaker Cache (AOSC) + FIFO cache-state + tensors so a multi-chunk session can be replayed exactly. Both are + stored in this companion ``.npz`` (see :func:`save_diarization_golden`), + mirroring the ``*_inputs.npz`` pattern used by drafter tasks. + + Maps ``testdata/cases//.yaml`` + to ``testdata/golden//_probs.npz``. + """ + task_dir = case.yaml_path.parent.name + return golden_dir / task_dir / f"{case.case_id}_probs.npz" + + +def save_diarization_golden( + json_path: Path, + *, + arrays: dict[str, np.ndarray], + provenance: dict[str, object] | None = None, +) -> None: + """Save a diarization test case's golden reference. + + The full-precision arrays (per-frame speaker probabilities and, for + streaming cases, the per-chunk cache-state tensors) are stored in a + companion ``*_probs.npz`` next to ``json_path`` (see + :func:`diarization_probs_path_for_case`); ``json_path`` itself records + only lightweight, human-diffable shape/provenance metadata so a reviewer + can sanity-check a regenerated golden without loading the ``.npz``. + + Args: + json_path: Destination path for the metadata ``.json`` file. + arrays: Named reference arrays, e.g. ``{"probs": ...}`` for an + offline case or ``{"chunk0_probs": ..., "chunk0_past_fifo": + ..., ...}`` for a streaming session. + provenance: Immutable source identities used to generate the golden. + """ + json_path = Path(json_path) + json_path.parent.mkdir(parents=True, exist_ok=True) + npz_path = json_path.with_name(json_path.stem + "_probs.npz") + np.savez(npz_path, **arrays) # type: ignore[arg-type] + + data: dict = {"array_shapes": {k: list(v.shape) for k, v in arrays.items()}} + if provenance: + data["provenance"] = provenance + + with open(json_path, "w", encoding="utf-8") as f: + json.dump(data, f, indent=2) + f.write("\n") # trailing newline for POSIX compliance + + +def load_diarization_golden( + case: GoldenTestCase, + golden_dir: Path = GOLDEN_DIR, +) -> dict[str, np.ndarray] | None: + """Load a diarization test case's golden reference arrays. + + Returns the ``.npz`` companion arrays saved by + :func:`save_diarization_golden` (e.g. ``{"probs": ...}`` or the + per-chunk ``chunk{i}_*`` streaming arrays), or ``None`` if either the + metadata ``.json`` or its companion ``.npz`` is missing. + + Args: + case: The test case whose golden to load. + golden_dir: Root directory for golden files. + """ + json_path = golden_path_for_case(case, golden_dir) + if not json_path.exists(): + return None + npz_path = diarization_probs_path_for_case(case, golden_dir) + if not npz_path.exists(): + return None + with np.load(npz_path, allow_pickle=False) as data: + return {k: data[k] for k in data.files} + + def discover_test_cases( task_type: str | None = None, level: str | None = None, diff --git a/src/mobius/_testing/parity.py b/src/mobius/_testing/parity.py index e21ce3692..61f8696b1 100644 --- a/src/mobius/_testing/parity.py +++ b/src/mobius/_testing/parity.py @@ -6,6 +6,9 @@ Provides ParityReport dataclass and level-appropriate comparison functions: - compare_synthetic(): atol/rtol-gated (for L3 synthetic parity) - compare_golden(): argmax-gated (for L4 golden comparison) +- compare_diarization_golden(): atol/rtol-gated over continuous per-frame + per-speaker probabilities (for diarization L4/L5 golden comparison, where + argmax-over-vocab is not a meaningful gate) """ from __future__ import annotations @@ -347,3 +350,118 @@ def compare_golden( level="L4", message=message, ) + + +# Default tolerances for diarization probability comparison, keyed by dtype. +_DIARIZATION_TOLERANCES: dict[str, tuple[float, float]] = { + "float32": (1e-3, 1e-2), + "float16": (1e-2, 5e-2), + "bfloat16": (5e-2, 1e-1), +} + + +def compare_diarization_golden( + onnx_probs: np.ndarray, + golden_probs: np.ndarray, + atol: float | None = None, + rtol: float | None = None, + dtype: str = "float32", + level: str = "L4", +) -> ParityReport: + """L4/L5 diarization golden comparison. Gate: elementwise allclose. + + Diarization models (``diarization`` / ``diarization-streaming`` task + types) emit continuous per-frame per-speaker sigmoid probabilities in + ``[0, 1]`` rather than a vocab logit vector, so ``compare_golden()``'s + argmax/top-K gate is not meaningful here. The gate is atol/rtol + allclose on the full probability tensor (mirroring ``compare_synthetic``'s + gate, but used at L4/L5 against a real-weights reference rather than a + synthetic one). Per-frame dominant-speaker argmax agreement and the + active-speaker-set Jaccard at the standard 0.5 decision threshold are + reported as diagnostics (repurposing ``argmax_match``/``top10_jaccard`` + for the diarization-appropriate quantities of the same *kind*). Unlike + ``compare_golden()``, there is no AMBIGUOUS downgrade for an allclose + failure: dominant-speaker argmax agreement is not a valid substitute for + full elementwise tolerance on a multi-label sigmoid output, since a + secondary (non-dominant) speaker's probability can cross the 0.5 + activation threshold -- a real diarization error -- while the dominant + argmax stays unchanged. Every allclose failure is a hard FAIL. + """ + assert onnx_probs.shape == golden_probs.shape, ( + f"Shape mismatch: ONNX {onnx_probs.shape} vs golden {golden_probs.shape}" + ) + default_atol, default_rtol = _DIARIZATION_TOLERANCES.get(dtype, (1e-3, 1e-2)) + atol = default_atol if atol is None else atol + rtol = default_rtol if rtol is None else rtol + + a64 = onnx_probs.astype(np.float64) + b64 = golden_probs.astype(np.float64) + + abs_diff = np.abs(a64 - b64) + max_abs_diff = float(abs_diff.max()) + mean_abs_diff = float(abs_diff.mean()) + atol_pass = bool(np.allclose(a64, b64, atol=atol, rtol=0)) + rtol_pass = bool(np.allclose(a64, b64, atol=0, rtol=rtol)) + allclose_pass = bool(np.allclose(a64, b64, atol=atol, rtol=rtol)) + + # Per-frame dominant-speaker argmax agreement (diagnostic). + onnx_dominant = np.argmax(a64, axis=-1) + golden_dominant = np.argmax(b64, axis=-1) + dominant_match_ratio = float(np.mean(onnx_dominant == golden_dominant)) + dominant_match = dominant_match_ratio >= 1.0 - 1e-9 + + # Active-speaker-set Jaccard at the 0.5 decision threshold, averaged per + # frame (the diarization analogue of compare_golden's top10_jaccard). + onnx_active = a64 > 0.5 + golden_active = b64 > 0.5 + union = np.logical_or(onnx_active, golden_active).sum(axis=-1) + intersection = np.logical_and(onnx_active, golden_active).sum(axis=-1) + frame_jaccard = np.where(union > 0, intersection / np.maximum(union, 1), 1.0) + active_speaker_jaccard = float(np.mean(frame_jaccard)) + + norm_onnx = np.linalg.norm(a64.reshape(-1)) + norm_golden = np.linalg.norm(b64.reshape(-1)) + cosine_similarity = ( + float(np.dot(a64.reshape(-1), b64.reshape(-1)) / (norm_onnx * norm_golden)) + if norm_onnx > 0 and norm_golden > 0 + else 0.0 + ) + + if allclose_pass: + result = ParityResult.PASS + message = ( + f"{level} PASS: diarization probs match within atol={atol}, " + f"rtol={rtol} (max_abs_diff={max_abs_diff:.4g})" + ) + else: + # No AMBIGUOUS downgrade here: dominant-speaker argmax agreement is + # not a valid substitute for full elementwise tolerance on a + # multi-label sigmoid output -- a secondary speaker's probability + # can cross the 0.5 activation threshold (a real diarization error) + # while the single dominant argmax happens to stay unchanged. Report + # dominant-speaker match and active-speaker-set Jaccard as + # diagnostics only; any allclose failure is a hard FAIL. + result = ParityResult.FAIL + message = ( + f"{level} FAIL: diarization probs diverge " + f"(max_abs_diff={max_abs_diff:.4g} > atol={atol}), " + f"dominant-speaker match={dominant_match_ratio:.2%}, " + f"active-speaker Jaccard={active_speaker_jaccard:.2%}" + ) + + return ParityReport( + argmax_match=dominant_match, + atol_pass=atol_pass, + rtol_pass=rtol_pass, + max_abs_diff=max_abs_diff, + mean_abs_diff=mean_abs_diff, + top10_jaccard=active_speaker_jaccard, + cosine_similarity=cosine_similarity, + top1_logit=float("nan"), + top2_logit=float("nan"), + near_tie=False, + near_tie_margin=0.0, + result=result, + level=level, + message=message, + ) diff --git a/src/mobius/_testing/parity_test.py b/src/mobius/_testing/parity_test.py index 60553a35d..e9d1a486e 100644 --- a/src/mobius/_testing/parity_test.py +++ b/src/mobius/_testing/parity_test.py @@ -11,6 +11,7 @@ from mobius._testing.parity import ( ParityReport, ParityResult, + compare_diarization_golden, compare_golden, compare_synthetic, ) @@ -221,3 +222,48 @@ def test_identical_top10_low_token_promoted_still_fails(self): golden_top10_ids=golden_top10, ) assert report.result == ParityResult.FAIL + + +class TestCompareDiarizationGolden: + """Tests for L4/L5 diarization golden comparison.""" + + def test_identical_probs_pass(self): + probs = np.array([[[0.9, 0.1], [0.2, 0.8]]], dtype=np.float32) + report = compare_diarization_golden(probs, probs.copy()) + assert report.result == ParityResult.PASS + + def test_small_diff_within_tolerance_passes(self): + golden = np.array([[[0.9, 0.1], [0.2, 0.8]]], dtype=np.float32) + onnx = golden + 1e-5 + report = compare_diarization_golden(onnx, golden) + assert report.result == ParityResult.PASS + + def test_secondary_speaker_threshold_crossing_fails(self): + """Regression test for the review-flagged AMBIGUOUS-masking bug. + + A single frame's dominant speaker (index 0) is unchanged in both + ONNX and golden, but the *secondary* speaker's probability crosses + the 0.5 decision threshold (0.6 -> golden considers it active, + 0.1 -> ONNX does not). This is a real diarization error -- the + active-speaker set differs -- and must FAIL, not be downgraded to + AMBIGUOUS merely because the dominant-speaker argmax still agrees. + """ + golden = np.array([[[0.9, 0.6]]], dtype=np.float32) + onnx = np.array([[[0.9, 0.1]]], dtype=np.float32) + report = compare_diarization_golden(onnx, golden, atol=1e-3, rtol=1e-2) + assert report.result == ParityResult.FAIL + # Diagnostics are still reported even though the result is FAIL. + assert report.argmax_match is True + + def test_large_diff_with_dominant_mismatch_fails(self): + golden = np.array([[[0.9, 0.1], [0.2, 0.8]]], dtype=np.float32) + onnx = np.array([[[0.1, 0.9], [0.8, 0.2]]], dtype=np.float32) + report = compare_diarization_golden(onnx, golden) + assert report.result == ParityResult.FAIL + assert report.argmax_match is False + + def test_shape_mismatch_raises(self): + golden = np.zeros((1, 2, 3), dtype=np.float32) + onnx = np.zeros((1, 2, 4), dtype=np.float32) + with pytest.raises(AssertionError): + compare_diarization_golden(onnx, golden) diff --git a/src/mobius/models/__init__.py b/src/mobius/models/__init__.py index 11b9ef922..638ae99db 100644 --- a/src/mobius/models/__init__.py +++ b/src/mobius/models/__init__.py @@ -181,6 +181,7 @@ "SEMambaSpeechEnhancementModel", "SortformerConfig", "SortformerDiarizationModel", + "Nemotron3DiarizationModel", "Qwen3TTSCodePredictorModel", "Qwen3TTSCodecDecoderModel", "Qwen3TTSCodecEncoderModel", @@ -371,6 +372,7 @@ from mobius.models.nanochat import NanoChatCausalLMModel from mobius.models.nemo_rnnt import EncDecRNNTModel from mobius.models.nemotron import NemotronCausalLMModel +from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel from mobius.models.nemotron_h import NemotronHCausalLMModel from mobius.models.nemotron_parse import NemotronParseForConditionalGeneration from mobius.models.olmo import OLMo2CausalLMModel, OLMoCausalLMModel diff --git a/src/mobius/models/nemotron3_diarization.py b/src/mobius/models/nemotron3_diarization.py new file mode 100644 index 000000000..f2c39f408 --- /dev/null +++ b/src/mobius/models/nemotron3_diarization.py @@ -0,0 +1,1265 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Nemotron 3 Diarization: streaming Sortformer speaker-diarization model (HF). + +Replicates HuggingFace's ``Nemotron3DiarizationForAudioFrameClassification`` +(e.g. ``nvidia/Nemotron-3-Diarization``). The exported ONNX graph(s) consume +mel-spectrogram features and produce per-frame speaker-activity probabilities +for up to ``config.num_speakers`` speakers, ordered by their first arrival in +the audio. + +Pipeline (matching ``Nemotron3DiarizationAudioModel`` + ``Nemotron3DiarizationModel`` ++ ``Nemotron3DiarizationForAudioFrameClassification`` in HuggingFace's +``modeling_nemotron3_diarization.py``): + + input_features [B, feat, T] + -> transpose [B, T, feat] + -> feature stacking (subsampling) [B, T/8, feat*8] -> Linear -> hidden + -> bidirectional RoPE Transformer encoder (pre-LN) + -> proj (Linear) [B, T/8, head_hidden] + -> sub-pixel upsampler (Conv1d) [B, T, head_hidden] + -> classification head (relu-dense-relu-out_proj) -> sigmoid + +Two forwards/tasks are exported: + +* **Offline** (``diarization`` task, :class:`~mobius.tasks.DiarizationTask`): + the whole input is embedded once, then an ONNX ``Loop`` iterates over + fixed-size ``config.chunk_length`` embed chunks with + ``config.chunk_right_context`` look-ahead, reusing the same Arrival-Order + Speaker Cache (AOSC) + FIFO bookkeeping as streaming but with + offline-specific cache sizes. Matches HuggingFace's offline forward for + recordings of any length, not just single-chunk ones. +* **Streaming** (``diarization-streaming`` task, + :class:`~mobius.tasks.DiarizationStreamingTask`): a per-chunk, stateful + forward. Every call consumes one chunk of audio (plus a few look-ahead + frames) and the previous call's Arrival-Order Speaker Cache (AOSC) + FIFO + queue state, and returns this chunk's speaker probabilities plus updated + cache state. The AOSC/FIFO buffers are fixed-size graph inputs/outputs + (like a static KV cache) with explicit scalar occupancy counters, so the + same ONNX graph can be called repeatedly to process an arbitrarily long + recording — reproducing HuggingFace's ``Nemotron3DiarizationSpeakerCache`` + bookkeeping exactly, top-k score-based compression included. +""" + +from __future__ import annotations + +import math + +import onnx_ir as ir +from onnxscript import OpBuilder, nn + +from mobius._configs import Nemotron3DiarizationConfig +from mobius.components import ( + Conv1d, + LayerNorm, + Linear, + apply_rotary_pos_emb, + get_activation, + initialize_rope, +) + +# LayerNorm epsilon used throughout the HuggingFace reference (``nn.LayerNorm`` +# defaults), not exposed as a config field. +_LAYER_NORM_EPS = 1e-5 + + +class _FeatureStacking(nn.Module): + """Stacks ``subsampling_factor`` consecutive spectrogram frames and projects them. + + Matches ``Nemotron3DiarizationFeatureStacking``: zero-pads the time axis + to a multiple of ``subsampling_factor`` (the padding is a runtime-computed + length since it depends on the dynamic sequence length), then reshapes + consecutive frame groups into the feature axis before projecting to + ``hidden_size``. + """ + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self._factor = config.subsampling_factor + self._stacked_size = config.feat_in * config.subsampling_factor + self.projection = Linear(self._stacked_size, config.hidden_size, bias=False) + + def forward(self, op: OpBuilder, input_features: ir.Value) -> ir.Value: + # input_features: [B, T, feat_in]. + sequence_length = op.Squeeze(op.Shape(input_features, start=1, end=2), [0]) + factor = op.Constant(value_int=self._factor) + remainder = op.Mod(sequence_length, factor, fmod=0) + pad_length = op.Mod(op.Sub(factor, remainder), factor, fmod=0) + # Pad axis 1 (time) at the end by ``pad_length`` frames. + pads = op.Concat( + op.Constant(value_ints=[0, 0, 0, 0]), + op.Reshape(pad_length, [1]), + op.Constant(value_ints=[0]), + axis=0, + ) + stacked = op.Pad(input_features, pads) + # [B, T', feat_in] -> [B, T'/factor, feat_in*factor] (row-major + # reshape, matching PyTorch's contiguous ``.reshape``). + stacked = op.Reshape(stacked, [0, -1, self._stacked_size]) + return self.projection(op, stacked) + + +class _AudioAttention(nn.Module): + """Bidirectional GQA-capable audio attention with full RoPE. + + Matches ``Nemotron3DiarizationAttention``: unbiased Q/K/V projections + and a biased output projection, no causal masking (single unpadded + utterance, so no attention bias is needed either). + """ + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self._num_attention_heads = config.num_attention_heads + self._num_key_value_heads = config.num_key_value_heads + self._head_dim = config.head_dim + prf = config.partial_rotary_factor if config.partial_rotary_factor is not None else 1.0 + self._rotary_dim = 0 if math.isclose(prf, 1.0) else int(self._head_dim * prf) + self._scale = config.head_dim**-0.5 + + q_size = config.num_attention_heads * config.head_dim + kv_size = config.num_key_value_heads * config.head_dim + self.q_proj = Linear(config.hidden_size, q_size, bias=False) + self.k_proj = Linear(config.hidden_size, kv_size, bias=False) + self.v_proj = Linear(config.hidden_size, kv_size, bias=False) + self.o_proj = Linear(q_size, config.hidden_size, bias=True) + + def forward( + self, + op: OpBuilder, + hidden_states: ir.Value, + position_embeddings: tuple[ir.Value, ir.Value], + ) -> ir.Value: + query_states = self.q_proj(op, hidden_states) + key_states = self.k_proj(op, hidden_states) + value_states = self.v_proj(op, hidden_states) + + query_states = apply_rotary_pos_emb( + op, + query_states, + position_embeddings, + num_heads=self._num_attention_heads, + rotary_embedding_dim=self._rotary_dim, + ) + key_states = apply_rotary_pos_emb( + op, + key_states, + position_embeddings, + num_heads=self._num_key_value_heads, + rotary_embedding_dim=self._rotary_dim, + ) + + attention_output = op.Attention( + query_states, + key_states, + value_states, + None, + None, + None, + q_num_heads=self._num_attention_heads, + kv_num_heads=self._num_key_value_heads, + scale=self._scale, + is_causal=0, + ) + return self.o_proj(op, attention_output) + + +class _AudioMLP(nn.Module): + """Feed-forward network: ``fc1 -> activation -> fc2`` (both biased).""" + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.fc1 = Linear(config.hidden_size, config.intermediate_size, bias=True) + self.fc2 = Linear(config.intermediate_size, config.hidden_size, bias=True) + self._activation = get_activation(config.hidden_act) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + return self.fc2(op, self._activation(op, self.fc1(op, hidden_states))) + + +class _AudioLayer(nn.Module): + """Pre-norm bidirectional transformer layer (``Nemotron3DiarizationAudioLayer``).""" + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.self_attn = _AudioAttention(config) + self.layer_norm1 = LayerNorm(config.hidden_size, eps=_LAYER_NORM_EPS) + self.mlp = _AudioMLP(config) + self.layer_norm2 = LayerNorm(config.hidden_size, eps=_LAYER_NORM_EPS) + + def forward( + self, + op: OpBuilder, + hidden_states: ir.Value, + position_embeddings: tuple[ir.Value, ir.Value], + ) -> ir.Value: + residual = hidden_states + hidden_states = self.layer_norm1(op, hidden_states) + hidden_states = self.self_attn(op, hidden_states, position_embeddings) + hidden_states = op.Add(residual, hidden_states) + + residual = hidden_states + hidden_states = self.layer_norm2(op, hidden_states) + hidden_states = self.mlp(op, hidden_states) + return op.Add(residual, hidden_states) + + +class _AudioTower(nn.Module): + """RoPE transformer audio encoder (``Nemotron3DiarizationAudioModel``).""" + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.embedder = _FeatureStacking(config) + self.input_layer_norm = LayerNorm(config.hidden_size, eps=_LAYER_NORM_EPS) + self.layers = nn.ModuleList( + [_AudioLayer(config) for _ in range(config.num_hidden_layers)] + ) + self.layer_norm = LayerNorm(config.hidden_size, eps=_LAYER_NORM_EPS) + self.rotary_emb = initialize_rope(config) + + def embed(self, op: OpBuilder, input_features: ir.Value) -> ir.Value: + """Feature-stacking projection only (raw features -> encoder-frame embeddings). + + Split out from :meth:`encode` so the streaming forward can embed just + the new chunk's raw features, then concatenate them with cached + (already-embedded) frames before running the layer stack. + """ + return self.embedder(op, input_features) + + def encode(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + """Runs the RoPE transformer layer stack over already-embedded frames.""" + sequence_length = op.Squeeze(op.Shape(hidden_states, start=1, end=2), [0]) + position_ids = op.Unsqueeze( + op.Range( + op.Constant(value_int=0), + sequence_length, + op.Constant(value_int=1), + ), + [0], + ) + if self.rotary_emb is None: + raise ValueError("Nemotron3Diarization audio encoder requires rotary embeddings") + position_embeddings = self.rotary_emb(op, position_ids) + + hidden_states = self.input_layer_norm(op, hidden_states) + for layer in self.layers: + hidden_states = layer(op, hidden_states, position_embeddings) + return self.layer_norm(op, hidden_states) + + def forward(self, op: OpBuilder, input_features: ir.Value) -> ir.Value: + hidden_states = self.embed(op, input_features) + return self.encode(op, hidden_states) + + +class _SubpixelUpsampler(nn.Module): + """Upsamples the encoder frame rate back to the spectrogram frame rate. + + Matches ``Nemotron3DiarizationSubpixelUpsampler``: a Conv1d expands the + channel axis by ``subsampling_factor``, which is then reshaped into the + time axis (sub-pixel / depth-to-space upsampling). + """ + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self._factor = config.subsampling_factor + self._hidden_size = config.head_hidden_size + self.conv = Conv1d( + config.head_hidden_size, + config.head_hidden_size * config.subsampling_factor, + kernel_size=3, + padding=1, + ) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + # [B, T, C] -> [B, C, T] -> conv -> [B, C*factor, T] -> [B, T, C*factor] + hidden_states = op.Transpose(hidden_states, perm=[0, 2, 1]) + hidden_states = self.conv(op, hidden_states) + hidden_states = op.Transpose(hidden_states, perm=[0, 2, 1]) + # Row-major reshape spreads each time step's C*factor channels across + # ``factor`` consecutive upsampled time steps of width C. + return op.Reshape(hidden_states, [0, -1, self._hidden_size]) + + +class _DiarizationBackbone(nn.Module): + """Audio tower + projection + upsampler (``Nemotron3DiarizationModel``).""" + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.audio_tower = _AudioTower(config) + self.proj = Linear(config.hidden_size, config.head_hidden_size) + self.upsampler = _SubpixelUpsampler(config) + + def project_upsample(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + hidden_states = self.proj(op, hidden_states) + return self.upsampler(op, hidden_states) + + def forward(self, op: OpBuilder, input_features: ir.Value) -> ir.Value: + hidden_states = self.audio_tower(op, input_features) + return self.project_upsample(op, hidden_states) + + +class _ClassificationHead(nn.Module): + """Speaker activity head: ``relu -> dense -> relu -> out_proj``. + + Matches ``Nemotron3DiarizationClassificationHead`` exactly, including + the (unusual) placement of ReLU *before* both linear layers. + """ + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.dense = Linear(config.head_hidden_size, config.head_hidden_size) + self.out_proj = Linear(config.head_hidden_size, config.num_speakers) + + def forward(self, op: OpBuilder, hidden_states: ir.Value) -> ir.Value: + hidden_states = self.dense(op, op.Relu(hidden_states)) + return self.out_proj(op, op.Relu(hidden_states)) + + +def _prerealize_parameters(builder, module: nn.Module, prefix: str) -> None: + """Registers every parameter under ``module`` as a root-graph initializer. + + Uses ``named_parameters()``'s pre-computed dotted names directly. + + ``Parameter._realize`` normally qualifies a name via the *root* graph + builder's current module-scope stack (``push_module``/``pop_module``) + -- but that scope stack belongs to whichever builder happens to call + ``_realize`` first, which is wrong when a parameter is first referenced + from *inside* a ``Loop``/``If`` subgraph's own separate sub-builder (its + scope-stack pushes never touch the *root* builder's stack that + ``_realize`` actually reads). Bypassing that mechanism here -- while + still at plain root scope, before any subgraph is built -- realizes + every descendant parameter up front with the exact same fully-qualified + name ``Module.__call__``'s automatic realization would have produced. + ``_realize`` is idempotent, so every later (redirect) call becomes a + no-op regardless of which builder performs it. + """ + root = builder.root + for name, param in module.named_parameters(prefix=prefix): + if param._realized: # pylint: disable=protected-access + continue + param.name = name + root.graph.initializers[name] = param + param._realized = True # pylint: disable=protected-access + + +def _scalar(op: OpBuilder, value: ir.Value) -> ir.Value: + """Reshapes a scalar (rank-0) value to rank-1 shape ``[1]`` for Slice bounds.""" + return op.Reshape(value, op.Constant(value_ints=[1])) + + +def _pad_time_axis(op: OpBuilder, x: ir.Value, target_length: ir.Value) -> ir.Value: + """Zero-pads a ``[B, F, C]`` tensor at the end of axis 1 up to ``target_length``. + + ``target_length`` is a scalar (rank-0) ``int64`` value; ``F`` (the current + length) may be dynamic and is assumed ``<= target_length``. + """ + current_length = op.Shape(x, start=1, end=2) + pad_amount = op.Sub(_scalar(op, target_length), current_length) + zero = op.Constant(value_ints=[0]) + pads = op.Concat(zero, zero, zero, zero, pad_amount, zero, axis=0) + return op.Pad(x, pads, mode="constant") + + +def _avg_pool_probs(op: OpBuilder, logits: ir.Value, factor: int) -> ir.Value: + """Sigmoid + average-pool speaker logits down to the encoder frame rate. + + Matches ``Nemotron3DiarizationSpeakerCache._pool_probs`` (without the + optional padding mask, which streaming export does not support — see the + module docstring's streaming limitations). + """ + probs = op.Sigmoid(logits) + probs = op.Transpose(probs, perm=[0, 2, 1]) + probs = op.AveragePool(probs, kernel_shape=[factor], strides=[factor]) + return op.Transpose(probs, perm=[0, 2, 1]) + + +def _num_popped_frames( + op: OpBuilder, num_fifo_total: ir.Value, fifo_capacity: int, update_period: int +) -> ir.Value: + """Matches ``Nemotron3DiarizationSpeakerCache._num_popped_frames``.""" + over_capacity = op.Sub(num_fifo_total, op.Constant(value_int=fifo_capacity)) + popped = op.Max(op.Constant(value_int=update_period), over_capacity) + popped = op.Min(popped, num_fifo_total) + zero = op.Constant(value_int=0) + has_overflow = op.Greater(num_fifo_total, op.Constant(value_int=fifo_capacity)) + return op.Where(has_overflow, popped, zero) + + +def _get_frame_scores( + op: OpBuilder, probs: ir.Value, threshold: float, min_positive_scores: int +) -> ir.Value: + """Matches ``Nemotron3DiarizationSpeakerCache._get_frame_scores``. + + ``probs``: ``[B, F, S]``. Frames that are all-zero (this module's padding + convention for "no such frame") are never speech (``probs <= 0.5``), so + they always resolve to ``-inf`` and are naturally excluded from top-k + selection without extra masking. + + All float literals below are cast to ``probs``'s dtype: this scoring path + runs entirely in ``config.dtype`` (fp16/bf16 exports included), and ONNX + elementwise ops require both operands to share a dtype — a bare FLOAT32 + ``Constant`` combined with an fp16/bf16 tensor is an invalid graph. + """ + + def _lit(value: float) -> ir.Value: + return op.CastLike(op.Constant(value_float=value), probs) + + threshold_t = _lit(threshold) + log_probs = op.Log(op.Clip(probs, threshold_t)) + complements = op.Sub(_lit(1.0), probs) + log_complements = op.Log(op.Clip(complements, threshold_t)) + sum_log_complements = op.ReduceSum(log_complements, [-1], keepdims=1) + log_half = _lit(math.log(0.5)) + scores = op.Sub(op.Add(op.Sub(log_probs, log_complements), sum_log_complements), log_half) + + neg_inf = _lit(float("-inf")) + is_speech = op.Greater(probs, _lit(0.5)) + scores = op.Where(is_speech, scores, neg_inf) + + is_positive = op.Greater(scores, _lit(0.0)) + positive_count = op.ReduceSum(op.Cast(is_positive, to=ir.DataType.INT64), [1], keepdims=1) + has_enough_positive = op.GreaterOrEqual( + positive_count, op.Constant(value_int=min_positive_scores) + ) + extra_masked = op.And(op.And(op.Not(is_positive), is_speech), has_enough_positive) + return op.Where(extra_masked, neg_inf, scores) + + +def _boost_scores(op: OpBuilder, scores: ir.Value, num_boosted: int, boost: float) -> ir.Value: + """Matches ``Nemotron3DiarizationSpeakerCache._boost_scores`` (no-op if ``num_boosted <= 0``).""" + if num_boosted <= 0: + return scores + _, indices = op.TopK( + scores, + op.Constant(value_ints=[num_boosted]), + axis=1, + largest=1, + sorted=0, + _outputs=2, + ) + updates = op.Expand(op.Constant(value_float=boost), op.Shape(indices)) + updates = op.CastLike(updates, scores) + return op.ScatterElements(scores, indices, updates, axis=1, reduction="add") + + +def _sort_ascending(op: OpBuilder, values: ir.Value, length: int) -> ir.Value: + """Sorts a ``[B, length]`` int64 tensor ascending along axis 1 (no ``Sort`` op in ONNX). + + Uses the "negate + TopK(largest, sorted, k=length)" trick: requesting the + full length back in sorted-descending order of the negation is exactly + ascending order of the original values. + """ + negated = op.Cast(op.Neg(values), to=ir.DataType.FLOAT) + sorted_negated, _ = op.TopK( + negated, + op.Constant(value_ints=[length]), + axis=1, + largest=1, + sorted=1, + _outputs=2, + ) + return op.Cast(op.Neg(sorted_negated), to=ir.DataType.INT64) + + +def _compress_speaker_cache( + op: OpBuilder, + embeds: ir.Value, + probs: ir.Value, + silence_embeds: ir.Value, + config: Nemotron3DiarizationConfig, +) -> tuple[ir.Value, ir.Value]: + """Matches ``Nemotron3DiarizationSpeakerCache._compress``. + + ``embeds``: ``[B, F, H]``, ``probs``: ``[B, F, S]`` with ``F`` (dynamic) + strictly greater than ``speaker_cache_length`` (the caller only invokes + this when compression is actually needed). Returns exactly + ``speaker_cache_length`` frames, keyed by score and grouped by speaker + (highest-scoring frames per speaker, in original temporal order). + """ + num_speakers = config.num_speakers + cache_length = config.streaming_speaker_cache_length + num_silence_frames = config.streaming_silence_frames_per_speaker + threshold = config.streaming_prediction_score_threshold + # Per-speaker budget the score policy spends on boosting/positivity, + # excluding the reserved silence slots (mirrors the Python constants + # HF precomputes once in ``Nemotron3DiarizationSpeakerCache.__init__``). + budget = cache_length // num_speakers - num_silence_frames + min_positive_scores = math.floor(budget * config.streaming_min_positive_scores_rate) + num_strong_boosted = math.floor(budget * config.streaming_strong_boost_rate) + num_weak_boosted = math.floor(budget * config.streaming_weak_boost_rate) + + scores = _get_frame_scores(op, probs, threshold, min_positive_scores) + + num_frames = op.Squeeze(op.Shape(embeds, start=1, end=2), [0]) + frame_positions = op.Range(op.Constant(value_int=0), num_frames, op.Constant(value_int=1)) + is_tail = op.GreaterOrEqual(frame_positions, op.Constant(value_int=cache_length)) + tail_boost = op.Where( + is_tail, + op.Constant(value_float=config.streaming_latest_frames_score_boost), + op.Constant(value_float=0.0), + ) + # ``tail_boost`` is built from bare FLOAT32 literals (matching each + # other); cast once here to ``scores``'s dtype before combining, rather + # than casting each literal individually. + tail_boost = op.CastLike(tail_boost, scores) + scores = op.Add(scores, op.Unsqueeze(tail_boost, [0, 2])) + + scores = _boost_scores(op, scores, num_strong_boosted, -2.0 * math.log(0.5)) + scores = _boost_scores(op, scores, num_weak_boosted, -math.log(0.5)) + + # Append the shared silence embedding/prob row, and reserve + # ``num_silence_frames`` score slots per speaker with +inf (always kept). + batch = op.Shape(embeds, start=0, end=1) + hidden_size = op.Shape(embeds, start=2, end=3) + silence_row = op.Expand( + op.Reshape(silence_embeds, op.Constant(value_ints=[1, 1, -1])), + op.Concat(batch, op.Constant(value_ints=[1]), hidden_size, axis=0), + ) + embeds = op.Concat(embeds, silence_row, axis=1) + probs_pad = op.Concat( + op.Constant(value_ints=[0, 0, 0]), op.Constant(value_ints=[0, 1, 0]), axis=0 + ) + probs = op.Pad(probs, probs_pad, mode="constant") + scores_pad = op.Concat( + op.Constant(value_ints=[0, 0, 0]), + op.Constant(value_ints=[0, num_silence_frames, 0]), + axis=0, + ) + scores = op.Pad( + scores, + scores_pad, + op.CastLike(op.Constant(value_float=float("inf")), scores), + mode="constant", + ) + + num_scored_frames = op.Add(num_frames, op.Constant(value_int=num_silence_frames)) + sentinel = op.Mul(num_scored_frames, op.Constant(value_int=num_speakers)) + flat_scores = op.Reshape( + op.Transpose(scores, perm=[0, 2, 1]), + op.Concat(batch, op.Constant(value_ints=[-1]), axis=0), + ) + topk_scores, topk_indices = op.TopK( + flat_scores, + op.Constant(value_ints=[cache_length]), + axis=1, + largest=1, + sorted=0, + _outputs=2, + ) + is_masked = op.Equal( + topk_scores, op.CastLike(op.Constant(value_float=float("-inf")), topk_scores) + ) + topk_indices = op.Where(is_masked, op.Unsqueeze(sentinel, [0]), topk_indices) + topk_indices = _sort_ascending(op, topk_indices, cache_length) + + frame_indices = op.Mod(topk_indices, op.Unsqueeze(num_scored_frames, [0]), fmod=0) + frame_indices = op.Min(frame_indices, op.Unsqueeze(num_frames, [0])) + is_sentinel = op.Equal(topk_indices, op.Unsqueeze(sentinel, [0])) + frame_indices = op.Where(is_sentinel, op.Unsqueeze(num_frames, [0]), frame_indices) + + gather_indices = op.Unsqueeze(frame_indices, [-1]) + new_embeds = op.GatherND(embeds, gather_indices, batch_dims=1) + new_probs = op.GatherND(probs, gather_indices, batch_dims=1) + return new_embeds, new_probs + + +class Nemotron3DiarizationModel(nn.Module): + """Streaming Sortformer speaker-diarization model (HuggingFace port). + + Replicates HuggingFace's ``Nemotron3DiarizationForAudioFrameClassification`` + (e.g. ``nvidia/Nemotron-3-Diarization``): a bidirectional, partial-RoPE + Transformer audio encoder followed by a sub-pixel upsampler and a speaker + sigmoid head. Consumes mel-spectrogram features and returns per-frame + speaker-activity probabilities in ``[0, 1]``. + + Two forwards are exported (see ``tasks/_diarization.py`` and + ``tasks/_diarization_streaming.py``): + + * :meth:`forward` — the **offline** whole-recording pass (``diarization`` + task), chunked via an ONNX ``Loop`` exactly like HuggingFace's offline + forward (any recording length, not just single-chunk ones). + * :meth:`forward_streaming` — the **streaming**, per-chunk pass + (``diarization-streaming`` task): consumes one chunk of audio plus a + few look-ahead frames, together with the previous step's Arrival-Order + Speaker Cache (AOSC) + FIFO queue state (fixed-size buffers and scalar + occupancy counters, all graph inputs/outputs), and returns this chunk's + speaker probabilities plus the updated cache state. Repeated calls + reproduce HuggingFace's streaming ``Nemotron3DiarizationSpeakerCache`` + bookkeeping exactly, including its top-k score-based compression. + + Limitation: unlike HuggingFace, the streaming export does not accept a + padding ``attention_mask`` — every step's input window is assumed fully + valid (no silence padding within a chunk). This holds for all but + possibly the very last chunk of a recording, matching the precision + needed for real-time streaming use. + """ + + default_task: str = "diarization" + category: str = "Speech-to-Text" + config_class = Nemotron3DiarizationConfig + + def __init__(self, config: Nemotron3DiarizationConfig): + super().__init__() + self.config = config + self.model = _DiarizationBackbone(config) + self.classifier = _ClassificationHead(config) + # Learned embedding filling reserved silence slots when the AOSC is + # compressed. Shared by both the offline and streaming forwards: + # the offline ``forward``'s chunked ``Loop`` calls the same + # ``_run_chunk_and_update_cache`` compression branch as + # ``forward_streaming``, so this is realized and consumed whenever + # a multi-chunk offline recording triggers AOSC compression too. + self.silence_embeds = nn.Parameter([config.hidden_size]) + + def forward(self, op: OpBuilder, input_features: ir.Value) -> ir.Value: + """Offline (whole-recording) forward, chunked exactly like HuggingFace. + + Matches ``Nemotron3DiarizationForAudioFrameClassification.forward``'s + offline path: the whole input is embedded once, then an ``ONNX Loop`` + iterates over fixed-size ``config.chunk_length`` embed chunks (with + ``config.chunk_right_context`` look-ahead), reusing the same + Arrival-Order Speaker Cache (AOSC) + FIFO bookkeeping as the streaming + forward but with offline-specific cache sizes + (``config.offline_fifo_length`` / ``config.offline_speaker_cache_update_period``). + This makes the offline graph exact for recordings of any length, not + just ones that fit within a single chunk. + """ + config = self.config + factor = config.subsampling_factor + chunk_length = config.chunk_length + chunk_right_context = config.chunk_right_context + cache_length = config.streaming_speaker_cache_length + fifo_capacity = config.offline_fifo_length + update_period = config.offline_speaker_cache_update_period + hidden_size = config.hidden_size + num_speakers = config.num_speakers + + # ``Parameter._realize`` qualifies a parameter's name using the + # *root* graph builder's current module-scope stack, not the scope + # stack of whatever (sub-)builder happens to invoke it -- so calling + # ``encode``/``project_upsample``/``classifier`` for the first time + # from *inside* the ``Loop`` body below (a separate sub-builder, via + # ``op.builder.subgraph(...)``) would both lose hierarchical name + # qualification *and* (since the sub-builder's own node-name counter + # independently reaches the same count as an equivalent outer-scope + # trace) risk colliding with an unrelated node's auto-generated name. + # Pre-realize every parameter with its final, fully-qualified dotted + # name (from ``named_parameters()``) directly, while still at root + # scope, bypassing ``_realize``'s scope-stack-based qualification + # entirely -- ``_realize`` is idempotent, so this makes every later + # call (from ``embed`` here and from ``encode``/``project_upsample``/ + # ``classifier`` inside the loop body) a no-op. + self.silence_embeds._realize(op.builder) # type: ignore[attr-defined] + _prerealize_parameters(op.builder, self.model, "model") + _prerealize_parameters(op.builder, self.classifier, "classifier") + + # input_features: [B, feat_in, T] -> [B, T, feat_in], matching the + # shared ``DiarizationTask`` input contract (also used by sortformer). + input_features = op.Transpose(input_features, perm=[0, 2, 1]) + # Original (pre-padding) frame count, to trim the upsampled output. + raw_num_frames = op.Shape(input_features, start=1, end=2) + + # One embedding pass over the *whole* input (matches HuggingFace: + # ``inputs_embeds = embedder(input_features)`` computed once, then + # chunks are sliced directly from it -- not re-embedded per chunk). + all_chunk_embeds = self._call_backbone_scoped(op, "embed", input_features) + num_embeds = op.Squeeze(op.Shape(all_chunk_embeds, start=1, end=2), [0]) + batch = op.Shape(all_chunk_embeds, start=0, end=1) + + chunk_length_c = op.Constant(value_int=chunk_length) + num_iterations = op.Div( + op.Sub(op.Add(num_embeds, chunk_length_c), op.Constant(value_int=1)), + chunk_length_c, + ) + # Fixed total accumulator length (all frames the loop will ever + # produce), computed once outside the loop -- lets + # ``accumulated_logits`` be a *fixed-shape* loop-carried buffer + # (each iteration adds its own zero-padded, non-overlapping region + # via ``Add`` rather than growing the tensor via ``Concat``). This + # sidesteps a real optimizer pitfall: a naive shape-inference pass + # can mistake a ``Concat`` whose first operand starts out empty + # (shape ``[B, 0, S]``) for a compile-time identity and fold it away + # -- which would silently keep only the *last* iteration's chunk. + total_raw_length = op.Mul(num_embeds, op.Constant(value_int=factor)) + + def _loop_body( + body_op: OpBuilder, + iter_num: ir.Value, + cond_in: ir.Value, + start_idx: ir.Value, + accumulated_logits: ir.Value, + cache_embeds: ir.Value, + cache_probs: ir.Value, + fifo: ir.Value, + num_cache_frames: ir.Value, + num_fifo_frames: ir.Value, + is_compressed: ir.Value, + ): + end_idx = body_op.Min( + body_op.Add(start_idx, body_op.Constant(value_int=chunk_length)), + num_embeds, + ) + num_chunk_frames = body_op.Sub(end_idx, start_idx) + context_end_idx = body_op.Min( + body_op.Add(end_idx, body_op.Constant(value_int=chunk_right_context)), + num_embeds, + ) + chunk_embeds = body_op.Slice( + all_chunk_embeds, + _scalar(body_op, start_idx), + _scalar(body_op, context_end_idx), + body_op.Constant(value_ints=[1]), + ) + + cached_length = body_op.Add(num_cache_frames, num_fifo_frames) + cached_embeds = body_op.Concat( + body_op.Slice( + cache_embeds, + body_op.Constant(value_ints=[0]), + _scalar(body_op, num_cache_frames), + body_op.Constant(value_ints=[1]), + ), + body_op.Slice( + fifo, + body_op.Constant(value_ints=[0]), + _scalar(body_op, num_fifo_frames), + body_op.Constant(value_ints=[1]), + ), + axis=1, + ) + chunk_input_embeds = body_op.Concat(cached_embeds, chunk_embeds, axis=1) + + ( + chunk_logits, + new_cache_embeds, + new_cache_probs, + new_fifo, + new_num_cache_frames, + new_num_fifo_frames, + new_is_compressed, + ) = self._run_chunk_and_update_cache( + body_op, + chunk_input_embeds, + cached_length, + num_chunk_frames, + cache_embeds, + cache_probs, + fifo, + num_cache_frames, + num_fifo_frames, + is_compressed, + cache_length=cache_length, + fifo_capacity=fifo_capacity, + update_period=update_period, + ) + + start_logit_idx = body_op.Mul(cached_length, body_op.Constant(value_int=factor)) + end_logit_idx = body_op.Mul( + body_op.Add(cached_length, num_chunk_frames), + body_op.Constant(value_int=factor), + ) + chunk_region = body_op.Slice( + chunk_logits, + _scalar(body_op, start_logit_idx), + _scalar(body_op, end_logit_idx), + body_op.Constant(value_ints=[1]), + ) + # Place this iteration's (non-overlapping) contribution into the + # fixed-size global accumulator by zero-padding it out to the + # full length and adding -- see ``total_raw_length``'s comment + # above for why this avoids a growing ``Concat``. + global_pad_before = body_op.Mul(start_idx, body_op.Constant(value_int=factor)) + global_pad_after = body_op.Sub( + total_raw_length, body_op.Mul(end_idx, body_op.Constant(value_int=factor)) + ) + zero1d = body_op.Constant(value_ints=[0]) + global_pads = body_op.Concat( + zero1d, + _scalar(body_op, global_pad_before), + zero1d, + zero1d, + _scalar(body_op, global_pad_after), + zero1d, + axis=0, + ) + padded_chunk_region = body_op.Pad(chunk_region, global_pads, mode="constant") + new_accumulated_logits = body_op.Add(accumulated_logits, padded_chunk_region) + + cond_out = body_op.Constant(value=ir.tensor(True)) + return ( + cond_out, + end_idx, + new_accumulated_logits, + new_cache_embeds, + new_cache_probs, + new_fifo, + new_num_cache_frames, + new_num_fifo_frames, + new_is_compressed, + ) + + zero_logits = op.CastLike( + op.Expand( + op.Constant(value_float=0.0), + op.Concat( + batch, + _scalar(op, total_raw_length), + op.Constant(value_ints=[num_speakers]), + axis=0, + ), + ), + all_chunk_embeds, + ) + init_cache_embeds = op.CastLike( + op.Expand( + op.Constant(value_float=0.0), + op.Concat(batch, op.Constant(value_ints=[cache_length, hidden_size]), axis=0), + ), + all_chunk_embeds, + ) + init_cache_probs = op.CastLike( + op.Expand( + op.Constant(value_float=0.0), + op.Concat(batch, op.Constant(value_ints=[cache_length, num_speakers]), axis=0), + ), + all_chunk_embeds, + ) + init_fifo = op.CastLike( + op.Expand( + op.Constant(value_float=0.0), + op.Concat(batch, op.Constant(value_ints=[fifo_capacity, hidden_size]), axis=0), + ), + all_chunk_embeds, + ) + + # ``subgraph()`` snapshots the *calling* builder's current scope + # stack into the new sub-builder, and node/value names auto- + # generated with an EMPTY scope stack carry no qualifying prefix at + # all (just e.g. ``v_Constant_19``) -- so two graphs built back to + # back at the same (root) scope, each with their own independent + # per-graph node counter starting at 0, can trivially produce + # colliding names once their node counts happen to line up. Push a + # dedicated scope here so every node/value auto-named *inside* the + # loop body (including further-nested ``If`` branches it builds, by + # inheritance) gets a prefix that can never collide with root-scope + # or other-scope names. + op.builder.push_module("offline_chunk_loop_body") + loop_body = op.builder.subgraph( + _loop_body, + inputs=[ + ir.Value( + name="iter_num", type=ir.TensorType(ir.DataType.INT64), shape=ir.Shape([]) + ), + ir.Value( + name="cond_in", type=ir.TensorType(ir.DataType.BOOL), shape=ir.Shape([]) + ), + ir.Value(name="start_idx"), + ir.Value(name="accumulated_logits"), + ir.Value(name="cache_embeds"), + ir.Value(name="cache_probs"), + ir.Value(name="fifo"), + ir.Value(name="num_cache_frames"), + ir.Value(name="num_fifo_frames"), + ir.Value(name="is_compressed"), + ], + outputs=[ + ir.Value(name="cond_out"), + ir.Value(name="start_idx_out"), + ir.Value(name="accumulated_logits_out"), + ir.Value(name="cache_embeds_out"), + ir.Value(name="cache_probs_out"), + ir.Value(name="fifo_out"), + ir.Value(name="num_cache_frames_out"), + ir.Value(name="num_fifo_frames_out"), + ir.Value(name="is_compressed_out"), + ], + name="offline_chunk_loop_body", + ) + op.builder.pop_module() + + ( + _, + accumulated_logits, + _, + _, + _, + _, + _, + _, + ) = op.Loop( + num_iterations, + op.Constant(value=ir.tensor(True)), + op.Constant(value_int=0), + zero_logits, + init_cache_embeds, + init_cache_probs, + init_fifo, + op.Constant(value_int=0), + op.Constant(value_int=0), + op.Constant(value=ir.tensor(False)), + body=loop_body, + _outputs=8, + ) + + logits = op.Slice( + accumulated_logits, + op.Constant(value_ints=[0]), + raw_num_frames, + op.Constant(value_ints=[1]), + ) + return op.Sigmoid(logits) + + def _call_backbone_scoped(self, op: OpBuilder, method_name: str, *args): + """Call a bound method on ``self.model``/``self.model.audio_tower`` under proper module scope. + + ``forward_streaming`` calls ``embed``/``encode``/``project_upsample`` + directly (they are not full ``forward`` passes), bypassing + ``Module.__call__``'s automatic name-qualification and parameter + realization. Mirrors that qualification manually (see + ``components/_moe.py``'s ``_realize_gate_and_get_qmoe_routing`` for + the same pattern) so weight names match the offline export exactly. + """ + backbone = self.model + audio_tower = backbone.audio_tower + builder = op.builder + builder.push_module(backbone.name or "model", type(backbone).__qualname__) + try: + for param in backbone.parameters(recurse=False): + param._realize(builder) # pylint: disable=protected-access + if method_name == "project_upsample": + return backbone.project_upsample(op, *args) + builder.push_module( + audio_tower.name or "audio_tower", type(audio_tower).__qualname__ + ) + try: + for param in audio_tower.parameters(recurse=False): + param._realize(builder) # pylint: disable=protected-access + return getattr(audio_tower, method_name)(op, *args) + finally: + builder.pop_module() + finally: + builder.pop_module() + + def forward_streaming( + self, + op: OpBuilder, + input_features: ir.Value, + num_lookahead_frames: ir.Value, + past_cache_embeds: ir.Value, + past_cache_probs: ir.Value, + past_fifo: ir.Value, + past_num_cache_frames: ir.Value, + past_num_fifo_frames: ir.Value, + past_is_compressed: ir.Value, + ) -> tuple[ir.Value, ir.Value, ir.Value, ir.Value, ir.Value, ir.Value, ir.Value]: + """One streaming chunk step, matching HuggingFace's chunked forward loop body. + + See ``Nemotron3DiarizationForAudioFrameClassification.forward``'s + per-chunk logic and ``Nemotron3DiarizationSpeakerCache.update``. Every + cache tensor is a fixed-size buffer (``streaming_speaker_cache_length``/ + ``streaming_fifo_length`` capacity); ``past_num_cache_frames``/ + ``past_num_fifo_frames`` track how much of each buffer is valid, the + same convention as a static KV cache with an explicit occupancy count. + + Returns ``(speaker_probs, cache_embeds, cache_probs, fifo, + num_cache_frames, num_fifo_frames, is_compressed)``. + """ + config = self.config + factor = config.subsampling_factor + cache_length = config.streaming_speaker_cache_length + fifo_capacity = config.streaming_fifo_length + update_period = config.streaming_speaker_cache_update_period + + # ``forward_streaming`` is invoked directly (not via ``self(op, ...)``), + # bypassing ``Module.__call__``'s automatic parameter realization, so + # ``silence_embeds`` (only referenced deep inside the compress-branch + # subgraph) must be registered as a graph initializer explicitly here. + self.silence_embeds._realize(op.builder) # type: ignore[attr-defined] + + # input_features: [B, feat_in, W] -> [B, W, feat_in]. + input_features = op.Transpose(input_features, perm=[0, 2, 1]) + raw_num_frames = op.Shape(input_features, start=1, end=2) + + chunk_embeds = self._call_backbone_scoped(op, "embed", input_features) + num_new_embeds = op.Squeeze(op.Shape(chunk_embeds, start=1, end=2), [0]) + # Callers must supply ``0 <= num_lookahead_frames < num_new_embeds`` + # (see ``DiarizationStreamingTask``'s docstring); clamp defensively so + # a caller-supplied out-of-range value can't corrupt the FIFO/cache + # occupancy bookkeeping below or produce a negative slice length. + num_lookahead_frames = op.Max( + op.Constant(value_int=0), + op.Min(num_lookahead_frames, op.Sub(num_new_embeds, op.Constant(value_int=1))), + ) + num_chunk_frames = op.Sub(num_new_embeds, num_lookahead_frames) + + cached_length = op.Add(past_num_cache_frames, past_num_fifo_frames) + cached_embeds = op.Concat( + op.Slice( + past_cache_embeds, + op.Constant(value_ints=[0]), + _scalar(op, past_num_cache_frames), + op.Constant(value_ints=[1]), + ), + op.Slice( + past_fifo, + op.Constant(value_ints=[0]), + _scalar(op, past_num_fifo_frames), + op.Constant(value_ints=[1]), + ), + axis=1, + ) + chunk_input_embeds = op.Concat(cached_embeds, chunk_embeds, axis=1) + + ( + chunk_logits, + new_cache_embeds, + new_cache_probs, + new_fifo, + new_num_cache_frames, + new_num_fifo_frames, + new_is_compressed, + ) = self._run_chunk_and_update_cache( + op, + chunk_input_embeds, + cached_length, + num_chunk_frames, + past_cache_embeds, + past_cache_probs, + past_fifo, + past_num_cache_frames, + past_num_fifo_frames, + past_is_compressed, + cache_length=cache_length, + fifo_capacity=fifo_capacity, + update_period=update_period, + ) + + # This step's speaker-probability output: the chunk region only + # (excludes both the cached prefix and the look-ahead suffix), + # trimmed to this call's raw (pre-feature-stacking-padding) length. + start_logit_idx = op.Mul(cached_length, op.Constant(value_int=factor)) + end_logit_idx = op.Mul( + op.Add(cached_length, num_chunk_frames), op.Constant(value_int=factor) + ) + chunk_region = op.Slice( + chunk_logits, + _scalar(op, start_logit_idx), + _scalar(op, end_logit_idx), + op.Constant(value_ints=[1]), + ) + # Trim to this call's raw (pre-feature-stacking-padding) window + # length, matching HF's reference `logits[:, :num_frames]` exactly + # (see `Nemotron3DiarizationForAudioFrameClassification.forward`). + # No lookahead subtraction here: `num_chunk_frames * factor` (the + # length of `chunk_region` above) is already <= `raw_num_frames` + # whenever `num_lookahead_frames > 0`, because the one padded + # feature-stacking frame (if any) falls inside the excluded + # look-ahead suffix. This `Slice` is therefore a no-op except on a + # final chunk (lookahead == 0) whose window isn't a multiple of + # `subsampling_factor`, where it trims off the padding-derived + # frame(s) — relying on `Slice`'s documented clamping of an + # out-of-range `ends` value to the actual dimension size. + chunk_region = op.Slice( + chunk_region, + op.Constant(value_ints=[0]), + raw_num_frames, + op.Constant(value_ints=[1]), + ) + speaker_probs = op.Sigmoid(chunk_region) + + return ( + speaker_probs, + new_cache_embeds, + new_cache_probs, + new_fifo, + new_num_cache_frames, + new_num_fifo_frames, + new_is_compressed, + ) + + def _run_chunk_and_update_cache( + self, + op: OpBuilder, + chunk_input_embeds: ir.Value, + cached_length: ir.Value, + num_chunk_frames: ir.Value, + past_cache_embeds: ir.Value, + past_cache_probs: ir.Value, + past_fifo: ir.Value, + past_num_cache_frames: ir.Value, + past_num_fifo_frames: ir.Value, + past_is_compressed: ir.Value, + *, + cache_length: int, + fifo_capacity: int, + update_period: int, + ) -> tuple[ir.Value, ir.Value, ir.Value, ir.Value, ir.Value, ir.Value, ir.Value]: + """Runs the encoder + classifier over one (already cache-prefixed) chunk. + + Updates the Arrival-Order Speaker Cache (AOSC) + FIFO queue. + + Shared by :meth:`forward_streaming` (streaming-sized cache) and + :meth:`forward`'s offline ``Loop`` body (offline-sized cache) -- see + ``Nemotron3DiarizationSpeakerCache.update``. ``cache_length`` / + ``fifo_capacity`` / ``update_period`` are plain ints (not graph + values) since both callers know their cache sizes statically. + + Returns ``(chunk_logits, new_cache_embeds, new_cache_probs, new_fifo, + new_num_cache_frames, new_num_fifo_frames, new_is_compressed)`` where + ``chunk_logits`` is the *full* ``chunk_input_embeds``-length logits + (unsliced) -- callers slice the chunk-only region themselves. + """ + config = self.config + factor = config.subsampling_factor + + encoded = self._call_backbone_scoped(op, "encode", chunk_input_embeds) + upsampled = self._call_backbone_scoped(op, "project_upsample", encoded) + chunk_logits = self.classifier(op, upsampled) + + # --- Arrival-Order Speaker Cache (AOSC) + FIFO update --- + probs = _avg_pool_probs(op, chunk_logits, factor) + + new_chunk_embeds = op.Slice( + chunk_input_embeds, + _scalar(op, cached_length), + _scalar(op, op.Add(cached_length, num_chunk_frames)), + op.Constant(value_ints=[1]), + ) + old_fifo = op.Slice( + past_fifo, + op.Constant(value_ints=[0]), + _scalar(op, past_num_fifo_frames), + op.Constant(value_ints=[1]), + ) + fifo_embeds = op.Concat(old_fifo, new_chunk_embeds, axis=1) + num_fifo_total = op.Add(past_num_fifo_frames, num_chunk_frames) + num_popped = _num_popped_frames(op, num_fifo_total, fifo_capacity, update_period) + + fifo_probs_all = op.Slice( + probs, + _scalar(op, past_num_cache_frames), + _scalar(op, op.Add(past_num_cache_frames, num_fifo_total)), + op.Constant(value_ints=[1]), + ) + stored_probs_uncompressed = op.Slice( + probs, + op.Constant(value_ints=[0]), + _scalar(op, past_num_cache_frames), + op.Constant(value_ints=[1]), + ) + stored_probs_compressed = op.Slice( + past_cache_probs, + op.Constant(value_ints=[0]), + _scalar(op, past_num_cache_frames), + op.Constant(value_ints=[1]), + ) + stored_probs = op.Where( + op.Reshape(past_is_compressed, op.Constant(value_ints=[1, 1, 1])), + stored_probs_compressed, + stored_probs_uncompressed, + ) + + popped_embeds = op.Slice( + fifo_embeds, + op.Constant(value_ints=[0]), + _scalar(op, num_popped), + op.Constant(value_ints=[1]), + ) + popped_probs = op.Slice( + fifo_probs_all, + op.Constant(value_ints=[0]), + _scalar(op, num_popped), + op.Constant(value_ints=[1]), + ) + old_cache_embeds = op.Slice( + past_cache_embeds, + op.Constant(value_ints=[0]), + _scalar(op, past_num_cache_frames), + op.Constant(value_ints=[1]), + ) + cache_embeds_candidate = op.Concat(old_cache_embeds, popped_embeds, axis=1) + cache_probs_candidate = op.Concat(stored_probs, popped_probs, axis=1) + combined_count = op.Add(past_num_cache_frames, num_popped) + needs_compress = op.Greater(combined_count, op.Constant(value_int=cache_length)) + + def _compress_branch(branch_op: OpBuilder): + embeds_out, probs_out = _compress_speaker_cache( + branch_op, + cache_embeds_candidate, + cache_probs_candidate, + self.silence_embeds, + config, + ) + count_out = branch_op.Constant(value_int=cache_length) + return embeds_out, probs_out, count_out + + def _passthrough_branch(branch_op: OpBuilder): + embeds_out = _pad_time_axis( + branch_op, cache_embeds_candidate, branch_op.Constant(value_int=cache_length) + ) + probs_out = _pad_time_axis( + branch_op, cache_probs_candidate, branch_op.Constant(value_int=cache_length) + ) + # Recompute (rather than pass through) `combined_count`: a bare + # outer-scope value cannot be a subgraph output (and any + # `Identity` wrapper added purely for that purpose is eliminated + # by mobius's cleanup optimization pass, which also recurses into + # subgraphs), so this must be a genuine op inside the branch. + count_out = branch_op.Add(past_num_cache_frames, num_popped) + return embeds_out, probs_out, count_out + + # Each ``If`` branch is its own subgraph with an independent node + # counter; without a distinguishing scope, two branches unlucky + # enough to reach the same node count (e.g. both start with a + # ``Constant``) produce colliding auto-generated names -- push a + # unique scope per branch so this can never happen, even when this + # ``If`` is built at the same (or repeatedly re-entered, e.g. inside + # a ``Loop`` body) outer scope. See ``_prerealize_parameters``'s + # docstring for the related root/subgraph collision this mirrors. + op.builder.push_module("compress_speaker_cache") + then_branch = op.builder.subgraph( + _compress_branch, + inputs=[], + outputs=[ + ir.Value(name="compressed_cache_embeds"), + ir.Value(name="compressed_cache_probs"), + ir.Value(name="compressed_num_cache_frames"), + ], + name="compress_speaker_cache", + ) + op.builder.pop_module() + op.builder.push_module("passthrough_speaker_cache") + else_branch = op.builder.subgraph( + _passthrough_branch, + inputs=[], + outputs=[ + ir.Value(name="uncompressed_cache_embeds"), + ir.Value(name="uncompressed_cache_probs"), + ir.Value(name="uncompressed_num_cache_frames"), + ], + name="passthrough_speaker_cache", + ) + op.builder.pop_module() + new_cache_embeds, new_cache_probs, new_num_cache_frames = op.If( + needs_compress, then_branch=then_branch, else_branch=else_branch, _outputs=3 + ) + new_is_compressed = op.Or(past_is_compressed, needs_compress) + + new_fifo_content = op.Slice( + fifo_embeds, + _scalar(op, num_popped), + _scalar(op, num_fifo_total), + op.Constant(value_ints=[1]), + ) + new_num_fifo_frames = op.Sub(num_fifo_total, num_popped) + new_fifo = _pad_time_axis(op, new_fifo_content, op.Constant(value_int=fifo_capacity)) + + return ( + chunk_logits, + new_cache_embeds, + new_cache_probs, + new_fifo, + new_num_cache_frames, + new_num_fifo_frames, + new_is_compressed, + ) diff --git a/src/mobius/tasks/__init__.py b/src/mobius/tasks/__init__.py index 1677cd01e..72ca3bfd6 100644 --- a/src/mobius/tasks/__init__.py +++ b/src/mobius/tasks/__init__.py @@ -39,6 +39,7 @@ "Qwen4ExpVisionLanguageTask", "DenoisingTask", "DiarizationTask", + "DiarizationStreamingTask", "FeatureExtractionTask", "GGUFEncoderFeatureExtractionTask", "GGUFAudioProjectorModel", @@ -142,6 +143,7 @@ from mobius.tasks._denoising import DenoisingTask from mobius.tasks._dflash import DFlashDraftTask from mobius.tasks._diarization import DiarizationTask +from mobius.tasks._diarization_streaming import DiarizationStreamingTask from mobius.tasks._draft_target import DraftTargetCausalLMTask from mobius.tasks._eagle3 import Eagle3DraftTask from mobius.tasks._falcon_h1 import FalconH1CausalLMTask @@ -246,6 +248,7 @@ "controlnet": ControlNetTask, "denoising": DenoisingTask, "diarization": DiarizationTask, + "diarization-streaming": DiarizationStreamingTask, "feature-extraction": FeatureExtractionTask, "gguf-encoder-feature-extraction": GGUFEncoderFeatureExtractionTask, "gguf-embedding-feature-extraction": GGUFEmbeddingFeatureExtractionTask, diff --git a/src/mobius/tasks/_diarization.py b/src/mobius/tasks/_diarization.py index 3f2f44aec..b81523404 100644 --- a/src/mobius/tasks/_diarization.py +++ b/src/mobius/tasks/_diarization.py @@ -22,7 +22,11 @@ class DiarizationTask(ModelTask): Input: ``input_features`` — ``[batch, feat, time]`` mel spectrogram. Output: ``speaker_probs`` — ``[batch, frames, num_spks]`` sigmoid - probabilities (``frames = time / subsampling_factor``). + probabilities. The ``time`` -> ``frames`` relationship is model-defined + (e.g. Sortformer downsamples by ``subsampling_factor``, so + ``frames = time / subsampling_factor``; Nemotron3 upsamples its encoder + output back to per-mel-frame resolution, so ``frames == time``) -- + consult the specific model class's own docstring for its exact formula. """ model_roles: ClassVar[dict[str, str]] = {"model": "encoder"} diff --git a/src/mobius/tasks/_diarization_streaming.py b/src/mobius/tasks/_diarization_streaming.py new file mode 100644 index 000000000..27f904c3b --- /dev/null +++ b/src/mobius/tasks/_diarization_streaming.py @@ -0,0 +1,155 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Streaming (per-chunk, stateful) speaker-diarization task. + +Builds a single ONNX graph that processes one audio chunk per call, carrying +the Arrival-Order Speaker Cache (AOSC) + FIFO queue state across calls as +fixed-size cache buffers plus scalar occupancy counters — the same "static +cache" convention used for KV caches elsewhere in mobius (see +``tasks/_cache_utils.py``). +""" + +from __future__ import annotations + +from typing import ClassVar + +import onnx_ir as ir + +from mobius._configs import Nemotron3DiarizationConfig +from mobius._model_package import ModelPackage +from mobius.tasks._base import ModelTask, _make_graph, _make_model + + +class DiarizationStreamingTask(ModelTask): + """Build an incremental, per-chunk ONNX graph for streaming speaker diarization. + + Input: + ``input_features`` — ``[batch, feat, chunk_window_frames]`` raw + (pre-subsampling) mel-spectrogram window: the chunk plus its + look-ahead frames, at ``config.subsampling_factor`` frames per + encoder frame. ``chunk_window_frames`` is a symbolic (dynamic) + axis — it defaults to ``(chunk_length + chunk_right_context) * + subsampling_factor`` but callers may pass a shorter window (with + ``num_lookahead_frames=0``) for a recording's final, partial + chunk; there is no padding/validity contract, so every frame in + the window must be real audio. + ``num_lookahead_frames`` — scalar ``int64``: how many trailing encoder + frames of the window are look-ahead only (attended to, but not + emitted or pushed to the FIFO). ``0`` for the last chunk of a + recording. Contract: ``0 <= num_lookahead_frames < num_new_embeds`` + (the number of encoder frames the window embeds to, i.e. + ``chunk_window_frames // subsampling_factor``, rounded up); the + graph clamps out-of-range values into this range defensively + (see ``forward_streaming``) rather than failing, but callers + should not rely on the clamped result being meaningful. + ``past_cache_embeds`` / ``past_cache_probs`` / ``past_fifo`` — fixed + capacity AOSC + FIFO state buffers from the previous call (all + zeros for the first chunk of a stream). + ``past_num_cache_frames`` / ``past_num_fifo_frames`` — scalar + ``int64`` occupancy counters for the two buffers above. + ``past_is_compressed`` — scalar ``bool``: whether the AOSC has ever + been compressed (governs how cached-frame probabilities are + re-estimated on the next call). + + Output: + ``speaker_probs`` — ``[batch, frames, num_spks]`` sigmoid + probabilities for this chunk only. ``frames`` equals + ``min(chunk_window_frames, (num_new_embeds - num_lookahead_frames) * + subsampling_factor)`` — i.e. it excludes the look-ahead frames and is + capped to the actual raw window length (matching HuggingFace's + reference ``logits[:, :num_frames]`` trim exactly). This is always + exactly the number of frames committed to the FIFO/cache this call + (``num_chunk_frames * subsampling_factor``), so the output length + never desyncs from cache growth — including for a non-final chunk + whose window isn't a multiple of ``subsampling_factor``, where the + cap only ever activates on a final (``num_lookahead_frames == 0``) + chunk with a padding-derived trailing frame. + ``present_cache_embeds`` / ``present_cache_probs`` / ``present_fifo`` + / ``present_num_cache_frames`` / ``present_num_fifo_frames`` / + ``present_is_compressed`` — updated state, fed back as the ``past_*`` + inputs of the next call. + """ + + model_roles: ClassVar[dict[str, str]] = {"model": "encoder"} + + def build( + self, + module, + config: Nemotron3DiarizationConfig, + ) -> ModelPackage: + graph, builder = _make_graph(name="nemotron3_diarization_streaming") + op = builder.op + + cache_capacity = config.streaming_speaker_cache_length + fifo_capacity = config.streaming_fifo_length + + # The time axis is symbolic (not the default full-window + # ``chunk_length + chunk_right_context`` frame count): HuggingFace's + # streaming API allows a shorter final chunk (with + # ``num_lookahead_frames=0``, no padding contract). ``forward_streaming`` + # already derives every frame count it needs from this input's actual + # shape (``raw_num_frames = op.Shape(...)``), so no other graph logic + # assumes the default window length. + input_features = builder.input( + "input_features", + dtype=config.dtype, + shape=["batch", config.feat_in, "chunk_window_frames"], + ) + num_lookahead_frames = builder.input( + "num_lookahead_frames", dtype=ir.DataType.INT64, shape=[] + ) + past_cache_embeds = builder.input( + "past_cache_embeds", + dtype=config.dtype, + shape=["batch", cache_capacity, config.hidden_size], + ) + past_cache_probs = builder.input( + "past_cache_probs", + dtype=config.dtype, + shape=["batch", cache_capacity, config.num_speakers], + ) + past_fifo = builder.input( + "past_fifo", + dtype=config.dtype, + shape=["batch", fifo_capacity, config.hidden_size], + ) + past_num_cache_frames = builder.input( + "past_num_cache_frames", dtype=ir.DataType.INT64, shape=[] + ) + past_num_fifo_frames = builder.input( + "past_num_fifo_frames", dtype=ir.DataType.INT64, shape=[] + ) + past_is_compressed = builder.input( + "past_is_compressed", dtype=ir.DataType.BOOL, shape=[] + ) + + ( + speaker_probs, + present_cache_embeds, + present_cache_probs, + present_fifo, + present_num_cache_frames, + present_num_fifo_frames, + present_is_compressed, + ) = module.forward_streaming( + op, + input_features, + num_lookahead_frames, + past_cache_embeds, + past_cache_probs, + past_fifo, + past_num_cache_frames, + past_num_fifo_frames, + past_is_compressed, + ) + + builder.add_output(speaker_probs, "speaker_probs") + builder.add_output(present_cache_embeds, "present_cache_embeds") + builder.add_output(present_cache_probs, "present_cache_probs") + builder.add_output(present_fifo, "present_fifo") + builder.add_output(present_num_cache_frames, "present_num_cache_frames") + builder.add_output(present_num_fifo_frames, "present_num_fifo_frames") + builder.add_output(present_is_compressed, "present_is_compressed") + + return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/testdata/cases/diarization/nemotron3_diarization.yaml b/testdata/cases/diarization/nemotron3_diarization.yaml new file mode 100644 index 000000000..61bb5889a --- /dev/null +++ b/testdata/cases/diarization/nemotron3_diarization.yaml @@ -0,0 +1,30 @@ +model_id: "nvidia/Nemotron-3-Diarization" +model_type: "nemotron3_diarization" +revision: "f667ed73aee57d40cc39428eb768b4fd87a0a29e" +task_type: "diarization" +dtype: "float32" + +level: "L4" + +notes: > + Offline single-pass speaker diarization. The golden reference is a + deterministic synthetic mel fixture (seeded torch.randn, below + chunk_length * subsampling_factor so HuggingFace's own offline forward + does not internally re-chunk) run through the real HuggingFace + Nemotron3DiarizationForAudioFrameClassification checkpoint -- diarization + models take raw mel input_features directly (no audio-waveform processor + step), so a real recorded audio fixture is not required to exercise every + op in the exported graph. The reference stores per-frame per-speaker + sigmoid probabilities in the companion nemotron3_diarization_probs.npz + (see mobius._testing.golden.save_diarization_golden); compare_golden()'s + argmax/top-K gate does not apply to this continuous output, so this case + is compared with compare_diarization_golden() (elementwise atol/rtol, + with per-frame dominant-speaker argmax agreement and active-speaker-set + Jaccard as diagnostics). The reference was generated with an unreleased + transformers snapshot (Nemotron3Diarization was not yet in any released + version) -- the golden's provenance.transformers_commit records the + exact upstream git commit used, for reproducibility. Regenerate with a + matching pinned install, e.g. + `pip install "git+https://github.com/huggingface/transformers.git@"`, + then + `python scripts/generate_golden.py --case testdata/cases/diarization/nemotron3_diarization.yaml`. diff --git a/testdata/cases/diarization/nemotron3_diarization_multichunk.yaml b/testdata/cases/diarization/nemotron3_diarization_multichunk.yaml new file mode 100644 index 000000000..942c67881 --- /dev/null +++ b/testdata/cases/diarization/nemotron3_diarization_multichunk.yaml @@ -0,0 +1,40 @@ +model_id: "nvidia/Nemotron-3-Diarization" +model_type: "nemotron3_diarization" +revision: "f667ed73aee57d40cc39428eb768b4fd87a0a29e" +task_type: "diarization" +dtype: "float32" + +level: "L4" + +generation: + num_frames: 6000 + +notes: > + Offline multi-chunk speaker diarization (a longer companion to + nemotron3_diarization.yaml). 6000 mel frames subsample to 750 encoder + embeddings -- well above config.chunk_length (340) -- so HuggingFace's own + offline forward internally re-chunks into 3 iterations (340 + 340 + 70 + frames), threading its Arrival-Order Speaker Cache (AOSC) + FIFO queue + (config.fifo_length / config.speaker_cache_update_period, distinct from + streaming_config's sizes) across chunk boundaries and, since + speaker_cache_update_period (300) is exceeded partway through, exercising + cache compression too. This is the reference this case's mobius ONNX + export must match: an ONNX Loop-based chunked offline forward (see + Nemotron3DiarizationForAudioFrameClassification.forward's + `for start_idx in range(0, num_chunk_embeds, chunk_length)` loop), not the + single flat non-causal pass nemotron3_diarization.yaml's shorter (single- + chunk) fixture exercises. The reference stores per-frame per-speaker + sigmoid probabilities in the companion + nemotron3_diarization_multichunk_probs.npz (see + mobius._testing.golden.save_diarization_golden); compare_golden()'s + argmax/top-K gate does not apply to this continuous output, so this case + is compared with compare_diarization_golden() (elementwise atol/rtol, + with per-frame dominant-speaker argmax agreement and active-speaker-set + Jaccard as diagnostics). The reference was generated with an unreleased + transformers snapshot (Nemotron3Diarization was not yet in any released + version) -- the golden's provenance.transformers_commit records the + exact upstream git commit used, for reproducibility. Regenerate with a + matching pinned install, e.g. + `pip install "git+https://github.com/huggingface/transformers.git@"`, + then + `python scripts/generate_golden.py --case testdata/cases/diarization/nemotron3_diarization_multichunk.yaml`. diff --git a/testdata/cases/diarization/nemotron3_diarization_streaming.yaml b/testdata/cases/diarization/nemotron3_diarization_streaming.yaml new file mode 100644 index 000000000..3df2a4aa8 --- /dev/null +++ b/testdata/cases/diarization/nemotron3_diarization_streaming.yaml @@ -0,0 +1,33 @@ +model_id: "nvidia/Nemotron-3-Diarization" +model_type: "nemotron3_diarization" +revision: "f667ed73aee57d40cc39428eb768b4fd87a0a29e" +task_type: "diarization-streaming" +dtype: "float32" + +level: "L5" + +notes: > + 3-chunk streaming diarization session. Each chunk is a full + (chunk_length + chunk_right_context) window of deterministic synthetic mel + (seeded torch.randn per chunk) run through HuggingFace's real per-chunk + forward() + Nemotron3DiarizationSpeakerCache bookkeeping. With this + checkpoint's default streaming config (FIFO capacity 264, speaker-cache + capacity 264, update period 222) a single default-sized chunk (340 encoder + frames) already overflows the FIFO, so this exercises both the + FIFO-to-cache eviction path and the top-k Arrival-Order Speaker Cache + (AOSC) compression path by the second chunk. The companion + nemotron3_diarization_streaming_probs.npz stores, per chunk i: the mel + input (chunk{i}_mel), per-frame speaker probabilities (chunk{i}_probs), + and the full AOSC + FIFO cache state after that chunk (chunk{i}_cache_*, + chunk{i}_fifo, chunk{i}_num_cache_frames, chunk{i}_num_fifo_frames, + chunk{i}_is_compressed) so the L5 test can thread the real exported + streaming ONNX graph's own outputs back in as the next chunk's cache + inputs and verify each step against the independently-computed HuggingFace + reference. The reference was generated with an unreleased transformers + snapshot (Nemotron3Diarization was not yet in any released version) -- + the golden's provenance.transformers_commit records the exact upstream + git commit used, for reproducibility. Regenerate with a matching pinned + install, e.g. + `pip install "git+https://github.com/huggingface/transformers.git@"`, + then + `python scripts/generate_golden.py --case testdata/cases/diarization/nemotron3_diarization_streaming.yaml`. diff --git a/testdata/cases/diarization/sortformer.yaml b/testdata/cases/diarization/sortformer.yaml new file mode 100644 index 000000000..ea37790f5 --- /dev/null +++ b/testdata/cases/diarization/sortformer.yaml @@ -0,0 +1,26 @@ +model_id: "nvidia/diar_streaming_sortformer_4spk-v2.1" +model_type: "sortformer" +revision: "fafaab5faa1617a0ca52d38dd3dc4bd636800d3d" +task_type: "diarization" +dtype: "float32" +reference_loader: "nemo" + +level: "L4+L5" + +notes: > + Offline (non-streaming) speaker diarization. The golden reference is a + deterministic synthetic mel fixture (seeded torch.randn) run through the + real NeMo SortformerEncLabelModel (frontend_encoder -> forward_infer), + restored from the .nemo archive on HuggingFace Hub via nemo_toolkit (not a + mobius runtime dependency -- see scripts/generate_golden.py's sortformer + branch). The companion sortformer_probs.npz stores the mel input and the + per-frame per-speaker sigmoid probabilities; compare_golden()'s argmax/ + top-K gate does not apply to this continuous output, so L4 is compared + with compare_diarization_golden() (elementwise atol/rtol). L5 additionally + checks that the two downstream diarization decisions a consumer reads -- + the per-frame dominant-speaker argmax and the binarized (threshold 0.5) + active-speaker set -- reproduce NeMo's exactly over the whole utterance. + The golden's provenance.nemo_commit records the exact upstream nemo_toolkit + git commit used, for reproducibility. Regenerate with + `python scripts/generate_golden.py --case testdata/cases/diarization/sortformer.yaml` + inside an environment with `pip install "nemo_toolkit[asr]"`. diff --git a/testdata/cases/schema.json b/testdata/cases/schema.json index bf9d68a6b..407cdfbd2 100644 --- a/testdata/cases/schema.json +++ b/testdata/cases/schema.json @@ -46,6 +46,8 @@ "codec", "ctc-asr", "depth-estimation", + "diarization", + "diarization-streaming", "dflash-draft", "feature-extraction", "feature-ctc-asr", @@ -429,10 +431,28 @@ "target_model_id": { "type": "string", "description": "Paired target model HF ID for speculative-decoding drafter tasks (e.g. dflash-draft): the target supplies the hidden states / shared KV used to build the drafter's golden inputs at generation time. Only read by generate_golden.py; the e2e test replays stored tensors and never loads the target." + }, + "num_frames": { + "type": "integer", + "minimum": 1, + "description": "Number of synthetic mel frames to generate for audio/diarization fixtures. Only read by generate_golden.py when synthesizing input features; not a HuggingFace GenerationConfig field." + }, + "num_chunks": { + "type": "integer", + "minimum": 1, + "description": "Number of streaming chunks to synthesize for streaming audio/diarization fixtures. Only read by generate_golden.py; not a HuggingFace GenerationConfig field." + }, + "seed": { + "type": "integer", + "description": "Random seed used to synthesize deterministic fixture inputs (e.g. mel features). Only read by generate_golden.py; not a HuggingFace GenerationConfig field." + }, + "nemo_filename": { + "type": "string", + "description": "Filename of a NeMo-format checkpoint asset to load for fixture generation. Only read by generate_golden.py; not a HuggingFace GenerationConfig field." } }, "additionalProperties": false, - "description": "HuggingFace GenerationConfig overrides. Only applicable for L5 generation tests." + "description": "Generation-time overrides read by generate_golden.py. Most properties are HuggingFace GenerationConfig overrides applicable only to L5 generation tests; num_frames/num_chunks/seed/nemo_filename instead configure synthetic audio/diarization fixture generation and are unrelated to GenerationConfig." }, "trust_remote_code": { "type": "boolean", @@ -443,10 +463,11 @@ "type": "string", "enum": [ "causal-lm", - "multimodal" + "multimodal", + "nemo" ], "default": "causal-lm", - "description": "HuggingFace loader used to generate the reference. Use multimodal when a text-only export is sourced from a composite checkpoint." + "description": "HuggingFace loader used to generate the reference. Use multimodal when a text-only export is sourced from a composite checkpoint. Use nemo when the reference is a NeMo toolkit model restored from a .nemo archive rather than a HuggingFace transformers checkpoint." }, "skip_reason": { "type": [ diff --git a/testdata/golden/diarization/nemotron3_diarization.json b/testdata/golden/diarization/nemotron3_diarization.json new file mode 100644 index 000000000..d5a7f7230 --- /dev/null +++ b/testdata/golden/diarization/nemotron3_diarization.json @@ -0,0 +1,24 @@ +{ + "array_shapes": { + "mel": [ + 1, + 128, + 2000 + ], + "probs": [ + 1, + 2000, + 8 + ] + }, + "provenance": { + "model_id": "nvidia/Nemotron-3-Diarization", + "revision": "f667ed73aee57d40cc39428eb768b4fd87a0a29e", + "transformers_version": "5.18.0.dev0", + "transformers_commit": "0291458166a6834a37b179a24c5235f242dadea0", + "torch_version": "2.14.0+cu130", + "seed": 0, + "num_speakers": 8, + "mel_dim": 128 + } +} diff --git a/testdata/golden/diarization/nemotron3_diarization_multichunk.json b/testdata/golden/diarization/nemotron3_diarization_multichunk.json new file mode 100644 index 000000000..fe776f1f8 --- /dev/null +++ b/testdata/golden/diarization/nemotron3_diarization_multichunk.json @@ -0,0 +1,24 @@ +{ + "array_shapes": { + "mel": [ + 1, + 128, + 6000 + ], + "probs": [ + 1, + 6000, + 8 + ] + }, + "provenance": { + "model_id": "nvidia/Nemotron-3-Diarization", + "revision": "f667ed73aee57d40cc39428eb768b4fd87a0a29e", + "transformers_version": "5.18.0.dev0", + "transformers_commit": "0291458166a6834a37b179a24c5235f242dadea0", + "torch_version": "2.14.0+cu130", + "seed": 0, + "num_speakers": 8, + "mel_dim": 128 + } +} diff --git a/testdata/golden/diarization/nemotron3_diarization_multichunk_probs.npz b/testdata/golden/diarization/nemotron3_diarization_multichunk_probs.npz new file mode 100644 index 000000000..c0d2af16f Binary files /dev/null and b/testdata/golden/diarization/nemotron3_diarization_multichunk_probs.npz differ diff --git a/testdata/golden/diarization/nemotron3_diarization_probs.npz b/testdata/golden/diarization/nemotron3_diarization_probs.npz new file mode 100644 index 000000000..578254ac3 Binary files /dev/null and b/testdata/golden/diarization/nemotron3_diarization_probs.npz differ diff --git a/testdata/golden/diarization/nemotron3_diarization_streaming.json b/testdata/golden/diarization/nemotron3_diarization_streaming.json new file mode 100644 index 000000000..114382ea1 --- /dev/null +++ b/testdata/golden/diarization/nemotron3_diarization_streaming.json @@ -0,0 +1,109 @@ +{ + "array_shapes": { + "num_chunks": [], + "chunk0_mel": [ + 1, + 128, + 3040 + ], + "chunk0_lookahead": [], + "chunk0_probs": [ + 1, + 2720, + 8 + ], + "chunk0_cache_embeds": [ + 1, + 264, + 512 + ], + "chunk0_cache_probs": [ + 1, + 264, + 8 + ], + "chunk0_fifo": [ + 1, + 264, + 512 + ], + "chunk0_num_cache_frames": [], + "chunk0_num_fifo_frames": [], + "chunk0_is_compressed": [], + "chunk1_mel": [ + 1, + 128, + 3040 + ], + "chunk1_lookahead": [], + "chunk1_probs": [ + 1, + 2720, + 8 + ], + "chunk1_cache_embeds": [ + 1, + 264, + 512 + ], + "chunk1_cache_probs": [ + 1, + 264, + 8 + ], + "chunk1_fifo": [ + 1, + 264, + 512 + ], + "chunk1_num_cache_frames": [], + "chunk1_num_fifo_frames": [], + "chunk1_is_compressed": [], + "chunk2_mel": [ + 1, + 128, + 3040 + ], + "chunk2_lookahead": [], + "chunk2_probs": [ + 1, + 3040, + 8 + ], + "chunk2_cache_embeds": [ + 1, + 264, + 512 + ], + "chunk2_cache_probs": [ + 1, + 264, + 8 + ], + "chunk2_fifo": [ + 1, + 264, + 512 + ], + "chunk2_num_cache_frames": [], + "chunk2_num_fifo_frames": [], + "chunk2_is_compressed": [] + }, + "provenance": { + "model_id": "nvidia/Nemotron-3-Diarization", + "revision": "f667ed73aee57d40cc39428eb768b4fd87a0a29e", + "transformers_version": "5.18.0.dev0", + "transformers_commit": "0291458166a6834a37b179a24c5235f242dadea0", + "torch_version": "2.14.0+cu130", + "seed": 0, + "num_speakers": 8, + "mel_dim": 128, + "chunk_length": 340, + "chunk_right_context": 40, + "subsampling_factor": 8, + "streaming_fifo_length": 264, + "streaming_speaker_cache_length": 264, + "streaming_speaker_cache_update_period": 222, + "num_stream_chunks": 3 + } +} diff --git a/testdata/golden/diarization/nemotron3_diarization_streaming_probs.npz b/testdata/golden/diarization/nemotron3_diarization_streaming_probs.npz new file mode 100644 index 000000000..12d9d1e3c Binary files /dev/null and b/testdata/golden/diarization/nemotron3_diarization_streaming_probs.npz differ diff --git a/testdata/golden/diarization/sortformer.json b/testdata/golden/diarization/sortformer.json new file mode 100644 index 000000000..25ecfddcf --- /dev/null +++ b/testdata/golden/diarization/sortformer.json @@ -0,0 +1,32 @@ +{ + "array_shapes": { + "mel": [ + 1, + 128, + 400 + ], + "emb_seq": [ + 1, + 50, + 192 + ], + "emb_len": [ + 1 + ], + "probs": [ + 1, + 50, + 4 + ] + }, + "provenance": { + "model_id": "nvidia/diar_streaming_sortformer_4spk-v2.1", + "revision": "fafaab5faa1617a0ca52d38dd3dc4bd636800d3d", + "nemo_version": "3.1.0+abb8254da", + "nemo_commit": "abb8254dac2bf5a011e6069fcaa7df71c7e3b8c1", + "seed": 0, + "feat_dim": 128, + "num_frames": 400, + "num_speakers": 4 + } +} diff --git a/testdata/golden/diarization/sortformer_probs.npz b/testdata/golden/diarization/sortformer_probs.npz new file mode 100644 index 000000000..cd3918576 Binary files /dev/null and b/testdata/golden/diarization/sortformer_probs.npz differ diff --git a/testdata/golden/speech/sortformer_diarization.npz b/testdata/golden/speech/sortformer_diarization.npz deleted file mode 100644 index 905c64c01..000000000 Binary files a/testdata/golden/speech/sortformer_diarization.npz and /dev/null differ diff --git a/tests/_test_configs.py b/tests/_test_configs.py index 51401b492..25efe9fac 100644 --- a/tests/_test_configs.py +++ b/tests/_test_configs.py @@ -51,6 +51,7 @@ MoonshineStreamingConfig, MuseGlimmerConfig, NanoChatConfig, + Nemotron3DiarizationConfig, NemotronHConfig, NemotronParseConfig, ParakeetCTCConfig, @@ -3548,6 +3549,25 @@ def vl_overrides(model_type: str) -> dict: }, True, ), + # --- Nemotron 3 Diarization (bidirectional full-RoPE Sortformer head) --- + ( + "nemotron3_diarization", + { + "_config_cls": Nemotron3DiarizationConfig, + "feat_in": 16, + "subsampling_factor": 2, + "head_hidden_size": 8, + "num_speakers": 4, + "partial_rotary_factor": 1.0, + # HuggingFace's Nemotron3Diarization audio encoder is always MHA + # (checkpoint's audio_config sets num_key_value_heads == + # num_attention_heads); the shared tiny-config default is GQA + # (num_key_value_heads=TINY_KV_HEADS < TINY_HEADS), which builds a + # V-projection shape no real checkpoint could ever populate. + "num_key_value_heads": TINY_HEADS, + }, + True, + ), # --- Moonshine (raw-waveform RoPE encoder-decoder ASR) --- ( "moonshine", diff --git a/tests/build_graph/speech_test.py b/tests/build_graph/speech_test.py index 58a44ac81..c8f600dff 100644 --- a/tests/build_graph/speech_test.py +++ b/tests/build_graph/speech_test.py @@ -33,6 +33,7 @@ from mobius._configs import ( AudioConfig, CodePredictorConfig, + Nemotron3DiarizationConfig, SpeakerEncoderConfig, TTSConfig, ) @@ -1386,3 +1387,400 @@ def test_runs_with_random_weights(self): assert out.shape == (1, n_time // config.fc_subsampling_factor, config.num_spks) # Sigmoid output must lie in [0, 1]. assert out.min() >= 0.0 and out.max() <= 1.0 + + +class TestBuildGraphNemotron3DiarizationOffline: + """Verify Nemotron3Diarization builds a chunked (``Loop``-based) offline graph.""" + + def _config(self, **overrides): + defaults = dict( + feat_in=16, + subsampling_factor=2, + head_hidden_size=8, + num_speakers=4, + partial_rotary_factor=1.0, + chunk_length=4, + chunk_right_context=2, + streaming_speaker_cache_length=6, + offline_fifo_length=3, + offline_speaker_cache_update_period=4, + ) + defaults.update(overrides) + return _base_config(Nemotron3DiarizationConfig, **defaults) + + def test_package_builds(self): + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationTask + + config = self._config() + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationTask()) + + assert "model" in pkg + + def test_graph_contains_loop(self): + """Verify the offline graph is chunked via an ONNX ``Loop`` node.""" + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationTask + + config = self._config() + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationTask()) + model = pkg["model"] + + op_types = {node.op_type for node in model.graph} + assert "Loop" in op_types + + def _run_with_random_weights(self, config, num_raw_frames: int) -> np.ndarray: + import os + import tempfile + + import onnxruntime as ort + + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationTask + + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationTask()) + model = pkg["model"] + + for init in model.graph.initializers.values(): + if init.const_value is not None: + continue + shape = [d if isinstance(d, int) else 1 for d in init.shape] + arr = (np.random.randn(*shape) * 0.02).astype(np.float32) + init.const_value = ir.tensor(arr, name=init.name) + + with tempfile.TemporaryDirectory() as tmp: + path = os.path.join(tmp, "model.onnx") + ir.save(model, path, external_data="model.onnx.data") + sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) + feats = np.random.randn(1, config.feat_in, num_raw_frames).astype(np.float32) + return sess.run(None, {sess.get_inputs()[0].name: feats})[0] + + def test_single_chunk_matches_expected_shape(self): + """A recording that fits in one ``chunk_length`` still runs (1 Loop iteration).""" + config = self._config() + # One encoder-frame chunk's worth of raw frames (no lookahead needed + # since it's the only/final chunk). + num_raw_frames = config.chunk_length * config.subsampling_factor + out = self._run_with_random_weights(config, num_raw_frames) + + assert out.shape == (1, num_raw_frames, config.num_speakers) + assert out.min() >= 0.0 and out.max() <= 1.0 + + def test_multi_chunk_recording_runs_the_loop_multiple_times(self): + """A recording spanning several ``chunk_length``-sized chunks. + + This is the scenario the previous (non-chunked, single-pass) offline + graph got wrong for long recordings: exercise more than one ``Loop`` + iteration end-to-end through ORT. + """ + config = self._config() + # A little over 3 chunks' worth of raw frames -> 4 Loop iterations. + num_raw_frames = (config.chunk_length * 3 + 1) * config.subsampling_factor + out = self._run_with_random_weights(config, num_raw_frames) + + assert out.shape == (1, num_raw_frames, config.num_speakers) + assert out.min() >= 0.0 and out.max() <= 1.0 + assert not np.isnan(out).any() + + def test_multi_chunk_triggers_cache_compression(self): + """A long-enough recording that the AOSC actually needs to compress. + + Uses a small ``streaming_speaker_cache_length`` and several chunks so + the cache overflows and the ``If``-based compress branch (nested + inside the ``Loop`` body) is actually exercised. + """ + config = self._config( + streaming_speaker_cache_length=3, + offline_fifo_length=2, + offline_speaker_cache_update_period=2, + ) + num_raw_frames = (config.chunk_length * 4 + 1) * config.subsampling_factor + out = self._run_with_random_weights(config, num_raw_frames) + + assert out.shape == (1, num_raw_frames, config.num_speakers) + assert out.min() >= 0.0 and out.max() <= 1.0 + assert not np.isnan(out).any() + + +class TestBuildGraphNemotron3DiarizationStreaming: + """Verify Nemotron3Diarization builds a streaming, per-chunk graph.""" + + def _config(self): + return _base_config( + Nemotron3DiarizationConfig, + feat_in=16, + subsampling_factor=2, + head_hidden_size=8, + num_speakers=4, + partial_rotary_factor=1.0, + chunk_length=4, + chunk_right_context=2, + streaming_fifo_length=6, + streaming_speaker_cache_length=6, + streaming_speaker_cache_update_period=4, + ) + + def test_package_builds(self): + """Build the streaming graph and verify a single 'model' component.""" + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationStreamingTask + + config = self._config() + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationStreamingTask()) + + assert "model" in pkg + + def test_model_io(self): + """Verify streaming input/output names, including cache state I/O.""" + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationStreamingTask + + config = self._config() + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationStreamingTask()) + model = pkg["model"] + + input_names = {inp.name for inp in model.graph.inputs} + output_names = {out.name for out in model.graph.outputs} + assert input_names == { + "input_features", + "num_lookahead_frames", + "past_cache_embeds", + "past_cache_probs", + "past_fifo", + "past_num_cache_frames", + "past_num_fifo_frames", + "past_is_compressed", + } + assert output_names == { + "speaker_probs", + "present_cache_embeds", + "present_cache_probs", + "present_fifo", + "present_num_cache_frames", + "present_num_fifo_frames", + "present_is_compressed", + } + + def test_task_registry_lookup(self): + """Verify the 'diarization-streaming' task resolves to DiarizationStreamingTask.""" + from mobius.tasks import DiarizationStreamingTask, get_task + + assert isinstance(get_task("diarization-streaming"), DiarizationStreamingTask) + + def test_runs_multiple_chunks_with_random_weights(self): + """Fill random weights and run 3 sequential streaming steps through ORT. + + This is a fast structural/shape smoke test only (random weights, no + HuggingFace comparison). For numeric parity against the real + HuggingFace checkpoint — including the lookahead trimming and the + FIFO-eviction / top-k AOSC-compression transitions — see + ``tests/e2e_golden_test.py::TestL5DiarizationSession``, which threads + the real exported streaming graph's own cache outputs back in as the + next chunk's inputs and checks each step against a real-weights + HuggingFace golden session. + """ + import os + import tempfile + + import onnxruntime as ort + + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationStreamingTask + + config = self._config() + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationStreamingTask()) + model = pkg["model"] + + for init in model.graph.initializers.values(): + if init.const_value is not None: + continue + shape = [d if isinstance(d, int) else 1 for d in init.shape] + arr = (np.random.randn(*shape) * 0.02).astype(np.float32) + init.const_value = ir.tensor(arr, name=init.name) + + with tempfile.TemporaryDirectory() as tmp: + path = os.path.join(tmp, "model.onnx") + ir.save(model, path, external_data="model.onnx.data") + sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) + + window = config.chunk_length + config.chunk_right_context + raw_window = window * config.subsampling_factor + cache_len = config.streaming_speaker_cache_length + fifo_len = config.streaming_fifo_length + + state = { + "past_cache_embeds": np.zeros((1, cache_len, config.hidden_size), np.float32), + "past_cache_probs": np.zeros((1, cache_len, config.num_speakers), np.float32), + "past_fifo": np.zeros((1, fifo_len, config.hidden_size), np.float32), + "past_num_cache_frames": np.array(0, dtype=np.int64), + "past_num_fifo_frames": np.array(0, dtype=np.int64), + "past_is_compressed": np.array(False, dtype=bool), + } + output_names = [out.name for out in sess.get_outputs()] + for step in range(3): + lookahead = config.chunk_right_context if step < 2 else 0 + # The final chunk of a recording may be shorter than the + # default full window (no lookahead, no padding contract) — + # the graph's time axis is symbolic to allow this. + step_raw_window = raw_window if step < 2 else raw_window // 2 + feats = np.random.randn(1, config.feat_in, step_raw_window).astype(np.float32) + inputs = { + "input_features": feats, + "num_lookahead_frames": np.array(lookahead, dtype=np.int64), + **state, + } + outs = dict(zip(output_names, sess.run(output_names, inputs))) + + speaker_probs = outs["speaker_probs"] + assert speaker_probs.min() >= 0.0 and speaker_probs.max() <= 1.0 + # Output frames exclude the look-ahead suffix (see the task's + # docstring): chunk_window_frames - lookahead * subsampling_factor. + expected_frames = step_raw_window - lookahead * config.subsampling_factor + assert speaker_probs.shape[1] == expected_frames + assert outs["present_num_cache_frames"] <= cache_len + assert outs["present_num_fifo_frames"] <= fifo_len + + state = { + "past_cache_embeds": outs["present_cache_embeds"], + "past_cache_probs": outs["present_cache_probs"], + "past_fifo": outs["present_fifo"], + "past_num_cache_frames": outs["present_num_cache_frames"], + "past_num_fifo_frames": outs["present_num_fifo_frames"], + "past_is_compressed": outs["present_is_compressed"], + } + + def test_non_final_chunk_with_misaligned_window_has_exact_frame_count(self): + """Regression test for titaiwangms's PR #748 review. + + Streaming output length must never desync from the number of frames + committed to the FIFO/cache, even for a non-final, + non-``subsampling_factor``-aligned window. + + Verifies, across two consecutive calls (an unaligned non-final chunk + followed by a second chunk consuming the returned cache state), that + ``speaker_probs``'s frame count always exactly matches + ``(present_num_fifo_frames - past_num_fifo_frames) * + subsampling_factor`` — the number of raw frames actually committed to + the cache this call. A caller advancing a sliding window using either + number therefore always stays in sync; neither frame double-counting + nor frame dropping is possible. + """ + import os + import tempfile + + import onnxruntime as ort + + from mobius.models.nemotron3_diarization import Nemotron3DiarizationModel + from mobius.tasks import DiarizationStreamingTask + + config = _base_config( + Nemotron3DiarizationConfig, + feat_in=16, + subsampling_factor=2, + head_hidden_size=8, + num_speakers=4, + partial_rotary_factor=1.0, + chunk_length=4, + chunk_right_context=2, + # Deliberately larger than either call's FIFO growth (5, then 4) + # so FIFO eviction/compaction -- a separately-tested mechanism -- + # never triggers and can't confound this test's invariant check. + streaming_fifo_length=20, + streaming_speaker_cache_length=20, + streaming_speaker_cache_update_period=20, + ) + module = Nemotron3DiarizationModel(config) + pkg = build_from_module(module, config, task=DiarizationStreamingTask()) + model = pkg["model"] + + for init in model.graph.initializers.values(): + if init.const_value is not None: + continue + shape = [d if isinstance(d, int) else 1 for d in init.shape] + arr = (np.random.randn(*shape) * 0.02).astype(np.float32) + init.const_value = ir.tensor(arr, name=init.name) + + with tempfile.TemporaryDirectory() as tmp: + path = os.path.join(tmp, "model.onnx") + ir.save(model, path, external_data="model.onnx.data") + sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) + + cache_len = config.streaming_speaker_cache_length + fifo_len = config.streaming_fifo_length + lookahead = config.chunk_right_context + factor = config.subsampling_factor + assert factor == 2 + # Deliberately misaligned: one raw frame more than a multiple of + # subsampling_factor (2), i.e. an odd raw window length. + misaligned_raw_window = ( + config.chunk_length + config.chunk_right_context + ) * factor + 1 + + output_names = [out.name for out in sess.get_outputs()] + + def run_chunk(raw_window: int, num_lookahead: int, state: dict): + feats = np.random.randn(1, config.feat_in, raw_window).astype(np.float32) + inputs = { + "input_features": feats, + "num_lookahead_frames": np.array(num_lookahead, dtype=np.int64), + **state, + } + outs = dict(zip(output_names, sess.run(output_names, inputs))) + next_state = { + "past_cache_embeds": outs["present_cache_embeds"], + "past_cache_probs": outs["present_cache_probs"], + "past_fifo": outs["present_fifo"], + "past_num_cache_frames": outs["present_num_cache_frames"], + "past_num_fifo_frames": outs["present_num_fifo_frames"], + "past_is_compressed": outs["present_is_compressed"], + } + return outs, next_state + + state = { + "past_cache_embeds": np.zeros((1, cache_len, config.hidden_size), np.float32), + "past_cache_probs": np.zeros((1, cache_len, config.num_speakers), np.float32), + "past_fifo": np.zeros((1, fifo_len, config.hidden_size), np.float32), + "past_num_cache_frames": np.array(0, dtype=np.int64), + "past_num_fifo_frames": np.array(0, dtype=np.int64), + "past_is_compressed": np.array(False, dtype=bool), + } + + # Call 1: a non-final (lookahead > 0), misaligned window. The + # correct emitted length matches HuggingFace's reference trim + # (`logits[:, :num_frames]`, no lookahead subtraction) rather + # than the previously buggy `raw_window - lookahead * factor`. + num_new_embeds_1 = -(-misaligned_raw_window // factor) # ceil division + num_chunk_embeds_1 = num_new_embeds_1 - lookahead + expected_frames_1 = min(misaligned_raw_window, num_chunk_embeds_1 * factor) + + outs1, state = run_chunk(misaligned_raw_window, lookahead, state) + assert outs1["speaker_probs"].shape[1] == expected_frames_1 + + fifo_growth_1 = int(state["past_num_fifo_frames"]) - 0 + assert fifo_growth_1 == num_chunk_embeds_1 + # Core invariant: output length exactly matches the frames + # committed to the cache this call (no desync). + assert outs1["speaker_probs"].shape[1] == fifo_growth_1 * factor + + # Call 2: a second, final (lookahead == 0) aligned window, + # consuming the state returned by call 1 — verifies the + # invariant continues to hold across calls, not just in + # isolation. + aligned_raw_window = config.chunk_length * factor + num_new_embeds_2 = aligned_raw_window // factor + num_chunk_embeds_2 = num_new_embeds_2 # lookahead == 0 + expected_frames_2 = min(aligned_raw_window, num_chunk_embeds_2 * factor) + + prev_fifo_frames = int(state["past_num_fifo_frames"]) + outs2, state = run_chunk(aligned_raw_window, 0, state) + assert outs2["speaker_probs"].shape[1] == expected_frames_2 + + fifo_growth_2 = int(state["past_num_fifo_frames"]) - prev_fifo_frames + assert fifo_growth_2 == num_chunk_embeds_2 + assert outs2["speaker_probs"].shape[1] == fifo_growth_2 * factor diff --git a/tests/e2e_golden_test.py b/tests/e2e_golden_test.py index 8e4b55a6c..6ba8dd589 100644 --- a/tests/e2e_golden_test.py +++ b/tests/e2e_golden_test.py @@ -51,13 +51,14 @@ generation_json_path_for_case, golden_path_for_case, has_golden, + load_diarization_golden, load_drafter_inputs, load_generation_golden, load_golden_ref, load_tolerances, ) from mobius._testing.ort_inference import OnnxModelSession -from mobius._testing.parity import ParityResult, compare_golden +from mobius._testing.parity import ParityResult, compare_diarization_golden, compare_golden @functools.cache @@ -569,6 +570,13 @@ def _build_model_package(case: GoldenTestCase) -> ModelPackage: assert route["preserve_quantization"] == source["preserve_quantization"] return package + if case.reference_loader == "nemo": + # NeMo-toolkit reference (e.g. Sortformer): built from a .nemo + # archive rather than a HuggingFace transformers checkpoint. + from mobius.integrations.nemo import build_from_nemo + + return build_from_nemo(case.model_id, revision=case.revision) + module_class = None task = None if case.architecture: @@ -580,6 +588,12 @@ def _build_model_package(case: GoldenTestCase) -> ModelPackage: reg = registry.get_registration(case.architecture) module_class = reg.module_class task = reg.task or getattr(module_class, "default_task", None) + elif case.task_type in {"diarization", "diarization-streaming"}: + # Diarization checkpoints default to the offline "diarization" task + # (registry default_task); the streaming variant must be requested + # explicitly since it produces a different ONNX graph contract + # (cache-state inputs/outputs) from the same checkpoint. + task = case.task_type return build( case.model_id, revision=case.revision, @@ -2181,6 +2195,167 @@ def _run_qwen35_mtp_prefill( session.close() +def _run_diarization_offline_prefill( + pkg: ModelPackage, + golden: dict[str, np.ndarray], +) -> np.ndarray: + """Run the offline diarization ONNX graph and return ``speaker_probs``. + + ``golden["mel"]`` is stored channel-first ``[batch, feat, frames]``, + matching the mobius ``DiarizationTask`` input contract directly (no + runtime transpose needed). + """ + session = _open_decoder_session(pkg) + try: + outputs = session.run({"input_features": golden["mel"].astype(np.float32)}) + finally: + session.close() + return outputs["speaker_probs"] + + +def _assert_diarization_offline_golden(case: GoldenTestCase, level: str) -> None: + """L4 (and offline L5) diarization check: golden vs. real ONNX build. + + Shared by both the offline ``diarization`` task type's L4 test and (for + L5-eligible cases such as sortformer, which has no streaming task + variant) the additional exact dominant-speaker / active-speaker-set + decision checks a downstream consumer would read. + """ + golden = load_diarization_golden(case) + if golden is None: + pytest.skip(f"Diarization golden missing for {case.case_id}") + tolerances = load_tolerances("L4", case.dtype) + pkg = _build_model_package(case) + probs = _run_diarization_offline_prefill(pkg, golden) + + report = compare_diarization_golden( + onnx_probs=probs, + golden_probs=golden["probs"], + dtype=case.dtype, + level=level, + ) + if report.top10_jaccard < tolerances.top10_jaccard_warn: + warnings.warn( + f"Low active-speaker Jaccard for {case.case_id}: " + f"{report.top10_jaccard:.2f} < {tolerances.top10_jaccard_warn}", + stacklevel=1, + ) + # compare_diarization_golden() has no AMBIGUOUS downgrade (see its + # docstring), so PASS is the only non-FAIL result -- require it + # explicitly rather than merely excluding FAIL. + assert report.result == ParityResult.PASS, report.message + + if level == "L5": + # Downstream diarization decisions: dominant speaker per frame + # (argmax) and the binarized active-speaker set (threshold 0.5) + # must reproduce the reference exactly over the whole utterance. + np.testing.assert_array_equal(probs.argmax(axis=-1), golden["probs"].argmax(axis=-1)) + np.testing.assert_array_equal(probs > 0.5, golden["probs"] > 0.5) + + +def _run_diarization_streaming_session( + case: GoldenTestCase, +) -> None: + """L5: thread a multi-chunk streaming diarization session end to end. + + Feeds each chunk's own AOSC + FIFO cache-state outputs back in as the + next chunk's cache inputs (self-consistent chaining through the real + exported streaming ONNX graph) and checks every chunk's speaker + probabilities and cache state against the independently-computed + HuggingFace reference. + """ + golden = load_diarization_golden(case) + if golden is None: + pytest.skip(f"Diarization golden missing for {case.case_id}") + pkg = _build_model_package(case) + config = pkg.config + assert config is not None + + num_chunks = int(golden["num_chunks"]) + num_speakers = int(golden["chunk0_probs"].shape[-1]) + cache_capacity = int(golden["chunk0_cache_embeds"].shape[1]) + fifo_capacity = int(golden["chunk0_fifo"].shape[1]) + hidden_size = config.hidden_size + + # Fresh-stream initial state: all-zero cache/FIFO buffers, no occupancy. + past_cache_embeds = np.zeros((1, cache_capacity, hidden_size), dtype=np.float32) + past_cache_probs = np.zeros((1, cache_capacity, num_speakers), dtype=np.float32) + past_fifo = np.zeros((1, fifo_capacity, hidden_size), dtype=np.float32) + past_num_cache_frames = np.array(0, dtype=np.int64) + past_num_fifo_frames = np.array(0, dtype=np.int64) + past_is_compressed = np.array(False, dtype=np.bool_) + + saw_compression = False + session = _open_decoder_session(pkg) + try: + for i in range(num_chunks): + mel_cf = golden[f"chunk{i}_mel"].astype(np.float32) + feed = { + "input_features": mel_cf, + "num_lookahead_frames": golden[f"chunk{i}_lookahead"], + "past_cache_embeds": past_cache_embeds, + "past_cache_probs": past_cache_probs, + "past_fifo": past_fifo, + "past_num_cache_frames": past_num_cache_frames, + "past_num_fifo_frames": past_num_fifo_frames, + "past_is_compressed": past_is_compressed, + } + out = session.run(feed) + + report = compare_diarization_golden( + onnx_probs=out["speaker_probs"], + golden_probs=golden[f"chunk{i}_probs"], + dtype=case.dtype, + level="L5", + ) + # No AMBIGUOUS downgrade for diarization -- require PASS. + assert report.result == ParityResult.PASS, f"chunk {i}: {report.message}" + + num_cache = int(out["present_num_cache_frames"]) + num_fifo = int(out["present_num_fifo_frames"]) + is_compressed = bool(out["present_is_compressed"]) + assert num_cache == int(golden[f"chunk{i}_num_cache_frames"]) + assert num_fifo == int(golden[f"chunk{i}_num_fifo_frames"]) + assert is_compressed == bool(golden[f"chunk{i}_is_compressed"]) + saw_compression = saw_compression or is_compressed + + np.testing.assert_allclose( + out["present_cache_embeds"][:, :num_cache], + golden[f"chunk{i}_cache_embeds"][:, :num_cache], + atol=1e-3, + err_msg=f"chunk {i} cache_embeds", + ) + np.testing.assert_allclose( + out["present_cache_probs"][:, :num_cache], + golden[f"chunk{i}_cache_probs"][:, :num_cache], + atol=1e-4, + err_msg=f"chunk {i} cache_probs", + ) + np.testing.assert_allclose( + out["present_fifo"][:, :num_fifo], + golden[f"chunk{i}_fifo"][:, :num_fifo], + atol=1e-3, + err_msg=f"chunk {i} fifo", + ) + + # Feed this step's own outputs back in as the next call's cache + # state (self-consistent chaining through the real graph). + past_cache_embeds = out["present_cache_embeds"] + past_cache_probs = out["present_cache_probs"] + past_fifo = out["present_fifo"] + past_num_cache_frames = out["present_num_cache_frames"] + past_num_fifo_frames = out["present_num_fifo_frames"] + past_is_compressed = out["present_is_compressed"] + finally: + session.close() + + # The golden's chosen chunk/cache sizes are expected to force at least + # one AOSC compression event; fail loudly if a future config/golden + # change stops exercising that path so this test doesn't silently lose + # coverage. + assert saw_compression, "expected the streaming session to trigger AOSC compression" + + # --------------------------------------------------------------------------- # L4 Tests: Checkpoint Verified # --------------------------------------------------------------------------- @@ -2202,6 +2377,13 @@ def test_prefill_argmax_matches_golden(self, case: GoldenTestCase) -> None: "Continuous-token TTS uses tests/vibevoice_golden_test.py " "(control logits, latent frames, and waveform semantics)." ) + if case.task_type == "diarization": + # Diarization emits continuous per-frame per-speaker + # probabilities, not a vocab logit vector, so it is compared + # with compare_diarization_golden() instead of the + # argmax/top-K-gated compare_golden() used below. + _assert_diarization_offline_golden(case, level="L4") + return golden_path = golden_path_for_case(case) golden = load_golden_ref(golden_path) if golden is None: @@ -3148,3 +3330,41 @@ def test_generation_matches_golden(self, case: GoldenTestCase) -> None: f" Got ({actual_len} tokens): " f"{new_tokens.tolist()}" ) + + +# --------------------------------------------------------------------------- +# L5 Tests: Diarization Session +# --------------------------------------------------------------------------- + +# Diarization has no autoregressive token loop (OnnxGenerator does not apply): +# offline cases replay a single forward pass and check the downstream +# per-frame decisions exactly; streaming cases thread a multi-chunk AOSC + +# FIFO cache-state session through the real exported graph. Both share the +# continuous-probability golden format loaded via load_diarization_golden() +# and the compare_diarization_golden() comparator (see _assert_diarization_ +# offline_golden / _run_diarization_streaming_session above). +_DIARIZATION_L5_TASKS = frozenset({"diarization", "diarization-streaming"}) + + +@pytest.mark.generation +@pytest.mark.integration +class TestL5DiarizationSession: + """L5: diarization decision / streaming-session parity. + + Gate: compare_diarization_golden()'s allclose gate, plus (offline) + exact dominant-speaker and active-speaker-set decision equality, or + (streaming) exact per-chunk cache-state equality across the whole + multi-chunk session. + """ + + @pytest.mark.parametrize("case", _L5_CASES) + def test_diarization_session_matches_golden(self, case: GoldenTestCase) -> None: + if case.task_type not in _DIARIZATION_L5_TASKS: + pytest.skip( + f"Not a diarization task_type ({case.task_type!r}); " + "covered by TestL5GenerationE2E instead." + ) + if case.task_type == "diarization-streaming": + _run_diarization_streaming_session(case) + else: + _assert_diarization_offline_golden(case, level="L5") diff --git a/tests/model_coverage_test.py b/tests/model_coverage_test.py index 7167e670e..f05c691dc 100644 --- a/tests/model_coverage_test.py +++ b/tests/model_coverage_test.py @@ -319,9 +319,14 @@ def _all_registered_with_test_id() -> dict[str, str]: "whisper": "Speech-to-text — requires audio inputs", "mms": "CTC ASR model — tested via TestBuildMMSGraph", "fastconformer_rnnt": "NeMo .nemo RNN-T ASR — tested via tests/nemo_rnnt_integration_test.py", - "sortformer": "NeMo .nemo speaker diarization — tested via tests/sortformer_integration_test.py", + "sortformer": "NeMo .nemo checkpoint has no HuggingFace config.json, so it cannot use " + "the generic L1-L3 config-based build test or L2 test_model_id validation (same " + "limitation as fastconformer_rnnt). L4/L5 golden coverage IS provided generically via " + "testdata/cases/diarization/sortformer.yaml + scripts/generate_golden.py + " + "tests/e2e_golden_test.py (see load_diarization_golden/compare_diarization_golden) — " + "this skip covers L1-L3 only, not L4/L5.", "VibeVoiceForASRStreamingTraining": "Streaming ASR has host-owned dual-convolution " - "state, arbitrary-mask decoder, hotword, and speaker-attribution orchestration that " + "state, arbitrary-mask decoder, hotword, and speaker-attribution orchestration that " "the generic L4/L5 runner cannot drive. Pinned L1-L3 graph/config/source-parity and " "complete checkpoint-index routing are covered for the 1.5B and 7B checkpoints; " "real-weight goldens require a dedicated GPU workflow.", diff --git a/tests/sortformer_integration_test.py b/tests/sortformer_integration_test.py deleted file mode 100644 index fc0e243b5..000000000 --- a/tests/sortformer_integration_test.py +++ /dev/null @@ -1,103 +0,0 @@ -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT License. - -"""L4/L5 golden parity tests for the NeMo Sortformer speaker-diarization model. - -Builds the ONNX diarization graph from the real ``.nemo`` archive -(``nvidia/diar_streaming_sortformer_4spk-v2.1``, resolved from HuggingFace Hub) -via the generic :func:`build_from_nemo` pipeline and compares its output against -a pre-computed NeMo reference -(``testdata/golden/speech/sortformer_diarization.npz``). - -The reference was generated with ``nemo_toolkit`` — regenerate it with -``scripts/generate_sortformer_golden.py`` if the architecture changes. - -- **L4** (``test_sortformer_diarization_parity``): the ONNX ``speaker_probs`` - match the NeMo offline reference speaker sigmoids element-wise. -- **L5** (``test_sortformer_end_to_end_diarization``): the full offline - diarization decision (per-frame active-speaker assignment derived from the - sigmoids) reproduces the NeMo reference over the whole utterance. -""" - -from __future__ import annotations - -import json -import os -import tempfile - -import numpy as np -import onnx_ir as ir -import onnxruntime as ort -import pytest - -pytestmark = pytest.mark.integration - -_MODEL_ID = "nvidia/diar_streaming_sortformer_4spk-v2.1" -# Pin the HF revision so the test always validates against the exact model the -# committed golden was generated from (see scripts/generate_sortformer_golden.py). -_REVISION = "fafaab5faa1617a0ca52d38dd3dc4bd636800d3d" -_GOLDEN = os.path.join( - os.path.dirname(__file__), - "..", - "testdata", - "golden", - "speech", - "sortformer_diarization.npz", -) - - -def _run_diarizer(golden: np.lib.npyio.NpzFile) -> np.ndarray: - """Build the ONNX diarizer via build_from_nemo and run the golden mel input.""" - from mobius.integrations.nemo import build_from_nemo - - pkg = build_from_nemo(_MODEL_ID, revision=_REVISION) - assert "model" in pkg - - mel = golden["mel"].astype(np.float32) # (1, feat_in, T) - with tempfile.TemporaryDirectory() as td: - path = os.path.join(td, "model.onnx") - ir.save(pkg["model"], path, external_data="model.onnx.data") - sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) - out = sess.run(None, {sess.get_inputs()[0].name: mel})[0] - return out - - -@pytest.mark.integration_slow -def test_sortformer_diarization_parity(): - """L4: ONNX speaker probabilities match the NeMo offline reference.""" - golden = np.load(_GOLDEN) - preds_ref = golden["preds"] # (1, T', num_spks) sigmoid activations - - out = _run_diarizer(golden) - - assert out.shape == preds_ref.shape - # Offline fp32 forward path: tight parity (matches build_from_nemo golden). - np.testing.assert_allclose(out, preds_ref, atol=1e-4) - - -@pytest.mark.integration_slow -def test_sortformer_end_to_end_diarization(): - """L5: the end-to-end offline diarization decision matches NeMo. - - The diarization output is a per-frame speaker-activity sigmoid in ``[0, 1]``. - The task-level decisions a downstream consumer reads are (a) the dominant - speaker per frame (``argmax``) and (b) the binarized set of active speakers - per frame (sigmoid threshold at 0.5). Verify the ONNX pipeline reproduces - both over the whole utterance and that the raw probabilities are well-formed. - """ - golden = np.load(_GOLDEN) - preds_ref = golden["preds"] - meta = json.loads(str(golden["meta"])) - num_spks = int(meta["num_spks"]) - - out = _run_diarizer(golden) - - # Raw sigmoids must be valid probabilities of the expected speaker count. - assert out.shape[-1] == num_spks - assert out.min() >= 0.0 and out.max() <= 1.0 - - # Diarization decision 1: per-frame dominant-speaker assignment sequence. - np.testing.assert_array_equal(out.argmax(axis=-1), preds_ref.argmax(axis=-1)) - - # Diarization decision 2: binarized per-frame active-speaker set. - np.testing.assert_array_equal(out > 0.5, preds_ref > 0.5)