Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
225 changes: 225 additions & 0 deletions scripts/generate_golden.py
Original file line number Diff line number Diff line change
Expand Up @@ -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://...@<ref>``) -- 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."""
Expand Down Expand Up @@ -1907,6 +1931,205 @@ 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="cpu")
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)
mel_len = torch.tensor([num_frames], 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)

arrays = {
"mel": mel.numpy().astype(np.float32),
"emb_seq": emb_seq.numpy().astype(np.float32),
"emb_len": emb_len.numpy().astype(np.int64),
"probs": preds.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.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)
with torch.no_grad():
out = model(input_features=mel)

arrays = {
"mel": np.transpose(mel.numpy(), (0, 2, 1)).astype(np.float32),
"probs": out.logits.sigmoid().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.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)
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.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().numpy().copy()
arrays[f"chunk{i}_cache_embeds"] = cache.embeds.numpy().copy()
arrays[f"chunk{i}_cache_probs"] = cache.probs.numpy().copy()
arrays[f"chunk{i}_fifo"] = cache.fifo.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.
Expand All @@ -1930,6 +2153,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,
}


Expand Down
106 changes: 0 additions & 106 deletions scripts/generate_sortformer_golden.py

This file was deleted.

2 changes: 2 additions & 0 deletions src/mobius/_configs/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@
MoonshineStreamingConfig,
MuseGlimmerConfig,
NanoChatConfig,
Nemotron3DiarizationConfig,
NemotronHConfig,
NemotronParseConfig,
ParakeetCTCConfig,
Expand Down Expand Up @@ -167,6 +168,7 @@
"MoonshineStreamingConfig",
"MuseGlimmerConfig",
"NanoChatConfig",
"Nemotron3DiarizationConfig",
"NemotronParseConfig",
"NemotronHConfig",
"ParakeetCTCConfig",
Expand Down
Loading
Loading