Adding Nemotron 3 Diarization ONNX Model Support - #748
themason2011 wants to merge 6 commits into
Conversation
Performance Comparison
|
🏗️ Architecture Diff
No architecture changes detected. ✅ Legend: ⚪ No change · 🔵 Minor (attrs/inits) · 🟡 Moderate (nodes added/removed) · 🔴 Major (interface changed) |
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Critical and moderate issues remain in model compatibility, streaming behavior, configuration extraction, test coverage, and registry coverage.
Review effort: Lite
Findings: 4
Open (5)
What changed in this PR
Adds Nemotron 3 Diarization ONNX support, including configuration, model registration, offline/stateful streaming graphs, cache handling, and tests.
Changes:
- Adds Nemotron model architecture and configuration extraction.
- Adds streaming task wiring with cache/FIFO handling.
- Adds registry/exports and graph execution tests.
| File | Summary / findings |
|---|---|
tests/build_graph/speech_test.py |
Adds graph and runtime tests; needs HF numerical parity coverage for normal, final-chunk, and compressed-cache steps. |
tests/_test_configs.py |
Adds a tiny Nemotron test configuration. |
src/mobius/tasks/_diarization_streaming.py |
Adds streaming graph I/O; fixed dimensions and lookahead handling do not support shorter final chunks or processor streaming modes. |
src/mobius/tasks/__init__.py |
Exports the streaming task. |
src/mobius/models/nemotron3_diarization.py |
Implements the model and cache logic; contains MHA/KV shape incompatibility and dtype-invalid padding. |
src/mobius/models/__init__.py |
Exports the Nemotron model. |
src/mobius/_registry.py |
Registers the model; missing coverage YAML or documented skip will fail the coverage suite. |
src/mobius/_configs/_base.py |
Adds configuration extraction; nested configs must be normalized before field access. |
src/mobius/_configs/__init__.py |
Exports the configuration. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
…ngNemotron3DiarizationModelSupport
Nemotron3DiarizationModel's diarization/diarization-streaming task strings are not part of the generic testdata/cases/ + generate_golden.py YAML pipeline (same limitation as the NeMo-based sortformer model), so tests/model_coverage_test.py's L4/L5 check was failing. Add: - scripts/generate_nemotron3_diarization_golden.py: regenerates testdata/golden/speech/nemotron3_diarization.npz from the real HuggingFace Nemotron3DiarizationForAudioFrameClassification reference (offline single-pass output, plus a 3-chunk streaming session that exercises FIFO overflow and AOSC top-k compression). - tests/nemotron3_diarization_integration_test.py: L4 offline parity and L5 multi-chunk streaming-session parity (speaker probabilities and full cache state) against that golden, building the real ONNX graphs via mobius.build(). - _COVERAGE_SKIP entry in tests/model_coverage_test.py pointing at the new integration test, mirroring the existing sortformer entry. Signed-off-by: themason2011 <masoncorey@microsoft.com>
Previously, nvidia/Nemotron-3-Diarization and sortformer had bespoke, one-off integration test scripts and golden generators that duplicated the shared L4/L5 golden-test framework (testdata/cases/, scripts/generate_golden.py, tests/e2e_golden_test.py, mobius._testing.golden/parity). This meant diarization models were not subject to the same coverage checks, schema validation, or discovery machinery as every other architecture, and required a model_coverage_test.py _COVERAGE_SKIP escape hatch to pass CI. This change extends the existing framework to natively support diarization's per-frame speaker-probability output, instead of building a separate one-off harness: - Add "diarization" / "diarization-streaming" task types and a "nemo" reference_loader to testdata/cases/schema.json. - Add save_diarization_golden()/load_diarization_golden() to mobius._testing.golden, storing full-precision probability arrays in a companion .npz (mirroring the existing drafter-inputs pattern) with lightweight metadata in the standard golden .json. - Add compare_diarization_golden() to mobius._testing.parity: an elementwise allclose gate over continuous frame probabilities, with dominant-speaker argmax agreement and active-speaker-set Jaccard as diagnostics, returning the same ParityReport shape used everywhere else in the framework. - Add _generate_diarization_offline()/_generate_diarization_streaming() generators to scripts/generate_golden.py, dispatching to either AutoModelForAudioFrameClassification (HF) or a NeMo .nemo checkpoint depending on case.model_type. - Wire diarization/diarization-streaming dispatch into tests/e2e_golden_test.py: a nemo-loader branch in _build_model_package(), new helpers for offline and streaming diarization runs, an early-exit branch in TestL4CheckpointVerified, and a new TestL5DiarizationSession class for full cache-threaded streaming verification. - Add 3 new YAML cases under testdata/cases/diarization/ covering nemotron3 offline (L4), nemotron3 streaming (L5), and sortformer offline (L4+L5, migrated from NeMo). - Migrate existing committed golden data for both models into the new storage format; nemotron3's HF-sourced mel (channel-last) is pre-transposed to the channel-first layout the ONNX graph expects, matching NeMo's native convention used by sortformer. - Remove the now-redundant bespoke test/generator scripts and their ad-hoc golden .npz files. - Remove the nemotron3_diarization _COVERAGE_SKIP entry entirely (now covered generically). Narrow the sortformer _COVERAGE_SKIP entry to only its pre-existing, unrelated L1-L3 gap (no test_model_id/tiny config, matching fastconformer_rnnt), since its L4/L5 coverage is now handled by this framework too. Verified: tests/yaml_schema_test.py, tests/model_coverage_test.py, and tests/e2e_golden_test.py -k diarization all pass; full baseline suite unaffected. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: themason2011 <masoncorey@microsoft.com>
Fixes 3 of the 5 outstanding review-bot comments on this PR (the other two -- missing L2/L4 coverage, and insufficient streaming test coverage -- were already resolved by the prior 'Integrate diarization L4/L5 testing' commit). - tests/_test_configs.py: set num_key_value_heads=TINY_HEADS on the nemotron3_diarization tiny config. HuggingFace's checkpoint is always MHA (audio_config sets num_key_value_heads == num_attention_heads); the shared tiny-config default is GQA (TINY_KV_HEADS < TINY_HEADS), which built a V-projection shape no real checkpoint could populate. - src/mobius/models/nemotron3_diarization.py: cast every FLOAT32 literal in the AOSC scoring/compression path (_get_frame_scores's threshold/complement/log-half/neg-inf/positivity constants, the tail-boost Where branches, the score-cache Pad fill value, and the top-k -inf sentinel comparison) to match the surrounding tensor's config.dtype via CastLike. ONNX elementwise ops require matching operand dtypes, so these bare FLOAT32 constants made fp16/bf16 streaming exports invalid. - src/mobius/tasks/_diarization_streaming.py: make the streaming input_features time axis symbolic instead of a fixed chunk_length+chunk_right_context window. forward_streaming already derives every frame count from the input's actual shape, so only the declared graph input shape was blocking HuggingFace's supported shorter final chunk. - tests/build_graph/speech_test.py: exercise a shorter final chunk (no lookahead) in the streaming execution test to prove the symbolic time axis actually accepts variable-length windows, and add a docstring note pointing to tests/e2e_golden_test.py:: TestL5DiarizationSession for real-weights numeric parity (including the lookahead trim and cache-compression transitions), since this test intentionally stays a fast random-weight structural smoke test. Verified: tests/build_graph/speech_test.py -k Nemotron3Diarization, tests/model_coverage_test.py, and tests/e2e_golden_test.py -k nemotron3_diarization -m integration (against the real nvidia/Nemotron-3-Diarization checkpoint) all pass. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: themason2011 <masoncorey@microsoft.com>
|
Full review of head Findings
Smaller follow-ups
Verification boundary: This was a static review, not an independent real-weight run. L1, L3, and lint passed on this head, but the L4 and L5 CI jobs each hit their one-hour timeout and were cancelled; that is not evidence that the model-specific goldens failed or passed. A targeted completed L4/L5 run is still needed. Short, non-subsampling-aligned final chunks and reduced-precision cache compression have not been independently compared with the HF reference here; these are open verification questions, not additional confirmed defects. |
Fixes the 6 findings confirmed with the user from titaiwangms's latest review pass: - src/mobius/_testing/parity.py, parity_test.py, tests/e2e_golden_test.py: remove compare_diarization_golden()'s AMBIGUOUS downgrade. Dominant-speaker argmax agreement is not a valid substitute for full elementwise tolerance on a multi-label sigmoid output -- a secondary (non-dominant) speaker's probability can cross the 0.5 activation threshold (a real diarization error) while the argmax stays unchanged. Every allclose failure is now a hard FAIL; dominant-speaker match and active-speaker-set Jaccard are reported as diagnostics only. Updated L4/L5 assertions to require ParityResult.PASS. - src/mobius/tasks/_diarization.py: made DiarizationTask's docstring model-agnostic -- "frames = time / subsampling_factor" is true for Sortformer but wrong for Nemotron3 (which upsamples back to per-mel-frame resolution, frames == time). - src/mobius/models/nemotron3_diarization.py (forward_streaming): fixed a real pad-leak bug -- the raw-frame trim bound included the lookahead window's frame count instead of just the current chunk's, so a shorter final chunk without lookahead retained trailing padding in its output. Added a regression test in tests/build_graph/speech_test.py. - testdata/golden/diarization/sortformer*: regenerated via real NeMo inference so field names match the current schema. - scripts/generate_golden.py: pinned an exact transformers commit (c8b81b63232be35ab1774dd3cabbf499d8b9808f) for reproducible Nemotron3Diarization golden generation (the model isn't in any released transformers version yet); added _best_effort_package_commit() to record the resolved git commit in golden provenance for git-installed dependencies (transformers, nemo_toolkit). Regenerated the Nemotron3 golden JSON/NPZ files against the pinned commit. - src/mobius/models/nemotron3_diarization.py (forward, the offline/ non-streaming graph -- the largest finding): replaced the single flat non-causal pass with an ONNX Loop-based chunked implementation matching HuggingFace's real offline forward, which embeds the whole input once then loops in fixed chunk_length-sized chunks (with chunk_right_context lookahead), reusing the AOSC + FIFO cache mechanism with offline-specific cache sizes (config.fifo_length, config.speaker_cache_update_period -- distinct from streaming_config's sizes). Extracted a shared _run_chunk_and_update_cache() helper out of forward_streaming so both callers share the same encode/classify/ cache-update logic parameterized by cache size. The loop accumulates logits into a fixed-size buffer (zero-padded per-iteration regions added in, rather than growing via Concat) to avoid an optimizer silently folding Concat(compile-time-empty, x) -> x and discarding all but the last iteration. Added _prerealize_parameters() to register every backbone/classifier parameter as a root-graph initializer with its final qualified name before the Loop body is built, and pushed distinguishing module scopes around each Loop/If subgraph -- both fixes for an ONNX SSA violation caused by nested subgraphs' independent per-graph node counters coincidentally generating identical auto-generated names. - src/mobius/_configs/_base.py: added offline_fifo_length (default 40) and offline_speaker_cache_update_period (default 300) to Nemotron3DiarizationConfig, extracted from HuggingFace's top-level config.fifo_length/config.speaker_cache_update_period. - src/mobius/_passes/_fold_transpose.py: fixed FoldTransposedInitializerPass removing a Transpose node via model.graph.remove(...) instead of node.graph.remove(...) -- this crashed whenever the node lived inside a subgraph (e.g. the new offline diarization Loop body), since ir.Graph.all_nodes() recurses into subgraphs but a node only belongs to its own (sub)graph. A general-purpose optimizer bug, not diarization-specific. - tests/build_graph/speech_test.py: added TestBuildGraphNemotron3DiarizationOffline (package_builds, Loop presence, single/multi-chunk shape and cache-compression coverage) since no offline-forward ORT-execution tests existed before. - testdata/cases/diarization/nemotron3_diarization_multichunk.yaml (new) + golden data: a 6000-mel-frame recording (750 encoder embeds, well above chunk_length=340) that forces HuggingFace's real offline forward to internally re-chunk into 3 iterations and exercise cache compression -- verifying full numeric parity of the new Loop-based offline graph against real HuggingFace chunked inference, not just structural ORT execution. Verified: tests/build_graph/speech_test.py (214 passed), tests/e2e_golden_test.py -k "diarization or sortformer" -m integration (25 passed, against real checkpoints), and the broader tests/build_graph + src/ suite (no new regressions; two pre-existing gemma4/gemma3n failures confirmed unrelated via git stash). Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: themason2011 <masoncorey@microsoft.com>
|
Thanks for the thorough re-review — addressed all findings in b574814. Summary of changes and verification for each: 1. Major — L4/L5 can pass with incorrect diarization probabilities. 2. Major — the default offline graph silently changes semantics for long recordings. Smaller follow-ups:
Verification for this pass: Full details in the commit message: b574814. |


Summary
Add ONNX graph support for NVIDIA Nemotron 3 Diarization, including both an offline full-sequence model and a stateful streaming model for per-frame speaker-activity probabilities.
Changes
Nemotron3DiarizationConfigto map the Hugging Face audio, head, and streaming settings, and registernemotron3_diarizationwith the existing offlinediarizationtask.diarization-streamingtask. Each graph invocation processes one mel-feature chunk with configurable look-ahead and returns that chunk’s speaker probabilities plus updated fixed-capacity Arrival-Order Speaker Cache (AOSC), FIFO buffers, occupancy counts, and compression state. The cache update includes score-based top-k compression.Validation
Scope / limitations
attention_maskfor a potentially padded final chunk.