From d7b232c48d196ad4803e2cb66e18a31e5f22525f Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 12:27:42 +0200 Subject: [PATCH] Add Nemotron-3 diarization export and inference sample Signed-off-by: contentis --- CMakeLists.txt | 1 + THIRD_PARTY_NOTICES.md | 2 + asr/diarization/CMakeLists.txt | 6 + asr/diarization/README.md | 39 ++++ asr/diarization/diarization.cpp | 179 ++++++++++++++++++ asr/diarization/diarization.h | 51 +++++ asr/diarization/main.cpp | 118 ++++++++++++ .../model_export/export_diarization.py | 69 +++++++ asr/diarization/model_export/mel.py | 31 +++ asr/diarization/model_export/model.py | 120 ++++++++++++ .../model_export/validate_diarization.py | 138 ++++++++++++++ 11 files changed, 754 insertions(+) create mode 100644 asr/diarization/CMakeLists.txt create mode 100644 asr/diarization/README.md create mode 100644 asr/diarization/diarization.cpp create mode 100644 asr/diarization/diarization.h create mode 100644 asr/diarization/main.cpp create mode 100644 asr/diarization/model_export/export_diarization.py create mode 100644 asr/diarization/model_export/mel.py create mode 100644 asr/diarization/model_export/model.py create mode 100644 asr/diarization/model_export/validate_diarization.py diff --git a/CMakeLists.txt b/CMakeLists.txt index 9392711..5ce9afe 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -101,6 +101,7 @@ add_subdirectory(asr/rnnt) add_subdirectory(vision/sam2) add_subdirectory(asr/whisper) add_subdirectory(asr/qwen3) +add_subdirectory(asr/diarization) if(DIN_BUILD_TESTING) set(_din_cli_test_audio "${DIN_DEPLOY_ARTIFACTS_ROOT}/audio/de/schakaleundaraber_elli.mp3") diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index a721e4e..eeb4fb4 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -80,6 +80,8 @@ For components supplied through an SDK or binary package, the corresponding vend ## Model/Artifact +- `nvidia/Nemotron-3-Diarization`: [OpenMDW 1.1](https://openmdw.ai/license/1-1/); model artifacts are not included. + - `black-forest-labs/FLUX.2-klein-4b`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-fp8`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-nvfp4`: Apache-2.0 diff --git a/asr/diarization/CMakeLists.txt b/asr/diarization/CMakeLists.txt new file mode 100644 index 0000000..5846780 --- /dev/null +++ b/asr/diarization/CMakeLists.txt @@ -0,0 +1,6 @@ +add_din_shared_library(din_diarization STATIC diarization.cpp) +target_include_directories(din_diarization PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}") +target_link_libraries(din_diarization PUBLIC din_common_ort din_common_io nlohmann_json::nlohmann_json) + +add_din_executable(din_diarization_cli main.cpp) +target_link_libraries(din_diarization_cli PRIVATE din_diarization din_asr_qwen3 argparse) diff --git a/asr/diarization/README.md b/asr/diarization/README.md new file mode 100644 index 0000000..1c72043 --- /dev/null +++ b/asr/diarization/README.md @@ -0,0 +1,39 @@ +# Nemotron diarization + +[`nvidia/Nemotron-3-Diarization`](https://huggingface.co/nvidia/Nemotron-3-Diarization): up to eight speakers, including overlaps, with 10 ms predictions. + +| Feature | Model / upstream toolkit | C++ sample | +| --- | :---: | :---: | +| Offline / long form | ✅ | ✅ | +| Live streaming | ✅ | — | +| Batched recordings | ✅ | — | +| Standalone diarization | ✅ | ✅ | +| ASR + forced alignment + speaker labels | ✅ | ✅ | +| CPU (FP32) / GPU (TensorRT RTX) | | ✅ | + +## Export and validate + +Tested with NeMo Speech commit `f613eed86ed4696db0891aac4e9104337a39142c`. + +```bash +pip install "nemo_toolkit[asr] @ git+https://github.com/NVIDIA-NeMo/Speech.git@f613eed86ed4696db0891aac4e9104337a39142c" +python asr/diarization/model_export/export_diarization.py --output artifacts/nemotron-diarization/onnx-bf16 --dtype bf16 +python asr/diarization/model_export/validate_diarization.py --checkpoint model.nemo --model-dir artifacts/nemotron-diarization/onnx-bf16 --provider trt-rtx --audio audio-16khz.wav +``` + +HF downloads the public checkpoint automatically. `--checkpoint` accepts an existing `.nemo` file. Export defaults to the checkpoint's BF16 precision; use `--dtype fp32` for CPU or `--dtype fp16` for FP16. Validation runs the NeMo reference on CUDA. + +## Run + +```bash +cmake --build build --target din_diarization_cli +din_diarization_cli audio.wav --model-dir artifacts/nemotron-diarization/onnx-bf16 +din_diarization_cli audio.wav --model-dir artifacts/nemotron-diarization/onnx-fp32 --provider cpu +din_diarization_cli audio.wav --asr-dir artifacts/qwen3/onnx-bf16 --aligner-dir artifacts/qwen3/aligner-onnx-bf16 +``` + +The combined command runs diarization alongside Qwen3 ASR/alignment. Each component has its own provider selection. Other ASR models can use the independent `Pipeline` and `Result::SpeakerAt` with their word timestamps. Overlapping speech is retained in speaker intervals; ASR is not source separation. + +One fixed-shape ONNX contains FP32 log-mel, model weights and speaker-cache updates. State stays on the selected device between chunks; recording length does not change the engine shape. FP32 validation checks probabilities strictly; reduced precision checks speaker/activity agreement against NeMo. + +Model license: [OpenMDW 1.1](https://openmdw.ai/license/1-1/). Model artifacts are downloaded separately. diff --git a/asr/diarization/diarization.cpp b/asr/diarization/diarization.cpp new file mode 100644 index 0000000..837ba38 --- /dev/null +++ b/asr/diarization/diarization.cpp @@ -0,0 +1,179 @@ +// SPDX-License-Identifier: Apache-2.0 +#include "diarization.h" + +#include +#include +#include +#include +#include +#include +#include + +#include "nvtx_helper.h" +#include "ort_session.h" +#include + +namespace din::asr::diarization +{ +using din::common::OrtRunner; +using din::common::TensorBuffer; + +struct Pipeline::Impl +{ + static constexpr int64_t kSamples = 3039 * 160 + 512 + 1; + Config config; + Ort::Env env{ORT_LOGGING_LEVEL_WARNING, "diarization"}; + std::unique_ptr model; + std::unique_ptr> history, history_preds, next_history, next_history_preds, probabilities, + samples; + std::unique_ptr> feature_length, history_length, sample_bounds; + Ort::IoBinding binding{nullptr}; + Ort::RunOptions options; + + explicit Impl(Config c) + : config(std::move(c)) + { + if (!std::isfinite(config.threshold) || config.threshold <= 0 || config.threshold >= 1) + throw std::invalid_argument("Diarization threshold must be between 0 and 1"); + std::ifstream file(config.model_dir / "metadata.json"); + if (!file) + throw std::runtime_error("Cannot read diarization metadata"); + const auto metadata = nlohmann::json::parse(file); + if (metadata.at("format_version") != 3 || metadata.at("chunk_frames") != 340 || + metadata.at("cache_frames") != 264 || metadata.at("speakers") != 8) + throw std::runtime_error("Incompatible diarization export; re-export the model"); + if (config.provider == "cpu" && metadata.at("dtype") != "fp32") + throw std::runtime_error("Use an FP32 diarization export with the CPU provider"); + const din::common::EpContextOptions context{config.ep_context_dir.string()}; + if (config.provider == "trt-rtx") + din::common::RegisterTensorRTRTXProvider(env); + din::common::ModelProfile profile; + profile.cache_subpath = "diarization_" + metadata.at("graph_hash").get(); + model = std::make_unique(env, (config.model_dir / "diarization.onnx").string(), config.provider, + config.ep_cache_dir.string(), context, profile); + const bool device = model->HasDeviceIo(); + history = std::make_unique>(*model, std::vector{1, 304, 512}, device); + next_history = std::make_unique>(*model, std::vector{1, 304, 512}, device); + history_preds = std::make_unique>(*model, std::vector{1, 264, 8}, device); + next_history_preds = std::make_unique>(*model, std::vector{1, 264, 8}, device); + probabilities = std::make_unique>(*model, std::vector{1, 2720, 8}, device); + feature_length = std::make_unique>(*model, std::vector{1}, device); + history_length = std::make_unique>(*model, std::vector{1}, device); + samples = std::make_unique>(*model, std::vector{1, kSamples}, device); + sample_bounds = std::make_unique>(*model, std::vector{2}, device); + binding = Ort::IoBinding(model->session); + binding.BindInput("samples", samples->BindingValue()); + binding.BindInput("sample_bounds", sample_bounds->BindingValue()); + binding.BindInput("feature_length", feature_length->BindingValue()); + binding.BindInput("history", history->BindingValue()); + binding.BindInput("history_preds", history_preds->BindingValue()); + binding.BindInput("history_length", history_length->BindingValue()); + binding.BindOutput("probabilities", probabilities->BindingValue()); + binding.BindOutput("next_history", next_history->BindingValue()); + binding.BindOutput("next_history_preds", next_history_preds->BindingValue()); + if (device) + options.AddConfigEntry("disable_synchronize_execution_providers", "1"); + } + + Result Run(const din::io::Audio& audio) + { + din::common::nvtx_scoped_range range{"diarization"}; + if (audio.sample_rate != 16000 || audio.samples.empty()) + throw std::invalid_argument("Expected nonempty 16 kHz mono audio"); + const auto start = std::chrono::steady_clock::now(); + Result result; + result.audio_seconds = static_cast(audio.Duration()); + const auto frames = audio.samples.size() / 160; + if (!frames) + return result; + history->Fill(0); + history_preds->Fill(0); + history->CopyAsyncToDevice(); + history_preds->CopyAsyncToDevice(); + std::array active; + active.fill(-1); + for (size_t offset = 0; offset < frames; offset += 2720) + { + din::common::nvtx_scoped_range chunk_range{"diarization.chunk"}; + const auto count = std::min(size_t{3040}, frames - offset); + const int64_t begin = static_cast(offset) * 160 - 257; + const int64_t lo = std::max(int64_t{0}, begin); + const int64_t hi = std::min(static_cast(audio.samples.size()), begin + kSamples); + samples->Fill(0); + std::copy_n(audio.samples.data() + lo, hi - lo, samples->HostData() + lo - begin); + sample_bounds->HostData()[0] = lo - begin; + sample_bounds->HostData()[1] = hi - begin; + *feature_length->HostData() = static_cast(count); + *history_length->HostData() = offset ? 304 : 0; + samples->CopyAsyncToDevice(); + sample_bounds->CopyAsyncToDevice(); + feature_length->CopyAsyncToDevice(); + history_length->CopyAsyncToDevice(); + model->session.Run(options, binding); + history->CopyFrom(*next_history); + history_preds->CopyFrom(*next_history_preds); + probabilities->CopyAsyncToHostWithNotification().Sync(); + const auto count_output = std::min(size_t{2720}, frames - offset); + for (size_t frame = 0; frame < count_output; ++frame) + { + const float time = std::min(result.audio_seconds, static_cast(offset + frame) * 0.01f); + for (int speaker = 0; speaker < 8; ++speaker) + { + const float probability = probabilities->HostData()[frame * 8 + speaker]; + if (probability > config.threshold && active[speaker] < 0) + active[speaker] = time; + else if (probability < config.threshold && active[speaker] >= 0) + { + if (time > active[speaker]) + result.segments.push_back({active[speaker], time, speaker}); + active[speaker] = -1; + } + } + } + } + for (int speaker = 0; speaker < 8; ++speaker) + if (active[speaker] >= 0 && active[speaker] < result.audio_seconds) + result.segments.push_back({active[speaker], static_cast(frames) * 0.01f, speaker}); + std::sort(result.segments.begin(), result.segments.end(), + [](const auto& a, const auto& b) + { + return a.start_time < b.start_time || (a.start_time == b.start_time && a.speaker < b.speaker); + }); + result.process_seconds = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); + return result; + } +}; + +int Result::SpeakerAt(float start, float end) const +{ + if (!std::isfinite(start) || !std::isfinite(end) || end < start) + return -1; + // Aligners can emit zero-duration words; use the speaker active at that instant. + if (start == end) + end = std::nextafter(end, std::numeric_limits::infinity()); + std::array overlap{}; + for (const auto& segment : segments) + { + if (segment.start_time >= end) + break; + overlap[segment.speaker] += + std::max(0.0f, std::min(end, segment.end_time) - std::max(start, segment.start_time)); + } + const auto best = std::max_element(overlap.begin(), overlap.end()); + return *best > 0 ? static_cast(best - overlap.begin()) : -1; +} + +Pipeline::Pipeline(Config config) + : impl_(std::make_unique(std::move(config))) +{ +} +Pipeline::~Pipeline() = default; +Result Pipeline::Diarize(const din::io::Audio& audio) +{ + return impl_->Run(audio); +} +Result Pipeline::DiarizeFile(const std::filesystem::path& path) +{ + return Diarize(din::io::LoadAudio(path, 16000)); +} +} // namespace din::asr::diarization diff --git a/asr/diarization/diarization.h b/asr/diarization/diarization.h new file mode 100644 index 0000000..f8b1680 --- /dev/null +++ b/asr/diarization/diarization.h @@ -0,0 +1,51 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include +#include + +#include "audio.h" + +namespace din::asr::diarization +{ +struct Config +{ + std::string provider = "trt-rtx"; + std::filesystem::path model_dir = "artifacts/nemotron-diarization/onnx-bf16"; + std::filesystem::path ep_cache_dir = "artifacts/nemotron-diarization/rt_cache"; + std::filesystem::path ep_context_dir = "artifacts/nemotron-diarization/ep_context"; + float threshold = 0.5f; +}; + +struct Segment +{ + float start_time = 0; + float end_time = 0; + int speaker = -1; +}; + +struct Result +{ + std::vector segments; + float audio_seconds = 0; + float process_seconds = 0; + // Greatest cumulative overlap; -1 means no active speaker. + int SpeakerAt(float start, float end) const; +}; + +// One synchronous stream per instance. Speaker state resets between recordings. +class Pipeline +{ +public: + explicit Pipeline(Config config = {}); + ~Pipeline(); + Result Diarize(const din::io::Audio& audio); + Result DiarizeFile(const std::filesystem::path& path); + +private: + struct Impl; + std::unique_ptr impl_; +}; +} // namespace din::asr::diarization diff --git a/asr/diarization/main.cpp b/asr/diarization/main.cpp new file mode 100644 index 0000000..214f2fe --- /dev/null +++ b/asr/diarization/main.cpp @@ -0,0 +1,118 @@ +// SPDX-License-Identifier: Apache-2.0 +#include +#include +#include + +#include "argparse/argparse.hpp" +#include "diarization.h" +#include "forced_aligner.h" +#include "qwen3.h" +#include + +int main(int argc, char** argv) +{ + using namespace din::asr; + try + { + diarization::Config config; + argparse::ArgumentParser parser("din_diarization_cli"); + parser.add_argument("audiofile"); + parser.add_argument("--model-dir").default_value(config.model_dir.string()); + parser.add_argument("--provider").default_value(config.provider).choices("cpu", "trt-rtx"); + parser.add_argument("--ep-cache").default_value(config.ep_cache_dir.string()); + parser.add_argument("--ep-context-dir").default_value(config.ep_context_dir.string()); + parser.add_argument("--threshold").default_value(config.threshold).scan<'g', float>(); + parser.add_argument("--asr-dir").default_value(std::string{}).help("Optional Qwen3 ASR export"); + parser.add_argument("--aligner-dir").default_value(std::string{}).help("Optional Qwen3 forced aligner export"); + parser.add_argument("--asr-provider").default_value(std::string{"trt-rtx"}).choices("cpu", "trt-rtx"); + parser.add_argument("--aligner-provider").default_value(std::string{}).choices("", "cpu", "trt-rtx"); + parser.add_argument("--lang-id").default_value(std::string{"auto"}); + parser.add_argument("--repeat").default_value(1).scan<'i', int>().help("Reuse loaded engines for timing"); + parser.parse_args(argc, argv); + config.model_dir = parser.get("--model-dir"); + config.provider = parser.get("--provider"); + config.ep_cache_dir = parser.get("--ep-cache"); + config.ep_context_dir = parser.get("--ep-context-dir"); + config.threshold = parser.get("--threshold"); + const auto audio = din::io::LoadAudio(parser.get("audiofile"), 16000); + const auto asr_dir = parser.get("--asr-dir"); + const auto aligner_dir = parser.get("--aligner-dir"); + if (!aligner_dir.empty() && asr_dir.empty()) + throw std::invalid_argument("--aligner-dir requires --asr-dir"); + const int repeat = parser.get("--repeat"); + if (repeat < 1) + throw std::invalid_argument("--repeat must be positive"); + diarization::Pipeline diarizer(config); + std::unique_ptr asr; + std::unique_ptr aligner; + if (!asr_dir.empty()) + { + qwen3::Qwen3Config c; + c.model_dir = asr_dir; + c.max_chunk_seconds = aligner_dir.empty() ? 0 : 175; + c.provider = parser.get("--asr-provider"); + c.lang_id = parser.get("--lang-id"); + asr = std::make_unique(c); + if (!aligner_dir.empty()) + { + qwen3::ForcedAlignerConfig alignment; + alignment.model_dir = aligner_dir; + alignment.provider = parser.get("--aligner-provider"); + if (alignment.provider.empty()) + alignment.provider = c.provider; + aligner = std::make_unique(alignment); + } + } + nlohmann::json output; + for (int iteration = 0; iteration < repeat; ++iteration) + { + const auto start = std::chrono::steady_clock::now(); + auto task = std::async(std::launch::async, + [&] + { + return diarizer.Diarize(audio); + }); + qwen3::TranscriptionResult transcript; + if (asr) + transcript = asr->Transcribe(audio); + std::vector timestamps; + if (aligner) + { + std::vector segments; + for (const auto& segment : transcript.segments) + segments.push_back({segment.text, segment.start_sample, segment.end_sample, segment.language}); + timestamps = aligner->AlignSegments(audio, segments); + } + const auto result = task.get(); + const float elapsed = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); + output = {{"audio_seconds", result.audio_seconds}, + {"diarize_seconds", result.process_seconds}, + {"process_seconds", elapsed}, + {"rtf", elapsed / result.audio_seconds}, + {"segments", nlohmann::json::array()}}; + for (const auto& segment : result.segments) + output["segments"].push_back( + {{"start_time", segment.start_time}, {"end_time", segment.end_time}, {"speaker", segment.speaker}}); + if (asr) + { + output["transcription"] = transcript.text; + output["reached_eos"] = transcript.reached_eos; + output["timestamps"] = nlohmann::json::array(); + for (const auto& word : timestamps) + output["timestamps"].push_back({{"text", word.text}, + {"start_time", word.start_time}, + {"end_time", word.end_time}, + {"speaker", result.SpeakerAt(word.start_time, word.end_time)}}); + } + std::cerr << "Run " << iteration + 1 << ": " << elapsed << " s, RTF " << elapsed / result.audio_seconds + << '\n'; + } + std::cout << output.dump(2) << '\n'; + return asr && !output.at("reached_eos").get() ? 1 : 0; + } + catch (const std::exception& error) + { + std::cerr << "Diarization: " << error.what() << '\n'; + return 1; + } +} diff --git a/asr/diarization/model_export/export_diarization.py b/asr/diarization/model_export/export_diarization.py new file mode 100644 index 0000000..497eeea --- /dev/null +++ b/asr/diarization/model_export/export_diarization.py @@ -0,0 +1,69 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Export Nemotron-3 diarization with fixed shapes and persistent speaker state.""" + +import argparse +import hashlib +import json +import sys +from pathlib import Path + +import onnx +import torch +from huggingface_hub import hf_hub_download + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) +from model import DiarizationStep, load_model # noqa: E402 + +MODEL_ID = "nvidia/Nemotron-3-Diarization" + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", type=Path) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--dtype", choices=("fp32", "fp16", "bf16"), default="bf16") + args = parser.parse_args() + checkpoint = args.checkpoint or hf_hub_download(MODEL_ID, "Nemotron-3-Diarization.nemo") + model = load_model(checkpoint) + args.output.mkdir(parents=True, exist_ok=True) + dtype = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}[args.dtype] + step = DiarizationStep(model).eval() + step.model.to(dtype=dtype) + with torch.inference_mode(): + torch.onnx.export( + step, + step.example_inputs(), + args.output / "diarization.onnx", + dynamo=True, + opset_version=23, + external_data=True, + input_names=step.input_names, + output_names=step.output_names, + ) + onnx.checker.check_model(str(args.output / "diarization.onnx")) + digest = hashlib.sha256() + for path in sorted(args.output.glob("diarization.onnx*")): + with path.open("rb") as file: + while block := file.read(1024 * 1024): + digest.update(block) + metadata = { + "model_id": MODEL_ID, + "format_version": 3, + "dtype": args.dtype, + "sample_rate": 16000, + "chunk_frames": 340, + "right_context": 40, + "subsampling": 8, + "mel_bins": 128, + "cache_frames": 264, + "fifo_frames": 40, + "hidden_size": 512, + "speakers": 8, + "frame_seconds": 0.01, + "graph_hash": digest.hexdigest()[:16], + } + (args.output / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/asr/diarization/model_export/mel.py b/asr/diarization/model_export/mel.py new file mode 100644 index 0000000..43c8fb2 --- /dev/null +++ b/asr/diarization/model_export/mel.py @@ -0,0 +1,31 @@ +# SPDX-License-Identifier: Apache-2.0 +import torch +from common.model_export.log_mel import dft_basis +from torch import nn +from torch.nn import functional as F + + +class LogMel(nn.Module): + # Raw PCM includes 256 left/right STFT samples and one preemphasis sample. + samples = 3039 * 160 + 512 + 1 + + def __init__(self, featurizer): + super().__init__() + assert featurizer.n_fft == 512 and featurizer.hop_length == 160 + assert featurizer.log_zero_guard_type == "add" and featurizer.mag_power == 2 + self.preemph = featurizer.preemph + self.guard = featurizer.log_zero_guard_value + real, imag = dft_basis(512) + window = F.pad(featurizer.window.float(), (56, 56)) + self.register_buffer("real", real * window) + self.register_buffer("imag", imag * window) + self.register_buffer("filters", featurizer.fb.squeeze(0).float()) + + def forward(self, samples, sample_bounds, feature_length): + x = samples[:, 1:] - self.preemph * samples[:, :-1] + positions = torch.arange(1, self.samples, device=x.device) + x = x * ((positions >= sample_bounds[0]) & (positions < sample_bounds[1])) + frames = x.unfold(1, 512, 160) + real, imag = frames @ self.real.T, frames @ self.imag.T + features = ((real.square() + imag.square()) @ self.filters.T + self.guard).log() + return features * (torch.arange(3040, device=x.device) < feature_length)[None, :, None] diff --git a/asr/diarization/model_export/model.py b/asr/diarization/model_export/model.py new file mode 100644 index 0000000..b7b8cc3 --- /dev/null +++ b/asr/diarization/model_export/model.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +import math + +import torch +from mel import LogMel +from torch import nn +from torch.nn import functional as F + + +def load_model(checkpoint): + from nemo.collections.asr.models import SortformerEncLabelModel + + model = SortformerEncLabelModel.restore_from(str(checkpoint), map_location="cpu", strict=True).eval() + if not model.high_resolution or model.encoder.self_attention_model != "rope": + raise ValueError("Expected the Nemotron-3 high-resolution RoPE diarizer") + m = model.sortformer_modules + m.chunk_len, m.chunk_right_context = 340, 40 + m.spkcache_len, m.fifo_len, m.spkcache_update_period = 264, 40, 300 + m._check_streaming_parameters() + model.preprocessor.featurizer.dither = 0 + return model + + +def gather_frames(x, indices): + return x.index_select(1, indices.flatten().clamp(0, x.shape[1] - 1)) + + +def compress_cache(m, embeddings, predictions): + # NeMo's AOSC selection, using functional scatter/where instead of dynamic boolean indexing. + per_speaker = m.spkcache_len // m.n_spk - m.spkcache_sil_frames_per_spk + scores = m._disable_low_scores( + predictions, m._get_log_pred_scores(predictions), math.floor(per_speaker * m.min_pos_scores_rate) + ) + latest = torch.arange(scores.shape[1], device=scores.device) >= m.spkcache_len + scores = scores + latest[None, :, None] * m.scores_boost_latest + for rate, scale in ((m.strong_boost_rate, 2), (m.weak_boost_rate, 1)): + indices = scores.topk(math.floor(per_speaker * rate), dim=1, sorted=False).indices + scores = scores.scatter(1, indices, scores.gather(1, indices) - scale * math.log(0.5)) + scores = F.pad(scores, (0, 0, 0, m.spkcache_sil_frames_per_spk), value=float("inf")) + values, indices = scores.transpose(1, 2).flatten(1).topk(m.spkcache_len, sorted=False) + indices = torch.where(values != float("-inf"), indices, m.max_index).sort(dim=1).values + disabled = (indices == m.max_index) | (indices % scores.shape[1] >= predictions.shape[1]) + indices = torch.where(disabled, 0, indices % scores.shape[1]) + silence = m.learnable_sil_emb.to(embeddings.dtype)[None] + return m._gather_spkcache_and_preds(embeddings, predictions, indices, disabled, silence) + + +class DiarizationStep(nn.Module): + input_names = ["samples", "sample_bounds", "feature_length", "history", "history_preds", "history_length"] + output_names = ["probabilities", "next_history", "next_history_preds"] + + def __init__(self, model): + super().__init__() + self.model = model + self.mel = LogMel(model.preprocessor.featurizer) + for layer in model.encoder.layers: + layer.attn.rope.extend_pe(684, "cpu", next(model.parameters()).dtype) + + def example_inputs(self): + return ( + torch.zeros(1, LogMel.samples), + torch.tensor([257, LogMel.samples]), + torch.tensor([3040]), + torch.zeros(1, 304, 512), + torch.zeros(1, 264, 8), + torch.tensor([0]), + ) + + def forward(self, samples, sample_bounds, feature_length, history, history_preds, history_length): + features = self.mel(samples, sample_bounds, feature_length) + model, m = self.model, self.model.sortformer_modules + dtype = model.encoder.pre_encode.proj.weight.dtype + chunk, chunk_length = model._call_pre_encode(features.to(dtype), feature_length) + cache_length = history_length.clamp(max=264) + fifo_length = history_length - cache_length + positions = torch.arange(684, device=features.device) + packed = torch.where( + (positions < history_length)[None, :, None], + gather_frames(history.to(dtype), positions), + gather_frames(chunk, positions - history_length), + ) + length = history_length + chunk_length + valid = positions < length + encoder = model.encoder + x = encoder.embed_norm(packed) + # NeMo's FlexAttention padding mask expressed as exportable SDPA. + mask = torch.zeros(684, dtype=dtype, device=x.device).masked_fill(~valid, float("-inf")) + mask = mask[None, None, None, :].expand(1, 1, 684, 684) + for layer in encoder.layers: + a = layer.attn + q, k, v = a.w_qkv(layer.norm1(x)).reshape(1, 684, 3, 8, 64).permute(2, 0, 3, 1, 4).unbind(0) + q, k = a.rope(q, k) + if dtype == torch.float32: + attended = ((q @ k.transpose(-1, -2)) * (64**-0.5) + mask).softmax(-1) @ v + else: + attended = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) + x = x + a.out_proj(attended.transpose(1, 2).reshape(1, 684, 512)) + x = x + layer.ffn(layer.norm2(x)) + x = m.encoder_proj(encoder.final_norm(x)) + # Padding must be zero before the upsampling convolution, including the final partial chunk. + x = x * valid[None, :, None] + probabilities = torch.sigmoid(m.forward_speaker_logits(m.upsample_hidden(x))) + probabilities = probabilities * valid.repeat_interleave(8)[None, :, None] + reduced = m.downsample_preds(probabilities, 8).float() + output = gather_frames(probabilities, torch.arange(2720, device=x.device) + history_length * 8).float() + + # Every non-final offline chunk is full. The final state is deliberately unused. + pop_length = fifo_length + 300 + candidate_length = cache_length + pop_length + candidate_positions = torch.arange(604, device=x.device) + candidate = gather_frames(packed, candidate_positions) + candidate_preds = torch.where( + (candidate_positions < cache_length)[None, :, None], + gather_frames(history_preds, candidate_positions), + gather_frames(reduced, candidate_positions), + ) + candidate_preds = candidate_preds * (candidate_positions < candidate_length)[None, :, None] + cache, cache_preds = compress_cache(m, candidate, candidate_preds) + fifo = gather_frames(packed, torch.arange(40, device=x.device) + candidate_length) + return output, torch.cat((cache, fifo), dim=1).float(), cache_preds.float() diff --git a/asr/diarization/model_export/validate_diarization.py b/asr/diarization/model_export/validate_diarization.py new file mode 100644 index 0000000..2a3d5d3 --- /dev/null +++ b/asr/diarization/model_export/validate_diarization.py @@ -0,0 +1,138 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Compare exported diarization probabilities to the official NeMo streaming pipeline.""" + +import argparse +import json +import sys +from pathlib import Path + +import onnxruntime as ort +import soundfile as sf +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) +from mel import LogMel # noqa: E402 +from model import DiarizationStep, load_model # noqa: E402 +from torch.nn import functional as F + + +def run(session, inputs): + values = {name: ort.OrtValue.from_dlpack(tensor.contiguous()) for name, tensor in inputs.items()} + return [torch.from_dlpack(value).clone() for value in session.run_with_ort_values(None, values)] + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", type=Path, required=True) + parser.add_argument("--model-dir", type=Path, required=True) + parser.add_argument("--audio", type=Path, required=True) + parser.add_argument("--seconds", type=float, default=95) + parser.add_argument("--provider", choices=("cpu", "trt-rtx"), default="cpu") + parser.add_argument("--output", type=Path) + parser.add_argument("--native-output", type=Path, help="Check C++ JSON output for the same audio and precision") + args = parser.parse_args() + with sf.SoundFile(args.audio) as file: + if file.samplerate != 16000: + raise ValueError("Validation audio must be 16 kHz; the C++ pipeline handles resampling") + samples = ( + torch.frombuffer( + bytearray(file.buffer_read(int(args.seconds * file.samplerate), dtype="float32")), dtype=torch.float32 + ) + .reshape(-1, file.channels) + .mean(1) + ) + samples = samples[None] + metadata = json.loads((args.model_dir / "metadata.json").read_text()) + dtype = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}[metadata["dtype"]] + if args.provider == "cpu" and dtype != torch.float32: + raise ValueError("Use an FP32 export for CPU validation") + model = load_model(args.checkpoint).cuda() + model.encoder.to(dtype=dtype) + model.sortformer_modules.to(dtype=dtype) + with torch.inference_mode(), torch.autocast("cuda", dtype=dtype, enabled=dtype != torch.float32): + reference = model( + audio_signal=samples.cuda(), audio_signal_length=torch.tensor([samples.numel()], device="cuda") + ) + options = ort.SessionOptions() + options.intra_op_num_threads = 8 + if args.provider == "trt-rtx": + import onnxruntime_ep_nv_tensorrt_rtx as ep + + ort.register_execution_provider_library(ep.get_ep_name(), ep.get_library_path()) + devices = [d for d in ort.get_ep_devices() if d.ep_name == ep.get_ep_name()] + options.add_session_config_entry("session.disable_cpu_ep_fallback", "1") + options.add_provider_for_devices(devices, {"enable_cuda_graph": "0"}) + session = ort.InferenceSession(str(args.model_dir / "diarization.onnx"), options) + length = samples.numel() // 160 + history, history_preds = torch.zeros(1, 304, 512), torch.zeros(1, 264, 8) + results = [] + for offset in range(0, length, 2720): + count = min(3040, length - offset) + begin = offset * 160 - 257 + lo, hi = max(0, begin), min(samples.numel(), begin + LogMel.samples) + pcm = F.pad(samples[:, lo:hi], (lo - begin, begin + LogMel.samples - hi)) + probabilities, history, history_preds = run( + session, + dict( + zip( + DiarizationStep.input_names, + ( + pcm, + torch.tensor([lo - begin, hi - begin]), + torch.tensor([count]), + history, + history_preds, + torch.tensor([304 if offset else 0]), + ), + strict=True, + ) + ), + ) + results.append(probabilities[:, : min(2720, int(length) - offset)]) + actual = torch.cat(results, dim=1) + reference = reference.cpu() + error = (actual - reference).abs() + speech = (actual.amax(-1) > 0.5) | (reference.amax(-1) > 0.5) + speaker_agreement = (actual.argmax(-1)[speech] == reference.argmax(-1)[speech]).float() + active_union = (actual > 0.5) | (reference > 0.5) + active_disagreement = ((actual > 0.5) != (reference > 0.5)).sum() / active_union.sum().clamp_min(1) + report = { + "dtype": metadata["dtype"], + "provider": args.provider, + "frames": actual.shape[1], + "max_probability_error": error.max().item(), + "mean_probability_error": error.mean().item(), + "activity_agreement": ((actual > 0.5) == (reference > 0.5)).float().mean().item(), + "active_disagreement": active_disagreement.item(), + "speaker_agreement": speaker_agreement.mean().item() if speaker_agreement.numel() else 1.0, + } + if args.output: + args.output.write_text(json.dumps(report, indent=2) + "\n") + print(json.dumps(report, indent=2)) + if dtype == torch.float32: + torch.testing.assert_close(actual.float(), reference.float(), atol=3e-3, rtol=3e-3) + else: + # Reduced-precision fusion and cache ranking can move individual boundary probabilities. + assert report["mean_probability_error"] < 0.005 + assert report["active_disagreement"] < 0.01 + assert report["speaker_agreement"] >= 0.99 + if args.native_output: + from nemo.collections.asr.parts.utils.vad_utils import binarization_vectorized + + expected = [] + for speaker in range(8): + spans = binarization_vectorized( + actual[0, :, speaker], {"onset": 0.5, "offset": 0.5, "frame_length_in_sec": 0.01} + ) + expected.extend((float(a), float(b), speaker) for a, b in spans) + expected.sort(key=lambda span: (span[0], span[2])) + text = args.native_output.read_text(encoding="utf-8") + native = json.loads(text[text.index("{") :]) + assert abs(native["audio_seconds"] - samples.numel() / 16000) < 0.01 + observed = [(s["start_time"], s["end_time"], s["speaker"]) for s in native["segments"]] + torch.testing.assert_close(torch.tensor(observed), torch.tensor(expected), rtol=0, atol=0.011) + print("Native timestamps match ONNX / NeMo postprocessing") + + +if __name__ == "__main__": + main()