diff --git a/AGENTS.md b/AGENTS.md index fa3334b..4e3ab2e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -76,12 +76,18 @@ src/ libparakeet implementation tdt.cpp / rnnt.cpp , TDT / RNNT greedy loops streaming_encoder.hpp/cpp, cache-aware streaming FastConformer encoder streaming.hpp/cpp , pk::StreamingSession (carried RNN-T + EOU events) + run_stream_over_pcm + diarization.hpp/cpp, pk::DiarizationModel: offline speaker diarization (Sortformer) + diarization_encoder/head.*, RoPE Transformer encoder + speaker head + diarization_streaming.*, NeMo cache-aware streaming diarization (speaker cache + FIFO) + sas_merge.hpp/cpp , ASR words x speaker segments -> speaker-attributed utterances examples/cli/ parakeet-cli binary subcommands: info, transcribe (+ --stream), quantize + diarize binary: diarize [--stream] scripts/ Python tooling convert_parakeet_to_gguf.py, .nemo/.hf -> GGUF (--dtype f32|f16|q8_0) gen_nemo_baseline.py , NeMo intermediates -> baseline.gguf gen_stream_baseline.py , NeMo cache-aware streaming encode+decode -> stream baseline.gguf + gen_diar_baseline.py , NeMo offline + streaming diarization -> diar baseline.gguf validate_vs_nemo.py , WER parity gate vs NeMo publish_hf.py , convert+quantize -> HF upload (dry-run default) requirements.txt , nemo_toolkit[asr] + gguf @@ -101,10 +107,15 @@ tests/ ctest targets test_streaming_decode.cpp , streaming RNN-T tokens == NeMo cache-aware streaming test_streaming_eou_reset.cpp, multi-utterance streaming: decoder resets on , transcript == NeMo reset-on-EOU (issue #13; PARAKEET_TEST_BASELINE_EOU_RESET) test_capi_stream.cpp , streaming C-API transcript == NeMo streaming (PARAKEET_TEST_BASELINE_EOU_STREAM) + test_diarization_accuracy.cpp, offline diarization == NeMo (PARAKEET_TEST_BASELINE_DIAR) + test_streaming_diarization.cpp, streaming diarization == NeMo streaming (same baseline) + test_combined_offline.cpp, SAS + streaming diarization/SAS through the C-API + test_sas_merge.cpp , SAS merge/grouping (model-independent) python/check_convert.py , converter round-trip (model-dependent) python/check_baseline.py, baseline dumper (model-dependent) fixtures/clip.wav , 2 s 16 kHz mono WAV for stage parity tests fixtures/speech.wav , LibriSpeech 2086-149220-0033, ~7.4 s + fixtures/two_speakers.wav, LibriSpeech 1272 + 2086 alternating A-B-A-B, 23.6 s third_party/ vendored deps ggml/ , submodule pinned at v0.13.0 dr_wav.h , vendored single header @@ -114,6 +125,7 @@ docs/ conversion.md , GGUF schema reference quantization.md , quantization allowlist, policy, measured size + WER per type parity.md , full model coverage matrix + per-stage tensor parity + diarization.md , speaker diarization + speaker-attributed ASR: parity, C-API, speed .github/workflows/ ci.yml , build job (per-push) + closed-loop job (pull_request + dispatch) ``` @@ -258,6 +270,18 @@ parakeet_capi_stream_finalize # flush the end-of-stream tail parakeet_capi_stream_free ``` +Speaker diarization (ABI v7, additive; not used by LocalAI yet). A +diarization GGUF loads into its own `parakeet_ctx`; see `docs/diarization.md`: + +``` +parakeet_capi_diarize_path / _pcm # offline, JSON segments +parakeet_capi_transcribe_and_diarize(_json) # speaker-attributed ASR (two contexts) +parakeet_capi_free_sas_results # frees the array and every .text +parakeet_capi_diarize_stream_begin / _feed / _free / _chunk_samples +parakeet_capi_free_diar_segments +parakeet_capi_sas_stream_begin / _feed / _free +``` + `parakeet_capi_transcribe_path_json(ctx, wav, decoder)` returns malloc'd UTF-8 JSON `{"text":..,"words":[{"w","start","end","conf"}],"tokens":[{"id","t","conf"}]}` (times in seconds, conf in `(0,1]`), built from @@ -300,6 +324,24 @@ drain, a word finalizes when the next `▁`-token arrives, the last word on ## Dumping NeMo baselines +Diarization (needs NeMo main / >= 3.1: NeMo 3.0 cannot load the RoPE +encoder of nvidia/Nemotron-3-Diarization): + +``` +.venv/bin/python scripts/convert_parakeet_to_gguf.py \ + --model nvidia/Nemotron-3-Diarization --output /tmp/diar.gguf +.venv/bin/python scripts/gen_diar_baseline.py \ + --model nvidia/Nemotron-3-Diarization \ + --audio tests/fixtures/two_speakers.wav --output /tmp/diar_baseline.gguf +PARAKEET_TEST_DIAR_GGUF=/tmp/diar.gguf PARAKEET_TEST_BASELINE_DIAR=/tmp/diar_baseline.gguf \ + ctest --test-dir build -R diar --output-on-failure +``` + +Quantized diarization GGUFs keep the same segments but move probabilities +more; set `PARAKEET_TEST_DIAR_PROB_TOL=0.05` for Q8_0. + +ASR: + Used by Phase 1 parity tests. Requires the venv and a 16 kHz mono WAV. ``` diff --git a/CMakeLists.txt b/CMakeLists.txt index 778309d..82e6374 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -92,7 +92,12 @@ set(PARAKEET_SRC src/transducer_batch.cpp src/tokenizer.cpp src/search.cpp - src/transcription.cpp) + src/transcription.cpp + src/diarization.cpp + src/diarization_encoder.cpp + src/diarization_head.cpp + src/sas_merge.cpp + src/diarization_streaming.cpp) if(PARAKEET_SHARED) add_library(parakeet SHARED ${PARAKEET_SRC}) diff --git a/docs/diarization.md b/docs/diarization.md new file mode 100644 index 0000000..6ddefaf --- /dev/null +++ b/docs/diarization.md @@ -0,0 +1,113 @@ +# Speaker diarization + +parakeet.cpp runs [nvidia/Nemotron-3-Diarization](https://huggingface.co/nvidia/Nemotron-3-Diarization), +a Sortformer model that answers "who spoke when" for up to 8 speakers, and +combines it with any Parakeet ASR model for speaker-attributed transcripts +("who said what"). + +## Model + +- Encoder: FeatureStacking (8 mel frames stacked, 80 ms per step) and a + 31-layer pre-norm Transformer with RoPE attention (d_model 512, 8 heads). +- Head: projection to 192, a subpixel Conv1d that upsamples 8x back to 10 ms + frames, two linear layers and a sigmoid per speaker. +- Output: per-speaker activity probabilities every 10 ms. Segments come from + thresholding at 0.5. +- Speakers are numbered in order of first appearance. + +## Converting + +``` +.venv/bin/python scripts/convert_parakeet_to_gguf.py \ + --model nvidia/Nemotron-3-Diarization --dtype q8_0 \ + --output models/nemotron-3-diarization.q8_0.gguf +``` + +`--model` also takes a local `.nemo`. The converter reads the checkpoint +directly, so it works with a NeMo that cannot instantiate the model (NeMo 3.0 +has no RoPE Transformer encoder). Sizes: F32 397 MB, F16 201 MB, Q8_0 109 MB. + +## Offline and streaming inference + +NeMo's `diarize()` for this checkpoint runs cache-aware streaming inference +(`streaming_mode: true` in the model config), and `parakeet_capi_diarize_*` +does the same: audio is processed in 264-step chunks (21.12 s) with a +speaker cache of 264 steps that keeps speaker identities stable across chunks. +This is also what keeps long recordings correct: the model is trained on +sessions of about 105 s, and attending over a whole long clip at once is +outside that range. On a 12 minute, 3 speaker recording the offline path +(which matches NeMo offline on short clips) agrees with NeMo's `diarize()` on +22% of speech frames and puts almost everything on one speaker; the streaming +path agrees on 100%. + +`DiarizationModel::run_offline` / `speaker_probs` still implement the offline +path (NeMo peak-normalizes the waveform there) for short clips and for parity +checks. + +The live streaming API (`parakeet_capi_diarize_stream_*`) takes 16 kHz PCM in +pieces of any size, computes the log-mel incrementally (bit-identical to the +whole-clip mel) and returns segments once per 21.12 s chunk. Segments that +continue past a chunk boundary are not split. + +## Parity with NeMo + +Measured against NeMo main (the reference must support the RoPE encoder), with +`scripts/gen_diar_baseline.py`: + +| Clip | Speakers | Offline prob max diff | Offline segments | Streaming segments | +|---|---|---|---|---| +| `tests/fixtures/two_speakers.wav`, 23.6 s | 2 | 0.004 | 5 / 5 identical | 5 / 5 identical | +| synthetic two-voice dialogue (VibeVoice sample), 68.5 s | 2 | 0.004 | 26 / 26 identical | 26 / 26 identical | +| synthetic three-voice dialogue (VibeVoice sample), 12.3 min | 3 | | | 100% frame agreement | + +Segment boundaries match to the 10 ms frame. F16 gives the same results; Q8_0 +keeps the same segments with a probability max diff of about 0.03. + +After the speaker cache compresses, NeMo picks cache frames with +`torch.topk`, whose order for tied scores is arbitrary. parakeet.cpp breaks ties +toward the earlier frame, so streaming probabilities can drift by up to about +0.02 on long clips while the segments stay the same. + +## Speaker-attributed ASR + +`parakeet_capi_transcribe_and_diarize(_json)` runs an ASR context and a +diarization context on the same PCM and assigns every ASR word to the speaker +whose segments overlap it most. A word that overlaps no segment takes the +nearest segment's speaker if that segment is within 0.5 s, otherwise -1. +Consecutive words from one speaker (gaps up to 0.5 s) form an utterance. + +`parakeet_capi_sas_stream_*` does the same live: when a diarization chunk +completes, the audio not yet committed is transcribed, and words that end at +least 1 s before the chunk edge are committed with their speakers. The rest is +transcribed again with the next chunk, so no word is cut at the edge. + +## Speed + +End to end with the `diarize` example (model load included), AMD Ryzen 9 +9950X3D, CPU: + +| Audio | F32 | Q8_0 | +|---|---|---| +| 23.6 s | 0.43 s | 0.25 s | +| 68.5 s | 0.80 s | 0.60 s | +| 12.3 min | 6.8 s | 6.0 s | + +Streaming cost grows linearly with length (each chunk attends over at most +528 steps), so long recordings run at about 110x real time. + +## Tests + +``` +PARAKEET_TEST_DIAR_GGUF=diar.gguf PARAKEET_TEST_BASELINE_DIAR=diar_baseline.gguf \ +PARAKEET_TEST_GGUF=asr.gguf ctest --test-dir build -R "diar|sas|combined" +``` + +- `test_diarization_accuracy`: offline probabilities and segments, and the + default `diarize_pcm` segments, against NeMo. +- `test_streaming_diarization`: streaming probabilities and segments against + NeMo streaming, from a whole-clip mel and from live 100 ms PCM pieces. +- `test_combined_offline`: speaker-attributed ASR and the streaming C-API on + the two-speaker fixture. +- `test_sas_merge`: word to speaker assignment (no model needed). + +Set `PARAKEET_TEST_DIAR_PROB_TOL=0.05` for Q8_0. diff --git a/examples/cli/CMakeLists.txt b/examples/cli/CMakeLists.txt index 3bb901f..ef14574 100644 --- a/examples/cli/CMakeLists.txt +++ b/examples/cli/CMakeLists.txt @@ -1,3 +1,7 @@ add_executable(parakeet-cli main.cpp) target_link_libraries(parakeet-cli PRIVATE parakeet) target_include_directories(parakeet-cli PRIVATE ${CMAKE_SOURCE_DIR}/src) + +add_executable(diarize diarize.cpp) +target_link_libraries(diarize PRIVATE parakeet) +target_include_directories(diarize PRIVATE ${CMAKE_SOURCE_DIR}/include ${CMAKE_SOURCE_DIR}/src) diff --git a/examples/cli/diarize.cpp b/examples/cli/diarize.cpp new file mode 100644 index 0000000..1e7233d --- /dev/null +++ b/examples/cli/diarize.cpp @@ -0,0 +1,79 @@ +// Standalone diarize tool: loads a diarization GGUF and diarizes a WAV. +// Usage: diarize [--stream] +// Prints {"speakers":N,"segments":[{"speaker","start","end"}, ...]}. +// --stream feeds the audio through the streaming C-API in 1 s pieces (NeMo +// cache-aware streaming) instead of the offline path. +#include "parakeet_capi.h" +#include "audio_io.hpp" + +#include +#include +#include +#include + +static int diarize_stream(parakeet_ctx* ctx, const char* wav) { + pk::Audio audio; + if (!pk::load_audio_16k_mono(wav, audio)) { + std::fprintf(stderr, "cannot read %s\n", wav); + return 1; + } + parakeet_diar_stream* s = parakeet_capi_diarize_stream_begin(ctx); + if (!s) { + std::fprintf(stderr, "stream_begin failed: %s\n", parakeet_capi_last_error(ctx)); + return 1; + } + std::vector all; + const int n = (int)audio.samples.size(); + for (int lo = 0; lo < n || lo == 0; lo += 16000) { + const int len = std::min(16000, n - lo); + parakeet_diar_segment* segs = nullptr; + int ns = 0; + if (parakeet_capi_diarize_stream_feed(s, audio.samples.data() + lo, len, + lo + len >= n, &segs, &ns) != 0) { + std::fprintf(stderr, "stream_feed failed: %s\n", parakeet_capi_last_error(ctx)); + parakeet_capi_diarize_stream_free(s); + return 1; + } + all.insert(all.end(), segs, segs + ns); + parakeet_capi_free_diar_segments(segs); + if (lo + len >= n) break; + } + parakeet_capi_diarize_stream_free(s); + std::sort(all.begin(), all.end(), [](const auto& a, const auto& b) { + return a.start != b.start ? a.start < b.start : a.speaker < b.speaker; + }); + std::printf("{\"speakers\":8,\"segments\":["); + for (size_t i = 0; i < all.size(); ++i) + std::printf("%s{\"speaker\":%d,\"start\":%.2f,\"end\":%.2f}", i ? "," : "", + all[i].speaker, all[i].start, all[i].end); + std::printf("]}\n"); + return 0; +} + +int main(int argc, char** argv) { + if (argc < 3) { + std::fprintf(stderr, "usage: %s [--stream]\n", argv[0]); + return 1; + } + const bool stream = argc > 3 && std::strcmp(argv[3], "--stream") == 0; + parakeet_ctx* ctx = parakeet_capi_load(argv[1]); + if (!ctx) { + std::fprintf(stderr, "failed to load %s\n", argv[1]); + return 1; + } + int rc = 0; + if (stream) { + rc = diarize_stream(ctx, argv[2]); + } else { + char* json = parakeet_capi_diarize_path(ctx, argv[2]); + if (!json) { + std::fprintf(stderr, "diarize failed: %s\n", parakeet_capi_last_error(ctx)); + rc = 1; + } else { + std::printf("%s\n", json); + parakeet_capi_free_string(json); + } + } + parakeet_capi_free(ctx); + return rc; +} diff --git a/include/parakeet_capi.h b/include/parakeet_capi.h index d4a5565..31c23b3 100644 --- a/include/parakeet_capi.h +++ b/include/parakeet_capi.h @@ -47,6 +47,11 @@ typedef struct parakeet_ctx parakeet_ctx; // KenLM) that need the raw distribution rather than this library's own // greedy/beam decode. Freed with the new parakeet_capi_free_logits. The // original entry points are unchanged. +// v7: added speaker diarization (parakeet_capi_diarize_*), speaker-attributed +// ASR (parakeet_capi_transcribe_and_diarize*, parakeet_capi_sas_stream_*) +// for nvidia/Nemotron-3-Diarization. A parakeet_ctx now holds either an +// ASR or a diarization model; parakeet_capi_load detects which. No +// existing signatures changed. int parakeet_capi_abi_version(void); // Load a GGUF model. Returns an owning context, or NULL on failure. @@ -323,7 +328,8 @@ char* parakeet_capi_stream_finalize_json(parakeet_stream* s); void parakeet_capi_stream_free(parakeet_stream* s); // Free a string previously returned by parakeet_capi_transcribe_* / -// parakeet_capi_stream_*. Safe on NULL. +// parakeet_capi_stream_* / parakeet_capi_diarize_* / +// parakeet_capi_transcribe_and_diarize_json. Safe on NULL. void parakeet_capi_free_string(char* s); // Human-readable description of the last error on `ctx`, or "" if none. @@ -331,6 +337,121 @@ void parakeet_capi_free_string(char* s); // it (or until parakeet_capi_free). Returns "" if `ctx` is NULL. const char* parakeet_capi_last_error(parakeet_ctx* ctx); +// --------------------------------------------------------------------------- +// Speaker diarization (nvidia/Nemotron-3-Diarization and compatible Sortformer +// models), ABI v7. +// +// A parakeet_ctx loaded from a diarization GGUF holds a diarization model +// instead of an ASR model (parakeet_capi_load detects the arch). The functions +// below are the only valid entry points for such a context; the transcribe_* +// and stream_* functions fail on it with a last_error message, and the +// diarize_* functions fail on an ASR context. +// +// Times are seconds from the start of the audio; speakers are 0-based indices +// in order of first appearance, up to the model's capacity (8). +// --------------------------------------------------------------------------- + +// Offline diarization of a WAV file. Returns a malloc'd UTF-8 JSON document +// (free with parakeet_capi_free_string): +// {"speakers":8,"segments":[{"speaker":0,"start":0.50,"end":5.52}, ...]} +// "speakers" is the model's capacity. Segments are sorted by start time then +// speaker, with times rounded to 10 ms. NULL on error (see last_error). +char* parakeet_capi_diarize_path(parakeet_ctx* ctx, const char* wav_path); + +// Same for in-memory mono float PCM; resampled to 16 kHz when +// `sample_rate != 16000`. +char* parakeet_capi_diarize_pcm(parakeet_ctx* ctx, const float* samples, + int n_samples, int sample_rate); + +// Speaker-attributed ASR ("who said what"): one utterance is a run of +// consecutive words from one speaker. +typedef struct parakeet_sas_result { + int speaker; // 0-based speaker index, -1 = no diarized speaker overlaps + char* text; // utterance text (space-joined words), owned by the array + float start; // first word start (seconds) + float end; // last word end (seconds) + float conf; // min word confidence +} parakeet_sas_result; + +// Run ASR (`asr_ctx`) and diarization (`diar_ctx`) on the same mono float PCM +// and assign each ASR word to the speaker whose segments overlap it most. +// On success returns 0 and sets *out (malloc'd array, free with +// parakeet_capi_free_sas_results) and *n_out; *out may be NULL when +// *n_out == 0. On error returns non-zero and sets last_error on the context +// that failed. +int parakeet_capi_transcribe_and_diarize(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, + const float* samples, int n_samples, + int sample_rate, + parakeet_sas_result** out, int* n_out); + +// Free an array from parakeet_capi_transcribe_and_diarize or +// parakeet_capi_sas_stream_feed, including every .text. Safe on NULL. +void parakeet_capi_free_sas_results(parakeet_sas_result* results, int n); + +// JSON variant with per-utterance and per-word detail (free with +// parakeet_capi_free_string; NULL on error): +// {"speakers":8, +// "utterances":[{"speaker":0,"text":"hello world","start":0.12,"end":0.85,"conf":0.95}], +// "words":[{"speaker":0,"text":"hello","start":0.12,"end":0.45,"conf":0.97}]} +char* parakeet_capi_transcribe_and_diarize_json(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, + const float* samples, int n_samples, + int sample_rate); + +// --- Streaming diarization ------------------------------------------------- +// NeMo cache-aware streaming (speaker cache + FIFO) over live 16 kHz mono +// float PCM. Audio is processed in the model's chunks +// (parakeet_capi_diarize_stream_chunk_samples; 21.12 s for +// Nemotron-3-Diarization), so segments arrive once per chunk. Speaker indices +// stay consistent across chunks. The stream borrows `diar_ctx`: free the +// stream first, and do not use one context from two threads at once. + +typedef struct parakeet_diar_segment { + int speaker; + float start; // seconds from stream start + float end; +} parakeet_diar_segment; + +typedef struct parakeet_diar_stream parakeet_diar_stream; + +// NULL on error (last_error on diar_ctx). +parakeet_diar_stream* parakeet_capi_diarize_stream_begin(parakeet_ctx* diar_ctx); + +// Samples per processing chunk (the segment latency). 0 on NULL. +int parakeet_capi_diarize_stream_chunk_samples(parakeet_diar_stream* s); + +// Feed PCM; `is_last` flushes the tail and closes open segments. Returns 0 and +// sets *out / *n_out to the segments that ENDED since the previous call +// (free with parakeet_capi_free_diar_segments; *out may be NULL when +// *n_out == 0). Non-zero on error (last_error on the stream's diar_ctx). +int parakeet_capi_diarize_stream_feed(parakeet_diar_stream* s, const float* pcm, + int n_samples, int is_last, + parakeet_diar_segment** out, int* n_out); + +void parakeet_capi_free_diar_segments(parakeet_diar_segment* segs); +void parakeet_capi_diarize_stream_free(parakeet_diar_stream* s); + +// --- Streaming speaker-attributed ASR --------------------------------------- +// Streaming diarization plus ASR over the same live 16 kHz PCM. Each time a +// diarization chunk completes, the not-yet-committed audio is transcribed; +// all words but the last (which may still be cut by the chunk edge) are +// committed with their speakers, and the rest is carried into the next +// chunk. `is_last` commits everything. Borrows both contexts. + +typedef struct parakeet_sas_stream parakeet_sas_stream; + +// NULL on error (last_error on the context that failed). +parakeet_sas_stream* parakeet_capi_sas_stream_begin(parakeet_ctx* asr_ctx, + parakeet_ctx* diar_ctx); + +// Returns 0 and sets *out / *n_out to the utterances committed by this call +// (free with parakeet_capi_free_sas_results(*out, *n_out)). Consecutive calls +// can each return an utterance from the same speaker. Non-zero on error. +int parakeet_capi_sas_stream_feed(parakeet_sas_stream* s, const float* pcm, + int n_samples, int is_last, + parakeet_sas_result** out, int* n_out); + +void parakeet_capi_sas_stream_free(parakeet_sas_stream* s); + #ifdef __cplusplus } // extern "C" #endif diff --git a/scripts/convert_parakeet_to_gguf.py b/scripts/convert_parakeet_to_gguf.py index 4c9ed45..a43bd59 100644 --- a/scripts/convert_parakeet_to_gguf.py +++ b/scripts/convert_parakeet_to_gguf.py @@ -41,6 +41,79 @@ print("PARAKEET_CONVERT_DEPS_MISSING", file=sys.stderr) sys.exit(2) +import io +import tarfile + +try: + import torch +except ImportError as e: # pragma: no cover - env guard + print(f"converter: missing dependency 'torch': {e}", file=sys.stderr) + print("PARAKEET_CONVERT_DEPS_MISSING", file=sys.stderr) + sys.exit(2) + +try: + import yaml +except ImportError as e: # pragma: no cover - env guard + print(f"converter: missing dependency 'pyyaml': {e}", file=sys.stderr) + print("PARAKEET_CONVERT_DEPS_MISSING", file=sys.stderr) + sys.exit(2) + + +def _load_diarization_from_tar(nemo_path): + """Load state_dict + config from a .nemo (POSIX tar) for diarization models. + + The .nemo tar contains model_config.yaml + model_weights.ckpt. We load + the state_dict directly with torch.load(weights_only=True) and parse the + YAML config, bypassing SortformerEncLabelModel.restore_from() which + fails on NeMo versions that don't support self_attention_model='rope'. + """ + with tarfile.open(nemo_path, "r") as tar: + # Find model_weights.ckpt and model_config.yaml + weight_names = [m.name for m in tar.getmembers() if "weights" in m.name] + config_names = [m.name for m in tar.getmembers() if "config" in m.name and m.name.endswith((".yaml", ".yml"))] + if not weight_names or not config_names: + raise ValueError(f"could not find weights/config in {nemo_path}") + # Extract weights + w_member = tar.extractfile(weight_names[0]) + buf = io.BytesIO(w_member.read()) + state_dict = torch.load(buf, map_location="cpu", weights_only=True) + # Extract config + c_member = tar.extractfile(config_names[0]) + cfg = yaml.safe_load(c_member) + return state_dict, cfg + + +def _is_diarization_nemo(nemo_path): + """Peek at a .nemo tar to check if it's a diarization model.""" + try: + with tarfile.open(nemo_path, "r") as tar: + config_names = [m.name for m in tar.getmembers() + if "config" in m.name and m.name.endswith((".yaml", ".yml"))] + if not config_names: + return False + c_member = tar.extractfile(config_names[0]) + cfg = yaml.safe_load(c_member) + # Diarization models have sortformer_modules or model.sortformer_modules + return "sortformer_modules" in cfg or ( + "model" in cfg and "sortformer_modules" in cfg.get("model", {}) + ) + except Exception: + return False + + +def _get_cfg_value(cfg, dotted_key, default=None): + """Get a value from a nested dict using dotted notation (a.b.c).""" + keys = dotted_key.split(".") + v = cfg + for k in keys: + if isinstance(v, dict): + v = v.get(k, default) + else: + return default + if v is None: + return default + return v + def _get(cfg, key, default=None): """Read ``key`` from an OmegaConf node or plain object, tolerating both.""" @@ -51,9 +124,8 @@ def _get(cfg, key, default=None): def detect_arch(m): - """Map a NeMo model to one of ctc/rnnt/tdt/hybrid_rnnt_ctc/hybrid_tdt_ctc.""" + """Map a NeMo ASR model to one of ctc/rnnt/tdt/hybrid_rnnt_ctc/hybrid_tdt_ctc.""" cfg = m.cfg - # An aux_ctc *config* block is necessary but not sufficient for a hybrid # model: prompt-conditioned RNNT checkpoints (nemotron) carry an unconfigured # aux_ctc stub (num_classes=-1, empty vocabulary) but NO ctc decoder and zero # ctc_decoder.* weights -- NeMo initializes them RNNT-only. Require an actual @@ -119,6 +191,30 @@ def prompt_config(cfg): # stays F32 -- it is intentionally NOT in this allowlist. r"^joint\.enc\.weight$", r"^joint\.pred\.weight$", + # Diarization speaker head linear weights (sortformer_modules). The + # encoder_proj (512->192), first_hidden_to_hidden (192->192), and + # single_hidden_to_spks (192->8) are all pure ggml_mul_mat inputs. + r"^sortformer_modules\.encoder_proj\.weight$", + r"^sortformer_modules\.first_hidden_to_hidden\.weight$", + r"^sortformer_modules\.single_hidden_to_spks\.weight$", + # Diarization transformer encoder linear weights (pre-LN RoPE Transformer). + # Fused QKV (w_qkv), attention output projection (out_proj), and FFN + # up/down linears (ffn.net.0, ffn.net.3) are all pure ggml_mul_mat inputs. + # FeatureStacking projection (encoder.pre_encode.proj) is also pure linear. + r"^encoder\.layers\.\d+\.attn\.w_qkv\.weight$", + r"^encoder\.layers\.\d+\.attn\.out_proj\.weight$", + r"^encoder\.layers\.\d+\.ffn\.net\.\d+\.weight$", + r"^encoder\.pre_encode\.proj\.weight$", +] + +# Weight names that are safe to skip (unused at inference) for diarization models. +DIAIRIZATION_SKIP = [ + r"^encoder\.pos_enc\.", # RoPE, no positional embedding table + r"^hidden_to_spks", # frozen/unused head variant + r"^spec_augmentation", # training-time augmentation + r"^loss", # training loss modules + r"^sortformer_modules\.hidden_to_spks", # unused 384->8 head + r"^sortformer_modules\.transformer_encoder", # None for this model ] _QUANTIZABLE_RE = [re.compile(p) for p in _QUANTIZABLE_PATTERNS] @@ -160,6 +256,166 @@ def main(): args = ap.parse_args() is_local = pathlib.Path(args.model).exists() + + # ------------------------------------------------------------------ + # Diarization path: load state_dict + config directly from the .nemo + # tar, bypassing SortformerEncLabelModel.restore_from() (which fails on + # NeMo versions that don't support self_attention_model='rope'). + # ------------------------------------------------------------------ + nemo_path = args.model if is_local and args.model.endswith(".nemo") else None + if nemo_path is None and not is_local and "/" in args.model: + # HF id: diarization repos ship .nemo; ASR ids fall through to + # ASRModel.from_pretrained below when this is absent. + try: + from huggingface_hub import hf_hub_download + nemo_path = hf_hub_download(args.model, args.model.split("/")[-1] + ".nemo") + except Exception: + nemo_path = None + is_diar = nemo_path is not None and _is_diarization_nemo(nemo_path) + + if is_diar: + sd, model_cfg = _load_diarization_from_tar(nemo_path) + arch = "diarization" + + w = gguf.GGUFWriter(args.output, "parakeet") + w.add_string("general.name", args.model) + w.add_string("parakeet.arch", arch) + + enc_cfg = model_cfg.get("encoder", {}) + sf_cfg = model_cfg.get("sortformer_modules", {}) + pre_cfg = model_cfg.get("preprocessor", {}) + + # Encoder KVs + d_model = int(_get_cfg_value(enc_cfg, "d_model", 512)) + n_layers = int(_get_cfg_value(enc_cfg, "n_layers", 31)) + n_heads = int(_get_cfg_value(enc_cfg, "n_heads", 8)) + ff_exp = float(_get_cfg_value(enc_cfg, "ff_expansion", 4.0)) + ff_dim = int(d_model * ff_exp) + sub_factor = int(_get_cfg_value(enc_cfg, "subsampling_factor", 8)) + + w.add_uint32("parakeet.encoder.feat_in", int(_get_cfg_value(enc_cfg, "feat_in", 128))) + w.add_uint32("parakeet.encoder.d_model", d_model) + w.add_uint32("parakeet.encoder.n_layers", n_layers) + w.add_uint32("parakeet.encoder.n_heads", n_heads) + w.add_uint32("parakeet.encoder.ff_dim", ff_dim) + w.add_uint32("parakeet.encoder.conv_kernel", 0) # N/A for transformer + w.add_string("parakeet.encoder.conv_norm_type", "layer_norm") + w.add_uint32("parakeet.encoder.subsampling_factor", sub_factor) + w.add_uint32("parakeet.encoder.subsampling_conv_channels", 0) + w.add_bool("parakeet.encoder.xscaling", + bool(_get_cfg_value(enc_cfg, "xscaling", False))) + w.add_uint32("parakeet.encoder.pos_emb_max_len", + int(_get_cfg_value(enc_cfg, "pos_emb_max_len", 5000))) + w.add_bool("parakeet.encoder.use_bias", + bool(_get_cfg_value(enc_cfg, "use_bias", False))) + + # Transformer-specific KVs (RoPE attention) + w.add_string("parakeet.encoder.self_attention_model", + str(_get_cfg_value(enc_cfg, "self_attention_model", "rope"))) + w.add_bool("parakeet.encoder.qkv_bias", + bool(_get_cfg_value(enc_cfg, "qkv_bias", False))) + w.add_bool("parakeet.encoder.pre_block_norm", + bool(_get_cfg_value(enc_cfg, "pre_block_norm", True))) + w.add_float32("parakeet.encoder.rope_base", + float(_get_cfg_value(enc_cfg, "rope_base", 10000.0))) + w.add_float32("parakeet.encoder.rotary_fraction", 1.0) + + # Preprocessor KVs (from flat config — no featurizer object) + sr = int(_get_cfg_value(pre_cfg, "sample_rate", 16000)) + n_mels = int(_get_cfg_value(pre_cfg, "features", 128)) + n_fft = int(_get_cfg_value(pre_cfg, "n_fft", 512)) + win_size = float(_get_cfg_value(pre_cfg, "window_size", 0.025)) + win_stride = float(_get_cfg_value(pre_cfg, "window_stride", 0.01)) + win_length = int(round(win_size * sr)) + hop_length = int(round(win_stride * sr)) + + w.add_uint32("parakeet.preprocessor.sample_rate", sr) + w.add_uint32("parakeet.preprocessor.n_mels", n_mels) + w.add_uint32("parakeet.preprocessor.n_fft", n_fft) + w.add_uint32("parakeet.preprocessor.win_length", win_length) + w.add_uint32("parakeet.preprocessor.hop_length", hop_length) + w.add_float32("parakeet.preprocessor.preemph", + float(_get_cfg_value(pre_cfg, "preemph", 0.97))) + w.add_float32("parakeet.preprocessor.mag_power", 2.0) + w.add_string("parakeet.preprocessor.normalize", + str(_get_cfg_value(pre_cfg, "normalize", "NA"))) + w.add_float32("parakeet.preprocessor.log_zero_guard", 2 ** -24) + + # Diarization head KVs + tf_d_model = int(_get_cfg_value(sf_cfg, "tf_d_model", 192)) + n_spk = int(_get_cfg_value(sf_cfg, "num_spks", 8)) + upsample = sub_factor # high_resolution → 10ms output + + w.add_uint32("parakeet.diar.n_speakers", n_spk) + w.add_uint32("parakeet.diar.tf_d_model", tf_d_model) + w.add_uint32("parakeet.diar.upsample_factor", upsample) + w.add_float32("parakeet.diar.frame_resolution_sec", 0.01) + w.add_float32("parakeet.diar.onset_threshold", 0.5) + w.add_float32("parakeet.diar.offset_threshold", 0.5) + + # NeMo diarize() runs streaming inference when streaming_mode is set. + w.add_bool("parakeet.diar.streaming_mode", + bool(_get_cfg_value(model_cfg, "streaming_mode", False))) + + # Streaming speaker-cache config (SortformerModules), in encoder frames. + # Defaults are the SortformerModules constructor defaults. + def sf(key, default): + return _get_cfg_value(sf_cfg, key, default) + w.add_uint32("parakeet.diar.chunk_len", int(sf("chunk_len", 188))) + w.add_uint32("parakeet.diar.spkcache_len", int(sf("spkcache_len", 188))) + w.add_uint32("parakeet.diar.fifo_len", int(sf("fifo_len", 0))) + w.add_uint32("parakeet.diar.spkcache_update_period", + int(sf("spkcache_update_period", 188))) + w.add_uint32("parakeet.diar.spkcache_sil_frames_per_spk", + int(sf("spkcache_sil_frames_per_spk", 3))) + w.add_float32("parakeet.diar.sil_threshold", float(sf("sil_threshold", 0.2))) + w.add_float32("parakeet.diar.pred_score_threshold", + float(sf("pred_score_threshold", 0.25))) + w.add_float32("parakeet.diar.scores_boost_latest", + float(sf("scores_boost_latest", 0.05))) + w.add_float32("parakeet.diar.strong_boost_rate", float(sf("strong_boost_rate", 0.75))) + w.add_float32("parakeet.diar.weak_boost_rate", float(sf("weak_boost_rate", 1.5))) + w.add_float32("parakeet.diar.min_pos_scores_rate", + float(sf("min_pos_scores_rate", 0.5))) + w.add_bool("parakeet.diar.use_learnable_sil_emb", + bool(sf("use_learnable_sil_emb", False))) + + # Write tensors from state_dict + written = 0 + quantized = 0 + skip_patterns = [re.compile(p) for p in DIAIRIZATION_SKIP] + for name, t in sd.items(): + if any(p.search(name) for p in skip_patterns): + continue + if not hasattr(t, "detach"): + continue + arr = t.detach().cpu().float().numpy() + if arr.ndim == 0: + continue + arr = np.ascontiguousarray(arr, dtype=np.float32) + ggml_ne = list(arr.shape[::-1]) + qtype = should_quantize(name, ggml_ne, args.dtype) + if qtype is None: + w.add_tensor(name, arr) + else: + raw = gguf.quantize(arr, qtype) + w.add_tensor(name, raw, raw_shape=raw.shape, raw_dtype=qtype) + quantized += 1 + written += 1 + + w.write_header_to_file() + w.write_kv_data_to_file() + w.write_tensors_to_file() + w.close() + print( + f"wrote {args.output}: arch={arch} tensors={written} " + f"dtype={args.dtype} quantized={quantized}" + ) + return + + # ------------------------------------------------------------------ + # ASR path: load via NeMo model class (as before) + # ------------------------------------------------------------------ try: if is_local: m = ASRModel.restore_from(args.model, map_location="cpu") @@ -365,6 +621,5 @@ def _int_list(v): f"dtype={args.dtype} quantized={quantized}" ) - if __name__ == "__main__": main() diff --git a/scripts/gen_diar_baseline.py b/scripts/gen_diar_baseline.py new file mode 100644 index 0000000..52977f9 --- /dev/null +++ b/scripts/gen_diar_baseline.py @@ -0,0 +1,100 @@ +#!/usr/bin/env python3 +"""Dump a NeMo speaker-diarization reference to a baseline GGUF. + +Used by tests/test_diarization_accuracy.cpp to check the C++ diarization +pipeline (nvidia/Nemotron-3-Diarization and compatible Sortformer models) +against NeMo on the same audio. + +Needs a NeMo with self_attention_model='rope' support (NeMo main / >= 3.1; +NeMo 3.0 cannot instantiate Nemotron-3-Diarization). + +Stored tensors (numpy shapes; the C++ side reads them outer..inner): + +* ``audio`` ``[S]`` the 16 kHz mono clip the reference used, + so the test does not depend on a wav path +* ``offline_probs`` ``[n_spk, T]`` offline forward() speaker probabilities + (``streaming_mode=False``), one frame per + 10 ms mel frame +* ``offline_segs`` ``[N, 3]`` offline diarize() segments as + (speaker, start_s, end_s) +* ``stream_probs`` ``[n_spk, T]`` streaming forward() probabilities + (``streaming_mode=True``, the model's + own chunk / speaker-cache config) +* ``stream_segs`` ``[N, 3]`` streaming diarize() segments + +``dither`` is forced to 0 so the mel is deterministic (the C++ side has no +dither). + +Usage: + python scripts/gen_diar_baseline.py \\ + --model /path/to/Nemotron-3-Diarization.nemo \\ + --audio tests/fixtures/two_speakers.wav \\ + --output /tmp/diar_baseline.gguf +""" +import argparse +import sys + +import numpy as np + +try: + import gguf + import soundfile as sf + import torch + from nemo.collections.asr.models import SortformerEncLabelModel +except ImportError as e: # pragma: no cover - env guard + print(f"gen_diar_baseline: missing dependency: {e}", file=sys.stderr) + sys.exit(2) + + +def _segments(model, path): + out = model.diarize(audio=[path], batch_size=1) + rows = [] + for s in out[0]: + start, end, spk = s.split() + rows.append([float(spk.split("_")[-1]), float(start), float(end)]) + return np.asarray(rows, dtype=np.float32).reshape(-1, 3) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--model", required=True, help=".nemo path or HF id") + ap.add_argument("--audio", required=True, help="16 kHz mono wav") + ap.add_argument("--output", required=True) + args = ap.parse_args() + + if args.model.endswith(".nemo"): + m = SortformerEncLabelModel.restore_from(args.model, map_location="cpu") + else: + m = SortformerEncLabelModel.from_pretrained(args.model, map_location="cpu") + m.eval() + m.preprocessor.featurizer.dither = 0.0 + + y, sr = sf.read(args.audio, dtype="float32") + if y.ndim != 1 or sr != 16000: + sys.exit(f"gen_diar_baseline: {args.audio} must be 16 kHz mono (got sr={sr}, shape={y.shape})") + x = torch.from_numpy(y)[None] + n = torch.tensor([len(y)]) + + results = {} + for name, streaming in (("offline", False), ("stream", True)): + m.streaming_mode = streaming + with torch.no_grad(): + preds = m.forward(x, n) # [1, T, n_spk] + results[f"{name}_probs"] = preds[0].numpy().T.copy() # [n_spk, T] + results[f"{name}_segs"] = _segments(m, args.audio) + print(f"{name}: probs {results[name + '_probs'].shape}, " + f"{len(results[name + '_segs'])} segments") + + w = gguf.GGUFWriter(args.output, "parakeet-diar-baseline") + w.add_tensor("audio", np.ascontiguousarray(y, dtype=np.float32)) + for k, v in results.items(): + w.add_tensor(k, np.ascontiguousarray(v, dtype=np.float32)) + w.write_header_to_file() + w.write_kv_data_to_file() + w.write_tensors_to_file() + w.close() + print(f"wrote {args.output}") + + +if __name__ == "__main__": + main() diff --git a/src/diarization.cpp b/src/diarization.cpp new file mode 100644 index 0000000..17d8915 --- /dev/null +++ b/src/diarization.cpp @@ -0,0 +1,248 @@ +#include "diarization.hpp" +#include "diarization_streaming.hpp" + +#include "audio_io.hpp" +#include "backend.hpp" +#include "ggml_graph.hpp" + +#include +#include +#include +#include + +namespace pk { + +std::unique_ptr DiarizationModel::load(const std::string& path) { + // unique_ptr via private ctor: construct then load. + std::unique_ptr m(new (std::nothrow) DiarizationModel()); + if (!m) return nullptr; + if (!m->loader_.load(path)) return nullptr; + + const auto& cfg = m->loader_.config(); + if (cfg.arch != "diarization") { + // Not a diarization model — caller should use Model for ASR. + return nullptr; + } + if (!cfg.diarization.present) { + return nullptr; + } + + // Give the weights a backend buffer ONCE so graphs reference them + // directly as leaves (zero per-call copy), same as Model::load. + ensure_weights_realized(m->loader_); + + // Construct the component objects (lightweight views over the ModelLoader). + m->mel_ = std::make_unique(m->loader_); + m->encoder_ = std::make_unique(m->loader_); + m->head_ = std::make_unique(m->loader_); + + return m; +} + +DiarizationResult DiarizationModel::diarize_path(const std::string& wav_path) { + Audio audio; + if (!load_audio_16k_mono(wav_path, audio)) { + throw std::runtime_error("parakeet: failed to load audio: " + wav_path); + } + // load_audio_16k_mono already resamples to 16 kHz mono. + return run(audio.samples); +} + +DiarizationResult DiarizationModel::diarize_pcm( + const std::vector& samples, int sample_rate) { + if (sample_rate <= 0) { + throw std::runtime_error("parakeet: invalid sample_rate"); + } + if (sample_rate == 16000) { + return run(samples); + } + std::vector pcm16k = resample_linear(samples, sample_rate, 16000); + return run(pcm16k); +} + +void DiarizationModel::speaker_probs(const std::vector& samples, + std::vector& probs, + int& n_spk, int& T) const { + const ParakeetConfig& cfg = loader_.config(); + n_spk = (int)cfg.diarization.n_speakers; + T = 0; + probs.clear(); + + // 1. Log-mel front end -> feats [n_mels, T] + // NeMo SortformerEncLabelModel.process_signal peak-normalizes the waveform + // in offline (non-streaming) mode: x * 1 / (max(x) + eps), eps = 1e-3. + // Note max(x), not max(|x|). + std::vector norm(samples); + if (!norm.empty()) { + const float peak = *std::max_element(norm.begin(), norm.end()); + const float scale = 1.0f / (peak + 1e-3f); + for (float& v : norm) v *= scale; + } + std::vector feats; + int n_mels = 0, T_mel = 0; + mel_->compute(norm, feats, n_mels, T_mel); + + // NeMo trims the features to the valid length floor(S / hop) before the + // encoder; the centered STFT yields one extra frame past it. + const int hop = (int)cfg.hop_length; + const int T_valid = hop > 0 ? std::min(T_mel, (int)(samples.size() / hop)) : T_mel; + if (T_valid < T_mel) { + std::vector trimmed((size_t)n_mels * T_valid); + for (int m = 0; m < n_mels; ++m) + std::copy_n(feats.begin() + (size_t)m * T_mel, T_valid, + trimmed.begin() + (size_t)m * T_valid); + feats.swap(trimmed); + T_mel = T_valid; + } + if (T_mel == 0) return; + + // 2. Diarization encoder -> enc_out [T_enc, d_model] (time-major) + std::vector enc_out; + int T_enc = 0; + encoder_->forward(feats, n_mels, T_mel, enc_out, T_enc); + + // 3. Diarization head -> probs [n_spk, T_out] (post-sigmoid) + int T_out = 0; + head_->forward(enc_out, T_enc, probs, n_spk, T_out); + + // High-resolution output has one frame per mel frame; drop the frames the + // FeatureStacking pad added past T_mel (NeMo slices preds to the mel length). + if (T_out > T_mel) { + std::vector cut((size_t)n_spk * T_mel); + for (int s = 0; s < n_spk; ++s) + std::copy_n(probs.begin() + (size_t)s * T_out, T_mel, + cut.begin() + (size_t)s * T_mel); + probs.swap(cut); + T_out = T_mel; + } + T = T_out; +} + +DiarizationResult DiarizationModel::run(const std::vector& samples) { + // NeMo's diarize() follows the checkpoint's streaming_mode (true for + // Nemotron-3-Diarization). The streaming path is also the one that holds + // up on long audio: offline attends over the whole clip, far beyond the + // training sessions (and quadratic in length). + return loader_.config().diarization.streaming_mode ? run_streaming(samples) + : run_offline(samples); +} + +DiarizationResult DiarizationModel::run_offline(const std::vector& samples) const { + std::vector probs; + int n_spk = 0, T = 0; + speaker_probs(samples, probs, n_spk, T); + DiarizationResult result; + result.segments = segments_from_probs(probs, n_spk, T); + result.n_speakers = (int)loader_.config().diarization.n_speakers; + return result; +} + +DiarizationResult DiarizationModel::run_streaming(const std::vector& samples) const { + const ParakeetConfig& cfg = loader_.config(); + DiarizationResult result; + result.n_speakers = (int)cfg.diarization.n_speakers; + + // Whole-clip log-mel, not peak-normalized in streaming mode, trimmed to + // floor(S / hop) frames like NeMo. + std::vector feats; + int n_mels = 0, T = 0; + mel_->compute(samples, feats, n_mels, T); + if (cfg.hop_length > 0) T = std::min(T, (int)(samples.size() / cfg.hop_length)); + if (T <= 0) return result; + const int T_full = (int)(feats.size() / n_mels); + + StreamingDiarization sd(loader_); + const int cm = sd.chunk_mel_frames(); + std::vector chunk; + for (int lo = 0; lo < T; lo += cm) { + const int n = std::min(cm, T - lo); + chunk.resize((size_t)n_mels * n); + for (int m = 0; m < n_mels; ++m) + std::copy_n(feats.begin() + (size_t)m * T_full + lo, n, chunk.begin() + (size_t)m * n); + for (const auto& g : sd.feed_mel_chunk(chunk, n_mels, n, lo + n >= T)) + result.segments.push_back({g.speaker, g.start, g.end}); + } + std::sort(result.segments.begin(), result.segments.end(), + [](const SpeakerSegment& a, const SpeakerSegment& b) { + return a.start != b.start ? a.start < b.start : a.speaker < b.speaker; + }); + return result; +} + +std::vector DiarizationModel::segments_from_probs( + const std::vector& probs, int n_spk, int T) const { + const auto& d = loader_.config().diarization; + return postprocess(probs, n_spk, T, d.frame_resolution_sec, d.onset_threshold, + d.offset_threshold); +} + +std::vector DiarizationModel::postprocess( + const std::vector& probs, int n_spk, int T_out, + float frame_sec, float onset, float offset) const { + + // NeMo predlist_to_timestamps (from nemo.collections.asr.parts.utils.vad_utils): + // + // 1. Hysteresis binarization per speaker: + // - OFF → ON when prob >= onset + // - ON → OFF when prob < offset + // When onset == offset (0.5 for Nemotron-3), this is a simple threshold. + // + // 2. Extract contiguous ON segments per speaker. + // + // 3. min_duration_on / min_duration_off filtering (defaults 0.0 → no-op). + // + // 4. merge_overlap_segment: merges same-speaker segments that overlap + // (from padded/chunked inference). Offline single-pass produces no + // overlaps, so this is a no-op here. + // + // 5. Round timestamps to 2 decimal places. + + // probs is row-major [n_spk, T_out]: probs[s * T_out + t] + std::vector segments; + + for (int s = 0; s < n_spk; ++s) { + const float* p = probs.data() + (size_t)s * T_out; + + bool active = false; + int start_frame = 0; + + for (int t = 0; t < T_out; ++t) { + const bool on = (p[t] >= onset); + if (on && !active) { + // Hysteresis: OFF → ON at onset threshold + start_frame = t; + active = true; + } else if (!on && active) { + // Hysteresis: ON → OFF when prob drops below offset. + // With onset == offset, p[t] < offset ⟺ p[t] < onset ⟺ !on. + float start_sec = start_frame * frame_sec; + float end_sec = t * frame_sec; + segments.push_back({s, start_sec, end_sec}); + active = false; + } + } + // Close any segment still open at the end of the audio. + if (active) { + float start_sec = start_frame * frame_sec; + float end_sec = T_out * frame_sec; + segments.push_back({s, start_sec, end_sec}); + } + } + + // Sort by start time, then by speaker (NeMo returns segments in start order). + std::sort(segments.begin(), segments.end(), + [](const SpeakerSegment& a, const SpeakerSegment& b) { + if (a.start != b.start) return a.start < b.start; + return a.speaker < b.speaker; + }); + + // Round to 2 decimal places (NeMo uses round(ts, 2)). + for (auto& seg : segments) { + seg.start = std::round(seg.start * 100.0f) / 100.0f; + seg.end = std::round(seg.end * 100.0f) / 100.0f; + } + + return segments; +} + +} // namespace pk diff --git a/src/diarization.hpp b/src/diarization.hpp new file mode 100644 index 0000000..4d39b75 --- /dev/null +++ b/src/diarization.hpp @@ -0,0 +1,90 @@ +#pragma once +#include "model_loader.hpp" +#include "diarization_head.hpp" +#include "diarization_encoder.hpp" +#include "mel.hpp" +#include +#include +#include + +namespace pk { + +// A speaker segment: which speaker was active, and when. +// Timestamps are in seconds, matching ASR Word timestamps for Phase 3 (SAS). +struct SpeakerSegment { + int speaker; // 0-indexed speaker ID (0..n_speakers-1) + float start; // segment start in seconds + float end; // segment end in seconds +}; + +// Diarization result: list of speaker segments + metadata. +struct DiarizationResult { + std::vector segments; + int n_speakers; // max speakers the model supports +}; + +// DiarizationModel — offline speaker diarization ("who spoke when"). +// +// Composes a DiarizationEncoder (pre-LN RoPE Transformer for Nemotron-3- +// Diarization) with a DiarizationHead (sortformer speaker sigmoid head). +// The mel frontend is shared with the ASR path. +// +// The model is loaded from a GGUF with arch="diarization". The C++ loader +// reuses ModelLoader (shared with ASR) and constructs a DiarizationEncoder + +// MelFrontend + DiarizationHead — no CTC/RNNT decoder. +class DiarizationModel { +public: + // Load a diarization GGUF. Returns nullptr on failure. + static std::unique_ptr load(const std::string& path); + + // Diarize an audio file (any format audio_io supports). Returns segments. + DiarizationResult diarize_path(const std::string& wav_path); + + // Diarize raw PCM samples (mono float, any sample rate — resampled to 16k). + DiarizationResult diarize_pcm(const std::vector& samples, + int sample_rate); + + // Per-frame speaker activity probabilities for already-16 kHz PCM, the + // same tensor NeMo's offline forward() returns: row-major [n_spk, T] + // (probs[s*T + t], post-sigmoid), one frame per 10 ms mel frame. + void speaker_probs(const std::vector& pcm16k, std::vector& probs, + int& n_spk, int& T) const; + + // Offline segments for probabilities from speaker_probs (hysteresis at the + // model's onset/offset, 10 ms frames, rounded to 10 ms). + std::vector segments_from_probs(const std::vector& probs, + int n_spk, int T) const; + + // The two pipelines diarize_* chooses between (config().diarization. + // streaming_mode, as NeMo's diarize() does). Input is 16 kHz PCM. + DiarizationResult run_offline(const std::vector& pcm16k) const; + DiarizationResult run_streaming(const std::vector& pcm16k) const; + + const ParakeetConfig& config() const { return loader_.config(); } + const ModelLoader& loader() const { return loader_; } + + // Access the mel frontend (for streaming diarization to compute mel features). + const MelFrontend& mel() const { return *mel_; } + +private: + DiarizationModel() = default; + + // Dispatch to run_offline / run_streaming on already-16 kHz PCM. + DiarizationResult run(const std::vector& samples); + + // Post-process per-frame speaker probabilities into speaker segments. + // Matches NeMo's predlist_to_timestamps: hysteresis binarization + // (onset/offset thresholds), segment extraction, min_duration filtering, + // merge overlapping segments, round to 2 decimal places. + std::vector postprocess(const std::vector& probs, + int n_spk, int T_out, + float frame_sec, + float onset, float offset) const; + + ModelLoader loader_; + std::unique_ptr mel_; + std::unique_ptr encoder_; + std::unique_ptr head_; +}; + +} // namespace pk diff --git a/src/diarization_encoder.cpp b/src/diarization_encoder.cpp new file mode 100644 index 0000000..8339dd2 --- /dev/null +++ b/src/diarization_encoder.cpp @@ -0,0 +1,185 @@ +#include "diarization_encoder.hpp" +#include "backend.hpp" +#include "graph_builder.hpp" +#include "ggml_graph.hpp" +#include "ggml.h" + +#include +#include +#include +#include + +namespace pk { + +DiarizationEncoder::DiarizationEncoder(const ModelLoader& ml) + : ml_(ml) { + const auto& cfg = ml.config(); + d_model_ = (int)cfg.d_model; + n_layers_ = (int)cfg.n_layers; + n_heads_ = (int)cfg.n_heads; + subsampling_factor_ = (int)cfg.subsampling_factor; + n_mels_ = (int)cfg.n_mels; + pre_block_norm_ = cfg.pre_block_norm; + rope_base_ = cfg.rope_base; + ln_eps_ = 1e-5f; + + if (n_layers_ <= 0 || d_model_ <= 0 || n_heads_ <= 0 || subsampling_factor_ <= 0 || + d_model_ % n_heads_ != 0) { + throw std::runtime_error("parakeet: invalid diarization encoder config"); + } + if (!cfg.self_attention_model.empty() && cfg.self_attention_model != "rope") { + throw std::runtime_error("parakeet: unsupported diarization self_attention_model '" + + cfg.self_attention_model + "'"); + } + head_dim_ = d_model_ / n_heads_; + n_rot_ = (int)(head_dim_ * cfg.rotary_fraction); +} + +static ggml_tensor* layer_norm(ggml_context* ctx, const ModelLoader& ml, ggml_tensor* x, + const std::string& name, float eps) { + ggml_tensor* g = pk::clone_weight(ctx, ml, (name + ".weight").c_str()); + ggml_tensor* b = pk::clone_weight(ctx, ml, (name + ".bias").c_str()); + return ggml_add(ctx, ggml_mul(ctx, ggml_norm(ctx, x, eps), g), b); +} + +static ggml_tensor* linear(ggml_context* ctx, const ModelLoader& ml, ggml_tensor* x, + const std::string& name) { + x = ggml_mul_mat(ctx, pk::clone_weight(ctx, ml, (name + ".weight").c_str()), x); + ggml_tensor* b = pk::clone_weight_opt(ctx, ml, (name + ".bias").c_str()); + return b ? ggml_add(ctx, x, b) : x; +} + +// mel: ne=[T_padded, n_mels] (the zero-padded [n_mels, T] frontend layout). +// Returns ne=[d_model, T_padded / factor]. +ggml_tensor* DiarizationEncoder::build_pre_encode(ggml_context* ctx, ggml_tensor* mel, + int T_padded) const { + const int factor = subsampling_factor_; + // FeatureStacking: [C, T] -> [T, C] -> reshape [T/f, C*f] (f frames stacked). + ggml_tensor* x = ggml_cont(ctx, ggml_transpose(ctx, mel)); // ne=[n_mels, T_padded] + x = ggml_reshape_2d(ctx, x, (int64_t)n_mels_ * factor, T_padded / factor); + return ggml_mul_mat(ctx, pk::clone_weight(ctx, ml_, "encoder.pre_encode.proj.weight"), x); +} + +// x: pre-encoded ne=[d_model, T]; pos: I32 [T]. embed_norm -> blocks -> +// final_norm, returning ne=[d_model, T]. embed_norm lives here, not in +// pre_encode, because NeMo applies it after the (bypassable) pre-encoder: +// the streaming speaker cache holds pre-norm embeddings. +ggml_tensor* DiarizationEncoder::build_blocks(ggml_context* ctx, ggml_tensor* x, + ggml_tensor* pos) const { + const int d = d_model_, H = n_heads_, hd = head_dim_; + const int64_t T = x->ne[1]; + const float scale = 1.0f / std::sqrt((float)hd); + if (pre_block_norm_) x = layer_norm(ctx, ml_, x, "encoder.embed_norm", ln_eps_); + + for (int i = 0; i < n_layers_; ++i) { + const std::string base = "encoder.layers." + std::to_string(i) + "."; + + // Attention: x = x + out_proj(attn(norm1(x))) + ggml_tensor* h = layer_norm(ctx, ml_, x, base + "norm1", ln_eps_); + ggml_tensor* qkv = linear(ctx, ml_, h, base + "attn.w_qkv"); // ne=[3d, T] + const size_t row = qkv->nb[1]; + ggml_tensor* q = ggml_view_3d(ctx, qkv, hd, H, T, hd * sizeof(float), row, 0); + ggml_tensor* k = ggml_view_3d(ctx, qkv, hd, H, T, hd * sizeof(float), row, + (size_t)d * sizeof(float)); + ggml_tensor* v = ggml_view_3d(ctx, qkv, hd, H, T, hd * sizeof(float), row, + (size_t)2 * d * sizeof(float)); + q = ggml_rope_ext(ctx, ggml_cont(ctx, q), pos, nullptr, n_rot_, GGML_ROPE_TYPE_NEOX, + 0, rope_base_, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + k = ggml_rope_ext(ctx, ggml_cont(ctx, k), pos, nullptr, n_rot_, GGML_ROPE_TYPE_NEOX, + 0, rope_base_, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + + // flash_attn_ext takes q/k/v as [hd, T, H] and returns [hd, H, T] + // (heads already interleaved per frame), which reshapes to [d, T]. + q = ggml_permute(ctx, q, 0, 2, 1, 3); + k = ggml_permute(ctx, k, 0, 2, 1, 3); + v = ggml_permute(ctx, ggml_cont(ctx, v), 0, 2, 1, 3); + ggml_tensor* attn = ggml_flash_attn_ext(ctx, q, k, v, nullptr, scale, 0.0f, 0.0f); + attn = ggml_reshape_2d(ctx, ggml_cont(ctx, attn), d, T); + x = ggml_add(ctx, x, linear(ctx, ml_, attn, base + "attn.out_proj")); + + // Feed-forward: x = x + W2 gelu(W1 norm2(x)) + h = layer_norm(ctx, ml_, x, base + "norm2", ln_eps_); + h = ggml_gelu(ctx, linear(ctx, ml_, h, base + "ffn.net.0")); + x = ggml_add(ctx, x, linear(ctx, ml_, h, base + "ffn.net.3")); + } + return layer_norm(ctx, ml_, x, "encoder.final_norm", ln_eps_); +} + +// Copy mel [n_mels, T] into a zero-padded [n_mels, T_padded] graph input. +static ggml_tensor* mel_input(ggml_context* ctx, GraphInputPool& pool, + const std::vector& mel, int n_mels, int T, int T_padded) { + std::vector& padded = pool.alloc_f32((size_t)n_mels * T_padded); + for (int m = 0; m < n_mels; ++m) + std::copy_n(mel.begin() + (size_t)m * T, T, padded.begin() + (size_t)m * T_padded); + int64_t ne[2] = {T_padded, n_mels}; + return pk::graph_input_tensor(ctx, GGML_TYPE_F32, 2, ne, padded.data(), + padded.size() * sizeof(float)); +} + +static std::vector positions(int T) { + std::vector p(T); + for (int i = 0; i < T; ++i) p[i] = i; + return p; +} + +void DiarizationEncoder::forward(const std::vector& mel, int n_mels, int T, + std::vector& enc_out, int& T_enc) const { + if (n_mels != n_mels_ || mel.size() != (size_t)n_mels * T || T <= 0) + throw std::runtime_error("parakeet: diarization encoder got a bad mel shape"); + const int f = subsampling_factor_; + const int T_padded = (T + f - 1) / f * f; + T_enc = T_padded / f; + const std::vector pos_data = positions(T_enc); + + pk::ensure_weights_realized(ml_); + GraphInputPool pool; + const bool ok = pk::run_graph(0, 0, [&](ggml_context* ctx) -> ggml_tensor* { + ggml_tensor* x = build_pre_encode(ctx, mel_input(ctx, pool, mel, n_mels, T, T_padded), + T_padded); + int64_t pos_ne[1] = {T_enc}; + ggml_tensor* pos = pk::graph_input_tensor(ctx, GGML_TYPE_I32, 1, pos_ne, + const_cast(pos_data.data()), + pos_data.size() * sizeof(int32_t)); + return build_blocks(ctx, x, pos); + }, enc_out); + if (!ok) throw std::runtime_error("parakeet: diarization encoder graph failed"); +} + +void DiarizationEncoder::pre_encode(const std::vector& mel, int n_mels, int T, + std::vector& emb, int& T_enc) const { + if (n_mels != n_mels_ || mel.size() != (size_t)n_mels * T || T <= 0) + throw std::runtime_error("parakeet: diarization pre_encode got a bad mel shape"); + const int f = subsampling_factor_; + const int T_padded = (T + f - 1) / f * f; + T_enc = T_padded / f; + + pk::ensure_weights_realized(ml_); + GraphInputPool pool; + const bool ok = pk::run_graph(0, 0, [&](ggml_context* ctx) -> ggml_tensor* { + return build_pre_encode(ctx, mel_input(ctx, pool, mel, n_mels, T, T_padded), T_padded); + }, emb); + if (!ok) throw std::runtime_error("parakeet: diarization pre_encode graph failed"); +} + +void DiarizationEncoder::transformer_forward(const std::vector& emb, int T_enc, + std::vector& enc_out) const { + if (emb.size() != (size_t)d_model_ * T_enc || T_enc <= 0) + throw std::runtime_error("parakeet: diarization transformer got a bad input shape"); + const std::vector pos_data = positions(T_enc); + + pk::ensure_weights_realized(ml_); + const bool ok = pk::run_graph(0, 0, [&](ggml_context* ctx) -> ggml_tensor* { + int64_t ne[2] = {d_model_, T_enc}; + ggml_tensor* x = pk::graph_input_tensor(ctx, GGML_TYPE_F32, 2, ne, + const_cast(emb.data()), + emb.size() * sizeof(float)); + int64_t pos_ne[1] = {T_enc}; + ggml_tensor* pos = pk::graph_input_tensor(ctx, GGML_TYPE_I32, 1, pos_ne, + const_cast(pos_data.data()), + pos_data.size() * sizeof(int32_t)); + return build_blocks(ctx, x, pos); + }, enc_out); + if (!ok) throw std::runtime_error("parakeet: diarization transformer graph failed"); +} + +} // namespace pk diff --git a/src/diarization_encoder.hpp b/src/diarization_encoder.hpp new file mode 100644 index 0000000..c559341 --- /dev/null +++ b/src/diarization_encoder.hpp @@ -0,0 +1,68 @@ +#pragma once +#include "model_loader.hpp" +#include + +struct ggml_context; +struct ggml_tensor; + +namespace pk { + +// DiarizationEncoder — pre-LN RoPE Transformer encoder for Nemotron-3-Diarization. +// +// This is NOT the FastConformer encoder used by ASR. Nemotron-3-Diarization uses +// a NeMo TransformerEncoder with: +// - FeatureStacking subsampling (8x stack of mel frames -> Linear, no bias) +// - embed_norm LayerNorm (pre_block_norm) +// - N x TransformerBlock (pre-norm): x = x + attn(norm1(x)); x = x + ffn(norm2(x)) +// - attention: fused QKV (optional bias) -> RoPE (GPT-NeoX) -> softmax attention +// -> out_proj (bias) +// - FeedForward: Linear -> GELU -> Linear (both with bias) +// - final_norm LayerNorm +// +// All sequences are TIME-MAJOR row-major [T, d_model] (x[t*d_model + c]), which +// is ggml's natural ne=[d_model, T] layout, so no transposes are needed between +// stages. The mel input keeps the frontend's [n_mels, T] layout. +class DiarizationEncoder { +public: + explicit DiarizationEncoder(const ModelLoader& ml); + + // Full encoder: mel [n_mels, T] (mel[m*T + t]) -> enc_out [T_enc, d_model]. + // T_enc = ceil(T / subsampling). + void forward(const std::vector& mel, int n_mels, int T, + std::vector& enc_out, int& T_enc) const; + + // Streaming split. pre_encode = FeatureStacking + projection (no + // embed_norm), i.e. what NeMo stores in the speaker cache / FIFO. + // mel [n_mels, T] -> emb [T_enc, d_model] + void pre_encode(const std::vector& mel, int n_mels, int T, + std::vector& emb, int& T_enc) const; + + // embed_norm + transformer blocks + final_norm over pre-encoded embeddings + // (NeMo frontend_encoder with bypass_pre_encode=True). + // emb [T_enc, d_model] -> enc_out [T_enc, d_model] + void transformer_forward(const std::vector& emb, int T_enc, + std::vector& enc_out) const; + + int subsampling() const { return subsampling_factor_; } + int n_mels() const { return n_mels_; } + int d_model() const { return d_model_; } + +private: + // Graph builders shared by the entry points above. + ggml_tensor* build_pre_encode(ggml_context* ctx, ggml_tensor* mel, int T_padded) const; + ggml_tensor* build_blocks(ggml_context* ctx, ggml_tensor* x, ggml_tensor* pos) const; + + const ModelLoader& ml_; + int d_model_; + int n_layers_; + int n_heads_; + int head_dim_; + int subsampling_factor_; + int n_mels_; + bool pre_block_norm_; + float rope_base_; + int n_rot_; // rotated dims per head (head_dim * rotary_fraction) + float ln_eps_; +}; + +} // namespace pk diff --git a/src/diarization_head.cpp b/src/diarization_head.cpp new file mode 100644 index 0000000..d61d003 --- /dev/null +++ b/src/diarization_head.cpp @@ -0,0 +1,94 @@ +#include "diarization_head.hpp" +#include "ggml_graph.hpp" +#include "backend.hpp" +#include "ggml.h" +#include +#include +#include + +namespace pk { + +DiarizationHead::DiarizationHead(const ModelLoader& ml) : ml_(ml) { + const auto& cfg = ml.config(); + d_model_ = (int)cfg.d_model; + tf_d_model_ = (int)cfg.diarization.tf_d_model; + n_spk_ = (int)cfg.diarization.n_speakers; + upsample_ = (int)cfg.diarization.upsample_factor; +} + +void DiarizationHead::forward(const std::vector& enc_out, int T_enc, + std::vector& probs, int& n_spk, int& T_out) const { + if (T_enc <= 0 || enc_out.size() != (size_t)d_model_ * T_enc) + throw std::runtime_error("parakeet: diarization head got a bad input shape"); + + n_spk = n_spk_; + const int up = upsample_; + T_out = T_enc * up; + + const ModelLoader& ml = ml_; + const int tf = tf_d_model_; + + pk::ensure_weights_realized(ml); + + const bool ok = pk::run_graph(0, 0, + [&](ggml_context* ctx) -> ggml_tensor* { + // enc_out is time-major [T_enc, d_model] = ggml ne=[d_model, T_enc]. + int64_t x_ne[2] = {d_model_, T_enc}; + ggml_tensor* x = pk::graph_input_tensor(ctx, GGML_TYPE_F32, 2, x_ne, + const_cast(enc_out.data()), + enc_out.size() * sizeof(float)); + + // ---- encoder_proj: Linear(d_model -> tf) ---- + ggml_tensor* ep_w = pk::clone_weight(ctx, ml, "sortformer_modules.encoder_proj.weight"); + ggml_tensor* proj = ggml_mul_mat(ctx, ep_w, x); // ne=[tf, T_enc] + ggml_tensor* ep_b = pk::clone_weight_opt(ctx, ml, "sortformer_modules.encoder_proj.bias"); + if (ep_b) proj = ggml_add(ctx, proj, ep_b); + + // ---- subpixel_upsample: Conv1d(tf -> tf*up, k=3, pad=1) + bias ---- + // im2col + mul_mat in F32 (ggml_conv_1d would force an F16 im2col). + // im2col wants data as ne=[T, IC, N]; proj is ne=[tf, T_enc]. + ggml_tensor* conv_in = ggml_cont(ctx, ggml_transpose(ctx, proj)); + conv_in = ggml_reshape_3d(ctx, conv_in, T_enc, tf, 1); + + // GGUF stores the PyTorch [OC, IC, k] weight as ggml ne=[k, IC, OC]. + ggml_tensor* spk_w = pk::clone_weight(ctx, ml, "sortformer_modules.subpixel_upsample.weight"); + ggml_tensor* cols = ggml_im2col(ctx, spk_w, conv_in, /*s0*/1, /*s1*/0, + /*p0*/1, /*p1*/0, /*d0*/1, /*d1*/0, + /*is_2D*/false, GGML_TYPE_F32); + // cols: ne=[k*IC, T_enc, 1] + cols = ggml_reshape_2d(ctx, cols, cols->ne[0], T_enc); + ggml_tensor* w2d = spk_w->type == GGML_TYPE_F32 + ? spk_w : ggml_cast(ctx, spk_w, GGML_TYPE_F32); + w2d = ggml_reshape_2d(ctx, w2d, spk_w->ne[0] * spk_w->ne[1], spk_w->ne[2]); + ggml_tensor* conv_out = ggml_mul_mat(ctx, w2d, cols); + // conv_out: ne=[OC=tf*up, T_enc], flat[c + t*OC] + + ggml_tensor* spk_b = pk::clone_weight_opt(ctx, ml, "sortformer_modules.subpixel_upsample.bias"); + if (spk_b) conv_out = ggml_add(ctx, conv_out, spk_b); + + // Subpixel shuffle, NeMo SortformerModules.upsample_hidden: + // conv(x).transpose(1,2).reshape(B, T, up, tf).reshape(B, T*up, tf) + // so output frame t*up+u, hidden h reads conv channel u*tf+h at + // frame t. With conv_out time-major (flat[u*tf + h + t*tf*up]) this + // is a plain reshape: element (h, t*up+u) = flat[h + (t*up+u)*tf]. + ggml_tensor* upsampled = ggml_reshape_2d(ctx, conv_out, tf, (int64_t)T_enc * up); + // ne[0]=tf, ne[1]=T_out + + // ---- forward_speaker_logits: relu -> Linear(tf->tf) -> relu -> Linear(tf->ns) -> sigmoid ---- + ggml_tensor* h = ggml_relu(ctx, upsampled); + h = ggml_mul_mat(ctx, pk::clone_weight(ctx, ml, "sortformer_modules.first_hidden_to_hidden.weight"), h); + ggml_tensor* fh_b = pk::clone_weight_opt(ctx, ml, "sortformer_modules.first_hidden_to_hidden.bias"); + if (fh_b) h = ggml_add(ctx, h, fh_b); + h = ggml_relu(ctx, h); + h = ggml_mul_mat(ctx, pk::clone_weight(ctx, ml, "sortformer_modules.single_hidden_to_spks.weight"), h); + ggml_tensor* ss_b = pk::clone_weight_opt(ctx, ml, "sortformer_modules.single_hidden_to_spks.bias"); + if (ss_b) h = ggml_add(ctx, h, ss_b); + // ne=[ns, T_out] -> speaker-major [ns][T_out] for the callers. + return ggml_cont(ctx, ggml_transpose(ctx, ggml_sigmoid(ctx, h))); + }, + probs); + + if (!ok) throw std::runtime_error("parakeet: diarization head graph failed"); +} + +} // namespace pk diff --git a/src/diarization_head.hpp b/src/diarization_head.hpp new file mode 100644 index 0000000..e4c7703 --- /dev/null +++ b/src/diarization_head.hpp @@ -0,0 +1,42 @@ +#pragma once +#include "model_loader.hpp" +#include + +namespace pk { + +// Diarization head — NeMo SortformerModules forward_speaker_logits + upsample. +// +// Architecture (offline path, transformer_encoder is None for Nemotron-3): +// enc_out [T_enc, d_model] (time-major, from DiarizationEncoder) +// → encoder_proj: Linear(d_model → tf_d_model) → [tf_d_model, T_enc] +// → subpixel_upsample: Conv1d(tf_d_model → tf_d_model*upsample, k=3, pad=1) +// → reshape → [tf_d_model, T_enc * upsample] +// → relu → first_hidden_to_hidden: Linear(tf_d_model → tf_d_model) +// → relu → single_hidden_to_spks: Linear(tf_d_model → n_speakers) +// → sigmoid +// → probs [n_speakers, T_enc * upsample] +// +// Weight names (verbatim from state dict): +// sortformer_modules.encoder_proj.{weight,bias} +// sortformer_modules.subpixel_upsample.{weight,bias} +// sortformer_modules.first_hidden_to_hidden.{weight,bias} +// sortformer_modules.single_hidden_to_spks.{weight,bias} +class DiarizationHead { +public: + explicit DiarizationHead(const ModelLoader& ml); + + // enc_out: time-major [T_enc, d_model] (enc_out[t*d_model + c]) + // probs: speaker-major [n_speakers, T_enc * upsample] (probs[s*T_out + t]), + // post-sigmoid + void forward(const std::vector& enc_out, int T_enc, + std::vector& probs, int& n_spk, int& T_out) const; + +private: + const ModelLoader& ml_; + int d_model_; // encoder d_model (512) + int tf_d_model_; // sortformer hidden (192) + int n_spk_; // number of speakers (8) + int upsample_; // upsample factor (8) +}; + +} // namespace pk diff --git a/src/diarization_streaming.cpp b/src/diarization_streaming.cpp new file mode 100644 index 0000000..4cf82fb --- /dev/null +++ b/src/diarization_streaming.cpp @@ -0,0 +1,330 @@ +#include "diarization_streaming.hpp" +#include "backend.hpp" +#include "ggml.h" +#include "ggml-backend.h" + +#include +#include +#include +#include +#include + +namespace pk { + +namespace { + +constexpr float kInf = std::numeric_limits::infinity(); + +// Append rows [lo, hi) of a row-major [*, width] buffer to `dst`. +void append_rows(std::vector& dst, const std::vector& src, int width, + int lo, int hi) { + dst.insert(dst.end(), src.begin() + (size_t)lo * width, src.begin() + (size_t)hi * width); +} + +// Drop the first `n` rows of a row-major [*, width] buffer. +void drop_rows(std::vector& v, int width, int n) { + v.erase(v.begin(), v.begin() + (size_t)n * width); +} + +float round2(float x) { return std::round(x * 100.0f) / 100.0f; } + +} // namespace + +StreamingDiarization::StreamingDiarization(const ModelLoader& ml) + : ml_(ml), encoder_(ml), head_(ml) { + const auto& cfg = ml.config(); + const auto& d = cfg.diarization; + d_model_ = (int)cfg.d_model; + n_spk_ = (int)d.n_speakers; + subsampling_ = encoder_.subsampling(); + upsample_ = (int)d.upsample_factor; + n_mels_ = encoder_.n_mels(); + chunk_len_ = d.chunk_len; + spkcache_len_ = d.spkcache_len; + fifo_len_ = d.fifo_len; + update_period_ = d.spkcache_update_period; + sil_frames_per_spk_ = d.spkcache_sil_frames_per_spk; + frame_sec_ = d.frame_resolution_sec; + onset_ = d.onset_threshold; + offset_ = d.offset_threshold; + sil_threshold_ = d.sil_threshold; + pred_score_threshold_ = d.pred_score_threshold; + scores_boost_latest_ = d.scores_boost_latest; + strong_boost_rate_ = d.strong_boost_rate; + weak_boost_rate_ = d.weak_boost_rate; + min_pos_scores_rate_ = d.min_pos_scores_rate; + + if (chunk_len_ <= 0 || spkcache_len_ <= 0 || fifo_len_ < 0 || update_period_ <= 0 || + upsample_ != subsampling_ || spkcache_len_ / n_spk_ - sil_frames_per_spk_ <= 0) { + throw std::runtime_error("parakeet: invalid diarization streaming config"); + } + + if (d.use_learnable_sil_emb) { + const ggml_tensor* t = ml.tensor("sortformer_modules.learnable_sil_emb"); + if (!t || t->type != GGML_TYPE_F32 || ggml_nelements(t) != d_model_) + throw std::runtime_error("parakeet: learnable_sil_emb missing or not F32"); + ensure_weights_realized(ml); + learnable_sil_emb_.resize(d_model_); + ggml_backend_tensor_get(t, learnable_sil_emb_.data(), 0, d_model_ * sizeof(float)); + use_learnable_sil_emb_ = true; + } + reset(); +} + +void StreamingDiarization::reset() { + spkcache_.clear(); spkcache_preds_.clear(); + fifo_.clear(); fifo_preds_.clear(); + spkcache_compressed_ = false; + mean_sil_emb_.assign(d_model_, 0.0f); + n_sil_frames_ = 0; + frames_done_ = 0; + last_probs_.clear(); + last_frames_ = 0; + active_.assign(n_spk_, 0); + start_frame_.assign(n_spk_, 0); +} + +std::vector StreamingDiarization::feed_mel_chunk( + const std::vector& mel, int n_mels, int n_frames, bool is_last) { + if (n_mels != n_mels_ || n_frames <= 0 || n_frames > chunk_mel_frames() || + mel.size() != (size_t)n_mels * n_frames) { + throw std::runtime_error("parakeet: bad streaming diarization chunk shape"); + } + + // 1. Pre-encode the chunk -> [cl, d]. + std::vector chunk_emb; + int cl = 0; + encoder_.pre_encode(mel, n_mels, n_frames, chunk_emb, cl); + + // 2. Transformer + head over [spkcache | fifo | chunk]. + const int S = (int)(spkcache_.size() / d_model_); + const int F = (int)(fifo_.size() / d_model_); + const int total = S + F + cl; + std::vector seq; + seq.reserve((size_t)total * d_model_); + seq.insert(seq.end(), spkcache_.begin(), spkcache_.end()); + seq.insert(seq.end(), fifo_.begin(), fifo_.end()); + seq.insert(seq.end(), chunk_emb.begin(), chunk_emb.end()); + + std::vector enc; + encoder_.transformer_forward(seq, total, enc); + std::vector hp; // [n_spk, total * up] + int n_spk = 0, T_hr = 0; + head_.forward(enc, total, hp, n_spk, T_hr); + + // 3. Encoder-resolution predictions [total, n_spk]: mean over each block of + // `up` high-resolution frames (NeMo downsample_preds). + std::vector preds((size_t)total * n_spk_); + for (int t = 0; t < total; ++t) + for (int s = 0; s < n_spk_; ++s) { + double acc = 0.0; + for (int u = 0; u < upsample_; ++u) acc += hp[(size_t)s * T_hr + t * upsample_ + u]; + preds[(size_t)t * n_spk_ + s] = (float)(acc / upsample_); + } + + // 4. This chunk's high-resolution slice, trimmed to the real mel frames. + const int base = (S + F) * upsample_; + last_frames_ = n_frames; + last_probs_.resize((size_t)n_spk_ * n_frames); + for (int s = 0; s < n_spk_; ++s) + std::copy_n(hp.begin() + (size_t)s * T_hr + base, n_frames, + last_probs_.begin() + (size_t)s * n_frames); + + // 5. Cache update for the next chunk. + if (!is_last) streaming_update(chunk_emb, cl, preds, S, F); + + std::vector out; + track_segments(last_probs_, n_frames, is_last, out); + return out; +} + +// SortformerModules.streaming_update (sync mode, lc = rc = 0). +void StreamingDiarization::streaming_update(const std::vector& chunk_emb, + int cl, const std::vector& preds, + int S, int F) { + const int d = d_model_, ns = n_spk_; + // FIFO predictions are refreshed from this step's output. + fifo_preds_.assign(preds.begin() + (size_t)S * ns, preds.begin() + (size_t)(S + F) * ns); + fifo_.insert(fifo_.end(), chunk_emb.begin(), chunk_emb.end()); + append_rows(fifo_preds_, preds, ns, S + F, S + F + cl); + + if (F + cl <= fifo_len_) return; + + int pop = std::max(update_period_, cl - fifo_len_ + F); + pop = std::min(pop, F + cl); + + if (!use_learnable_sil_emb_) { + // _get_silence_profile: running mean of popped frames whose summed + // speaker probability is below sil_threshold. + std::vector sum(d, 0.0); + long long count = 0; + for (int t = 0; t < pop; ++t) { + float p = 0.0f; + for (int s = 0; s < ns; ++s) p += fifo_preds_[(size_t)t * ns + s]; + if (p < sil_threshold_) { + ++count; + for (int c = 0; c < d; ++c) sum[c] += fifo_[(size_t)t * d + c]; + } + } + if (count > 0) { + const long long n_new = n_sil_frames_ + count; + for (int c = 0; c < d; ++c) + mean_sil_emb_[c] = (float)(((double)mean_sil_emb_[c] * n_sil_frames_ + sum[c]) / n_new); + n_sil_frames_ = n_new; + } + } + + if (!spkcache_compressed_) { + // Until the first compression the cache predictions are this step's. + spkcache_preds_.assign(preds.begin(), preds.begin() + (size_t)S * ns); + } + append_rows(spkcache_, fifo_, d, 0, pop); + append_rows(spkcache_preds_, fifo_preds_, ns, 0, pop); + drop_rows(fifo_, d, pop); + drop_rows(fifo_preds_, ns, pop); + + if ((int)(spkcache_.size() / d) > spkcache_len_) { + compress_spkcache(); + spkcache_compressed_ = true; + } +} + +// SortformerModules._compress_spkcache (eval: no speaker permutation, no noise). +void StreamingDiarization::compress_spkcache() { + const int d = d_model_, ns = n_spk_; + const int n = (int)(spkcache_preds_.size() / ns); + const int per_spk = spkcache_len_ / ns - sil_frames_per_spk_; + const int strong_k = (int)std::floor(per_spk * strong_boost_rate_); + const int weak_k = (int)std::floor(per_spk * weak_boost_rate_); + const int min_pos = (int)std::floor(per_spk * min_pos_scores_rate_); + const float* P = spkcache_preds_.data(); + + // _get_log_pred_scores + std::vector sc((size_t)n * ns); + for (int t = 0; t < n; ++t) { + float log1_sum = 0.0f; + for (int s = 0; s < ns; ++s) + log1_sum += std::log(std::max(1.0f - P[(size_t)t * ns + s], pred_score_threshold_)); + for (int s = 0; s < ns; ++s) { + const float p = P[(size_t)t * ns + s]; + sc[(size_t)t * ns + s] = std::log(std::max(p, pred_score_threshold_)) - + std::log(std::max(1.0f - p, pred_score_threshold_)) + + log1_sum - std::log(0.5f); + } + } + + // _disable_low_scores + for (int s = 0; s < ns; ++s) { + int pos = 0; + for (int t = 0; t < n; ++t) { + float& v = sc[(size_t)t * ns + s]; + if (!(P[(size_t)t * ns + s] > 0.5f)) v = -kInf; + if (v > 0.0f) ++pos; + } + if (pos >= min_pos) + for (int t = 0; t < n; ++t) { + float& v = sc[(size_t)t * ns + s]; + if (!(v > 0.0f) && P[(size_t)t * ns + s] > 0.5f) v = -kInf; + } + } + + // Boost frames newly added since the last compression. + if (scores_boost_latest_ > 0.0f) + for (size_t i = (size_t)spkcache_len_ * ns; i < sc.size(); ++i) sc[i] += scores_boost_latest_; + + // _boost_topk_scores: add -scale*log(0.5) to each speaker's top-k frames. + auto boost = [&](int k, float scale) { + k = std::min(k, n); + if (k <= 0) return; + const float add = -scale * std::log(0.5f); + std::vector idx(n); + for (int s = 0; s < ns; ++s) { + for (int t = 0; t < n; ++t) idx[t] = t; + // Ties (common: confident frames clamp to the same score) go to + // the earlier frame, so the result is deterministic. + std::nth_element(idx.begin(), idx.begin() + (k - 1), idx.end(), [&](int a, int b) { + const float va = sc[(size_t)a * ns + s], vb = sc[(size_t)b * ns + s]; + return va != vb ? va > vb : a < b; + }); + for (int i = 0; i < k; ++i) sc[(size_t)idx[i] * ns + s] += add; + } + }; + boost(strong_k, 2.0f); + boost(weak_k, 1.0f); + + // Append sil_frames_per_spk frames of +inf per speaker (reserved silence slots). + const int n_tot = n + sil_frames_per_spk_; + // _get_topk_indices over the speaker-major flattening (index = s*n_tot + t). + std::vector> flat((size_t)ns * n_tot); + for (int s = 0; s < ns; ++s) + for (int t = 0; t < n_tot; ++t) + flat[(size_t)s * n_tot + t] = {t < n ? sc[(size_t)t * ns + s] : kInf, s * n_tot + t}; + const int K = spkcache_len_; + std::nth_element(flat.begin(), flat.begin() + (K - 1), flat.end(), + [](const std::pair& a, const std::pair& b) { + return a.first != b.first ? a.first > b.first : a.second < b.second; + }); + constexpr int kMaxIndex = std::numeric_limits::max(); + std::vector top(K); + for (int i = 0; i < K; ++i) top[i] = flat[i].first == -kInf ? kMaxIndex : flat[i].second; + std::sort(top.begin(), top.end()); + + // _gather_spkcache_and_preds + const std::vector& sil = use_learnable_sil_emb_ ? learnable_sil_emb_ : mean_sil_emb_; + std::vector new_embs((size_t)K * d), new_preds((size_t)K * ns, 0.0f); + for (int i = 0; i < K; ++i) { + const int t = top[i] == kMaxIndex ? -1 : top[i] % n_tot; + if (t < 0 || t >= n) { + std::copy(sil.begin(), sil.end(), new_embs.begin() + (size_t)i * d); + } else { + std::copy_n(spkcache_.begin() + (size_t)t * d, d, new_embs.begin() + (size_t)i * d); + std::copy_n(spkcache_preds_.begin() + (size_t)t * ns, ns, new_preds.begin() + (size_t)i * ns); + } + } + spkcache_.swap(new_embs); + spkcache_preds_.swap(new_preds); +} + +// Hysteresis binarization carried across chunks, matching the offline +// DiarizationModel::postprocess on the concatenated probabilities. +void StreamingDiarization::track_segments(const std::vector& probs, int n_frames, + bool is_last, + std::vector& out) { + const size_t first = out.size(); + for (int s = 0; s < n_spk_; ++s) { + const float* p = probs.data() + (size_t)s * n_frames; + for (int t = 0; t < n_frames; ++t) { + const long long f = frames_done_ + t; + if (!active_[s] && p[t] >= onset_) { + active_[s] = 1; + start_frame_[s] = f; + } else if (active_[s] && p[t] < offset_) { + active_[s] = 0; + out.push_back({s, round2(start_frame_[s] * frame_sec_), round2(f * frame_sec_)}); + } + } + } + frames_done_ += n_frames; + if (is_last) { + for (int s = 0; s < n_spk_; ++s) + if (active_[s]) { + active_[s] = 0; + out.push_back({s, round2(start_frame_[s] * frame_sec_), + round2(frames_done_ * frame_sec_)}); + } + } + std::sort(out.begin() + first, out.end(), + [](const StreamingSpeakerSegment& a, const StreamingSpeakerSegment& b) { + return a.start != b.start ? a.start < b.start : a.speaker < b.speaker; + }); +} + +std::vector StreamingDiarization::open_segments() const { + std::vector out; + for (int s = 0; s < n_spk_; ++s) + if (active_[s]) + out.push_back({s, round2(start_frame_[s] * frame_sec_), round2(frames_done_ * frame_sec_)}); + return out; +} + +} // namespace pk diff --git a/src/diarization_streaming.hpp b/src/diarization_streaming.hpp new file mode 100644 index 0000000..21e8007 --- /dev/null +++ b/src/diarization_streaming.hpp @@ -0,0 +1,98 @@ +#pragma once +#include "model_loader.hpp" +#include "diarization_encoder.hpp" +#include "diarization_head.hpp" +#include + +namespace pk { + +// A finished speaker segment on the stream's timeline (seconds from stream start). +struct StreamingSpeakerSegment { + int speaker; + float start; + float end; +}; + +// StreamingDiarization — NeMo Sortformer cache-aware streaming ("AOSC"), +// synchronous mode (SortformerEncLabelModel.forward_streaming_step + +// SortformerModules.streaming_update), for nvidia/Nemotron-3-Diarization. +// +// Each chunk of mel frames is pre-encoded (FeatureStacking + projection + +// embed_norm) and the transformer + speaker head run over +// [speaker cache | FIFO | chunk]. The chunk's slice of the high-resolution +// output is the result for that chunk; the downsampled predictions drive the +// FIFO -> speaker-cache update and the score-based cache compression that +// keeps the speaker identities stable across chunks. +// +// chunk_len / spkcache_len / fifo_len / spkcache_update_period come from the +// GGUF and are in ENCODER frames (80 ms), as in NeMo. With the Nemotron-3 +// config a chunk is 264 encoder frames = 2112 mel frames = 21.12 s. +// +// The mel must be the un-normalized log-mel of the stream (NeMo does not +// peak-normalize in streaming mode), e.g. from pk::StreamingMel. +class StreamingDiarization { +public: + explicit StreamingDiarization(const ModelLoader& ml); + + void reset(); + + // Feed the next chunk: row-major [n_mels, n_frames] (mel[m*n_frames + t]), + // 0 < n_frames <= chunk_mel_frames(). Only the final chunk may be short. + // Returns the speaker segments that ENDED in this chunk; with is_last, + // every still-open segment is closed at the end of the stream. + std::vector feed_mel_chunk( + const std::vector& mel, int n_mels, int n_frames, bool is_last); + + // Speaker probabilities of the last fed chunk, speaker-major + // [n_speakers, last_chunk_frames()] (one frame per mel frame, 10 ms). + const std::vector& last_chunk_probs() const { return last_probs_; } + int last_chunk_frames() const { return last_frames_; } + + // Segments that are still active at the current end of the stream, with + // `end` set to the stream time consumed so far. + std::vector open_segments() const; + + int chunk_mel_frames() const { return chunk_len_ * subsampling_; } + int n_mels() const { return n_mels_; } + int n_speakers() const { return n_spk_; } + float frame_sec() const { return frame_sec_; } + // Mel frames consumed so far. + long long frames_done() const { return frames_done_; } + +private: + void streaming_update(const std::vector& chunk_emb, int chunk_frames, + const std::vector& preds, int spkcache_frames, + int fifo_frames); + void compress_spkcache(); + void track_segments(const std::vector& probs, int n_frames, bool is_last, + std::vector& out); + + const ModelLoader& ml_; + DiarizationEncoder encoder_; + DiarizationHead head_; + + int d_model_, n_spk_, subsampling_, upsample_, n_mels_; + int chunk_len_, spkcache_len_, fifo_len_, update_period_, sil_frames_per_spk_; + float frame_sec_, onset_, offset_; + float sil_threshold_, pred_score_threshold_, scores_boost_latest_; + float strong_boost_rate_, weak_boost_rate_, min_pos_scores_rate_; + bool use_learnable_sil_emb_ = false; + std::vector learnable_sil_emb_; // [d_model] + + // Streaming state. Embeddings are time-major [frames, d_model]; + // predictions are [frames, n_spk] at encoder resolution. + std::vector spkcache_, spkcache_preds_; + std::vector fifo_, fifo_preds_; + bool spkcache_compressed_ = false; + std::vector mean_sil_emb_; // [d_model] + long long n_sil_frames_ = 0; + + // Output state. + long long frames_done_ = 0; // mel frames consumed + std::vector last_probs_; + int last_frames_ = 0; + std::vector active_; // per speaker + std::vector start_frame_; // per speaker, when active +}; + +} // namespace pk diff --git a/src/model.cpp b/src/model.cpp index b9d812a..80a4120 100644 --- a/src/model.cpp +++ b/src/model.cpp @@ -45,6 +45,11 @@ std::unique_ptr Model::load(const std::string& gguf_path) { if (!m->loader_.load(gguf_path)) { return nullptr; } + // Model is the ASR entry point — reject diarization models so the C-API + // can fall through to DiarizationModel::load. + if (m->loader_.config().arch == "diarization") { + return nullptr; + } // Give the weights a CPU backend buffer ONCE so graphs reference them // directly as leaves (zero per-call copy). Done at load (vs. lazily on first // clone_weight) so the cost is paid up front, not per utterance. diff --git a/src/model_loader.cpp b/src/model_loader.cpp index 9d3bfe4..ae33991 100644 --- a/src/model_loader.cpp +++ b/src/model_loader.cpp @@ -154,6 +154,13 @@ bool ModelLoader::load(const std::string& path){ // encoder.use_bias: false for nemotron (the attention/FFN linear projections // carry no bias tensor). Defaults true so existing models are unaffected. cfg_.use_bias = kv_bool(gguf_, "parakeet.encoder.use_bias", true); + // Transformer encoder config (diarization models with RoPE attention). + // Absent for ASR (FastConformer) models → safe defaults. + cfg_.self_attention_model = kv_str(gguf_, "parakeet.encoder.self_attention_model", ""); + cfg_.qkv_bias = kv_bool(gguf_, "parakeet.encoder.qkv_bias", false); + cfg_.pre_block_norm = kv_bool(gguf_, "parakeet.encoder.pre_block_norm", true); + cfg_.rope_base = kv_f32(gguf_, "parakeet.encoder.rope_base", 10000.0f); + cfg_.rotary_fraction = kv_f32(gguf_, "parakeet.encoder.rotary_fraction", 1.0f); // Prompt conditioning (multilingual nemotron). Orthogonal capability flag; // absent -> present=false and the engine skips the prompt stage entirely. cfg_.prompt.present = kv_bool(gguf_, "parakeet.prompt.present", false); @@ -190,6 +197,34 @@ bool ModelLoader::load(const std::string& path){ cfg_.max_symbols = kv_u32(gguf_, "parakeet.decoding.max_symbols", 10); cfg_.vocab_size = kv_u32(gguf_, "parakeet.vocab_size"); cfg_.blank_id = kv_u32(gguf_, "parakeet.blank_id"); + // diarization config (absent for ASR models → present=false) + if (gguf_find_key(gguf_, "parakeet.diar.n_speakers") >= 0) { + auto& d = cfg_.diarization; + d.present = true; + d.n_speakers = kv_u32(gguf_, "parakeet.diar.n_speakers"); + d.tf_d_model = kv_u32(gguf_, "parakeet.diar.tf_d_model"); + d.upsample_factor = kv_u32(gguf_, "parakeet.diar.upsample_factor"); + d.frame_resolution_sec = kv_f32(gguf_, "parakeet.diar.frame_resolution_sec", 0.01f); + d.onset_threshold = kv_f32(gguf_, "parakeet.diar.onset_threshold", 0.5f); + d.offset_threshold = kv_f32(gguf_, "parakeet.diar.offset_threshold", 0.5f); + // Streaming (speaker cache) config, in ENCODER frames as in NeMo + // SortformerModules. Defaults are the Nemotron-3-Diarization values, + // for GGUFs converted before these keys were written. + d.streaming_mode = kv_bool(gguf_, "parakeet.diar.streaming_mode", true); + d.chunk_len = (int32_t)kv_u32(gguf_, "parakeet.diar.chunk_len", 264); + d.spkcache_len = (int32_t)kv_u32(gguf_, "parakeet.diar.spkcache_len", 264); + d.fifo_len = (int32_t)kv_u32(gguf_, "parakeet.diar.fifo_len", 0); + d.spkcache_update_period = (int32_t)kv_u32(gguf_, "parakeet.diar.spkcache_update_period", 264); + d.spkcache_sil_frames_per_spk = (int32_t)kv_u32(gguf_, "parakeet.diar.spkcache_sil_frames_per_spk", 1); + d.sil_threshold = kv_f32(gguf_, "parakeet.diar.sil_threshold", 0.2f); + d.pred_score_threshold = kv_f32(gguf_, "parakeet.diar.pred_score_threshold", 0.25f); + d.scores_boost_latest = kv_f32(gguf_, "parakeet.diar.scores_boost_latest", 0.05f); + d.strong_boost_rate = kv_f32(gguf_, "parakeet.diar.strong_boost_rate", 0.75f); + d.weak_boost_rate = kv_f32(gguf_, "parakeet.diar.weak_boost_rate", 1.5f); + d.min_pos_scores_rate = kv_f32(gguf_, "parakeet.diar.min_pos_scores_rate", 0.5f); + d.use_learnable_sil_emb = kv_bool(gguf_, "parakeet.diar.use_learnable_sil_emb", + gguf_find_tensor(gguf_, "sortformer_modules.learnable_sil_emb") >= 0); + } // durations array (stored as INT32 by the converter) { int64_t id = gguf_find_key(gguf_, "parakeet.tdt.durations"); if(id>=0 && gguf_get_arr_type(gguf_,id)==GGUF_TYPE_INT32){ @@ -207,7 +242,7 @@ bool ModelLoader::load(const std::string& path){ const int64_t nt = gguf_get_n_tensors(gguf_); for(int64_t i=0;i0 && cfg_.vocab_size>0; + return cfg_.d_model>0 && (cfg_.vocab_size>0 || cfg_.arch=="diarization"); } ggml_tensor* ModelLoader::tensor(const std::string& n) const { auto it = tensors_.find(n); return it==tensors_.end()? nullptr : it->second; diff --git a/src/model_loader.hpp b/src/model_loader.hpp index f181f24..7cfedae 100644 --- a/src/model_loader.hpp +++ b/src/model_loader.hpp @@ -71,6 +71,40 @@ struct ParakeetConfig { // vocab uint32_t vocab_size=0, blank_id=0; std::vector tokenizer_pieces; + // diarization (SortformerEncLabelModel). present=false for ASR models. + // The diarization head sits after the transformer encoder: encoder_proj + // → subpixel_upsample → speaker sigmoid head. See docs/diarization-plan.md. + struct DiarizationCfg { + bool present=false; + uint32_t n_speakers=0; // max speakers (8 for Nemotron-3-Diarization) + uint32_t tf_d_model=0; // sortformer hidden dim (192) + uint32_t upsample_factor=0; // = subsampling_factor (8 → 10ms frames) + float frame_resolution_sec=0.01f; // output frame duration + float onset_threshold=0.5f; // hysteresis onset + float offset_threshold=0.5f; // hysteresis offset + bool streaming_mode=true; // NeMo diarize() default: streaming + // --- streaming (speaker cache) config, in ENCODER frames (80 ms) --- + int32_t chunk_len=264; // encoder frames per chunk + int32_t spkcache_len=264; // speaker cache size + int32_t fifo_len=0; // FIFO size (0 for Nemotron-3) + int32_t spkcache_update_period=264; // frames popped FIFO -> cache per update + int32_t spkcache_sil_frames_per_spk=1; // reserved silence slots per speaker + float sil_threshold=0.2f; // silence detection threshold + float pred_score_threshold=0.25f; // log-score clamp floor + float scores_boost_latest=0.05f; // boost for latest frames + float strong_boost_rate=0.75f; // strong top-K fraction + float weak_boost_rate=1.5f; // weak top-K fraction + float min_pos_scores_rate=0.5f; // min positive scores fraction + bool use_learnable_sil_emb=false; // silence embedding is a model param + } diarization; + // Diarization encoder config (Nemotron-3-Diarization uses a TransformerEncoder + // with RoPE, not a FastConformer). These are read from parakeet.encoder.* KVs + // but only meaningful when arch == "diarization". + std::string self_attention_model; // "rope" for Nemotron-3-Diarization + bool qkv_bias=false; // QKV projection bias (false) + bool pre_block_norm=true; // embed_norm before blocks (true) + float rope_base=10000.0f; // RoPE theta + float rotary_fraction=1.0f; // fraction of head_dim rotated }; class ModelLoader { public: diff --git a/src/parakeet_capi.cpp b/src/parakeet_capi.cpp index 1a6c9b5..c3eba28 100644 --- a/src/parakeet_capi.cpp +++ b/src/parakeet_capi.cpp @@ -1,12 +1,18 @@ #include "parakeet_capi.h" #include "parakeet.h" // pk::Decoder #include "model.hpp" // pk::Model +#include "diarization.hpp" // pk::DiarizationModel +#include "diarization_streaming.hpp" // pk::StreamingDiarization #include "streaming.hpp" // pk::StreamingSession #include "mel.hpp" // pk::MelFrontend +#include "sas_merge.hpp" // pk::merge_asr_diarization, pk::group_speaker_words #include "transcription.hpp" // pk::Transcription, pk::Word #include "transcription_json.hpp" +#include +#include +#include #include #include #include @@ -35,11 +41,17 @@ // v6: transcribe_pcm_logits, exposing the CTC head's log-prob matrix (row-major // [T, vocab+1], already log-softmaxed) instead of decoded text, freed with // the new free_logits. Original entry points unchanged. -#define PARAKEET_CAPI_ABI_VERSION 6 +// v7: speaker diarization (diarize_*), speaker-attributed ASR +// (transcribe_and_diarize*, sas_stream_*) and streaming diarization +// (diarize_stream_*). A context holds either an ASR or a diarization model. +#define PARAKEET_CAPI_ABI_VERSION 7 // The opaque context: a loaded model plus a buffer for the last error message. +// Exactly one of `model` / `diar` is non-null: ASR models use `model`, +// diarization models (Sortformer) use `diar`. struct parakeet_ctx { std::unique_ptr model; + std::unique_ptr diar; std::string last_error; }; @@ -130,12 +142,28 @@ extern "C" int parakeet_capi_abi_version(void) { extern "C" parakeet_ctx* parakeet_capi_load(const char* gguf_path) { if (!gguf_path) return nullptr; try { - std::unique_ptr model = pk::Model::load(gguf_path); - if (!model) return nullptr; // load failure (bad/missing GGUF) auto* ctx = new (std::nothrow) parakeet_ctx(); if (!ctx) return nullptr; - ctx->model = std::move(model); - return ctx; + + // Try ASR first. Model::load returns nullptr if the GGUF is not a + // valid ASR model (bad/missing file, or arch=="diarization" which + // Model::load rejects). Then try diarization. + std::unique_ptr model = pk::Model::load(gguf_path); + if (model) { + ctx->model = std::move(model); + return ctx; + } + + // Not an ASR model — try diarization. + std::unique_ptr diar = pk::DiarizationModel::load(gguf_path); + if (diar) { + ctx->diar = std::move(diar); + return ctx; + } + + // Neither — load failed entirely. + delete ctx; + return nullptr; } catch (...) { // Never let an exception cross the boundary. return nullptr; @@ -150,7 +178,12 @@ extern "C" char* parakeet_capi_transcribe_path_lang(parakeet_ctx* ctx, const char* wav_path, int decoder, const char* target_lang) { if (!ctx) return nullptr; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return nullptr; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return nullptr; + } if (!wav_path) { ctx->last_error = "wav_path is NULL"; return nullptr; } // NULL / "" -> model default language (ignored by non-prompt models). const std::string lang = target_lang ? target_lang : ""; @@ -180,7 +213,12 @@ extern "C" char* parakeet_capi_transcribe_pcm_lang(parakeet_ctx* ctx, int sample_rate, int decoder, const char* target_lang) { if (!ctx) return nullptr; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return nullptr; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return nullptr; + } if (!samples || n_samples < 0) { ctx->last_error = "invalid samples buffer"; return nullptr; } // NULL / "" -> model default language (ignored by non-prompt models). const std::string lang = target_lang ? target_lang : ""; @@ -257,7 +295,12 @@ extern "C" int parakeet_capi_transcribe_pcm_batch_lang(parakeet_ctx* ctx, const char* target_lang, char** out) { if (!ctx) return 1; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return 1; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return 1; + } if (!samples || !n_samples || !out || n_clips < 0) { ctx->last_error = "invalid batch arguments"; return 1; @@ -314,7 +357,12 @@ extern "C" char* parakeet_capi_transcribe_path_json(parakeet_ctx* ctx, const char* wav_path, int decoder) { if (!ctx) return nullptr; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return nullptr; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return nullptr; + } if (!wav_path) { ctx->last_error = "wav_path is NULL"; return nullptr; } try { pk::Transcription tr = @@ -341,7 +389,12 @@ extern "C" char* parakeet_capi_transcribe_pcm_batch_json_lang(parakeet_ctx* ctx, const float* samples_concat, const int* n_samples, int n_clips, int sample_rate, int decoder, const char* target_lang) { if (!ctx) return nullptr; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return nullptr; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return nullptr; + } if (!samples_concat || !n_samples || n_clips < 0) { ctx->last_error = "invalid batch arguments"; return nullptr; } @@ -546,7 +599,12 @@ std::string feed_available(parakeet_stream* s, bool flush, int& eou_flag, extern "C" parakeet_stream* parakeet_capi_stream_begin_lang(parakeet_ctx* ctx, const char* target_lang) { if (!ctx) return nullptr; - if (!ctx->model) { ctx->last_error = "context has no loaded model"; return nullptr; } + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; use parakeet_capi_diarize_*" + : "context has no loaded model"; + return nullptr; + } if (!ctx->model->config().streaming.present) { ctx->last_error = "model is not a cache-aware streaming model"; return nullptr; @@ -806,3 +864,458 @@ extern "C" const char* parakeet_capi_last_error(parakeet_ctx* ctx) { if (!ctx) return ""; return ctx->last_error.c_str(); } + +// --------------------------------------------------------------------------- +// Speaker diarization + speaker-attributed ASR (ABI v7) +// --------------------------------------------------------------------------- + +namespace { + +bool require_diar(parakeet_ctx* ctx) { + if (!ctx) return false; + if (!ctx->diar) { + ctx->last_error = ctx->model + ? "context holds an ASR model; diarize_* needs a diarization model" + : "context has no loaded model"; + return false; + } + return true; +} + +bool require_asr(parakeet_ctx* ctx) { + if (!ctx) return false; + if (!ctx->model) { + ctx->last_error = ctx->diar + ? "context holds a diarization model; an ASR model is needed here" + : "context has no loaded model"; + return false; + } + return true; +} + +char* diar_result_to_json(const pk::DiarizationResult& r) { + std::string json = "{\"speakers\":"; + pk::append_json_int(json, r.n_speakers); + json += ",\"segments\":["; + for (size_t i = 0; i < r.segments.size(); ++i) { + if (i) json += ','; + json += "{\"speaker\":"; + pk::append_json_int(json, r.segments[i].speaker); + json += ",\"start\":"; + pk::append_json_float(json, "%.2f", r.segments[i].start); + json += ",\"end\":"; + pk::append_json_float(json, "%.2f", r.segments[i].end); + json += '}'; + } + json += "]}"; + return dup_to_c(json); +} + +template +void append_speaker_item(std::string& s, const T& x, const char* time_fmt) { + s += "{\"speaker\":"; + pk::append_json_int(s, x.speaker); + s += ",\"text\":"; + pk::append_json_string(s, x.text); + s += ",\"start\":"; + pk::append_json_float(s, time_fmt, x.start); + s += ",\"end\":"; + pk::append_json_float(s, time_fmt, x.end); + s += ",\"conf\":"; + pk::append_json_float(s, "%.3f", x.conf); + s += '}'; +} + +// Copy utterances into a malloc'd C array. Returns false on allocation failure. +bool to_c_results(const std::vector& utts, + parakeet_sas_result** out, int* n_out) { + *out = nullptr; + *n_out = 0; + if (utts.empty()) return true; + auto* r = static_cast(std::calloc(utts.size(), sizeof(parakeet_sas_result))); + if (!r) return false; + for (size_t i = 0; i < utts.size(); ++i) { + r[i].speaker = utts[i].speaker; + r[i].text = dup_to_c(utts[i].text); + r[i].start = utts[i].start; + r[i].end = utts[i].end; + r[i].conf = utts[i].conf; + if (!r[i].text) { parakeet_capi_free_sas_results(r, (int)i); return false; } + } + *out = r; + *n_out = (int)utts.size(); + return true; +} + +// ASR + diarization on the same audio, merged per word. +bool run_sas(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, + const float* samples, int n_samples, int sample_rate, + std::vector& words, int& n_speakers) { + if (!require_asr(asr_ctx) || !require_diar(diar_ctx)) return false; + if (!samples || n_samples < 0) { + asr_ctx->last_error = "invalid samples buffer"; + return false; + } + const std::vector pcm(samples, samples + n_samples); + pk::Transcription tr; + try { + tr = asr_ctx->model->transcribe_with_timestamps(pcm, sample_rate); + } catch (const std::exception& e) { + asr_ctx->last_error = e.what(); + return false; + } + pk::DiarizationResult dr; + try { + dr = diar_ctx->diar->diarize_pcm(pcm, sample_rate); + } catch (const std::exception& e) { + diar_ctx->last_error = e.what(); + return false; + } + n_speakers = dr.n_speakers; + words = pk::merge_asr_diarization(tr.words, dr.segments); + asr_ctx->last_error.clear(); + diar_ctx->last_error.clear(); + return true; +} + +} // namespace + +extern "C" char* parakeet_capi_diarize_path(parakeet_ctx* ctx, const char* wav_path) { + if (!require_diar(ctx)) return nullptr; + if (!wav_path) { ctx->last_error = "wav_path is NULL"; return nullptr; } + try { + char* out = diar_result_to_json(ctx->diar->diarize_path(wav_path)); + ctx->last_error.clear(); + return out; + } catch (const std::exception& e) { + ctx->last_error = e.what(); + } catch (...) { + ctx->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" char* parakeet_capi_diarize_pcm(parakeet_ctx* ctx, const float* samples, + int n_samples, int sample_rate) { + if (!require_diar(ctx)) return nullptr; + if (!samples || n_samples < 0) { ctx->last_error = "invalid samples buffer"; return nullptr; } + try { + const std::vector pcm(samples, samples + n_samples); + char* out = diar_result_to_json(ctx->diar->diarize_pcm(pcm, sample_rate)); + ctx->last_error.clear(); + return out; + } catch (const std::exception& e) { + ctx->last_error = e.what(); + } catch (...) { + ctx->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" int parakeet_capi_transcribe_and_diarize(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, + const float* samples, int n_samples, + int sample_rate, + parakeet_sas_result** out, int* n_out) { + if (!out || !n_out) return 1; + *out = nullptr; + *n_out = 0; + try { + std::vector words; + int n_speakers = 0; + if (!run_sas(asr_ctx, diar_ctx, samples, n_samples, sample_rate, words, n_speakers)) + return 1; + if (!to_c_results(pk::group_speaker_words(words), out, n_out)) { + asr_ctx->last_error = "out of memory"; + return 1; + } + return 0; + } catch (...) { + if (asr_ctx) asr_ctx->last_error = "unknown error"; + return 1; + } +} + +extern "C" void parakeet_capi_free_sas_results(parakeet_sas_result* results, int n) { + if (!results) return; + for (int i = 0; i < n; ++i) std::free(results[i].text); + std::free(results); +} + +extern "C" char* parakeet_capi_transcribe_and_diarize_json(parakeet_ctx* asr_ctx, + parakeet_ctx* diar_ctx, + const float* samples, int n_samples, + int sample_rate) { + try { + std::vector words; + int n_speakers = 0; + if (!run_sas(asr_ctx, diar_ctx, samples, n_samples, sample_rate, words, n_speakers)) + return nullptr; + const std::vector utts = pk::group_speaker_words(words); + std::string s = "{\"speakers\":"; + pk::append_json_int(s, n_speakers); + s += ",\"utterances\":["; + for (size_t i = 0; i < utts.size(); ++i) { + if (i) s += ','; + append_speaker_item(s, utts[i], "%.2f"); + } + s += "],\"words\":["; + for (size_t i = 0; i < words.size(); ++i) { + if (i) s += ','; + append_speaker_item(s, words[i], "%.3f"); + } + s += "]}"; + return dup_to_c(s); + } catch (...) { + if (asr_ctx) asr_ctx->last_error = "unknown error"; + return nullptr; + } +} + +// --- Streaming diarization ------------------------------------------------- + +struct parakeet_diar_stream { + parakeet_ctx* ctx = nullptr; + std::unique_ptr sd; + std::unique_ptr mel; + std::vector pending; // mel frames not yet diarized, frame-major [t][n_mels] + long long samples_in = 0; // PCM samples fed so far + bool finished = false; +}; + +namespace { + +// Append feat-major [n_mels, n] mel frames to a frame-major buffer. +void push_frames(std::vector& dst, const std::vector& fm, int n_mels, int n) { + const size_t base = dst.size(); + dst.resize(base + (size_t)n * n_mels); + for (int m = 0; m < n_mels; ++m) + for (int t = 0; t < n; ++t) dst[base + (size_t)t * n_mels + m] = fm[(size_t)m * n + t]; +} + +// Feed PCM to the stream's mel front end and run every full diarization chunk +// (and, with is_last, the tail). Closed segments are appended to `segs`. +// Returns the number of chunks run. +int diar_stream_advance(parakeet_diar_stream* s, const float* pcm, int n, bool is_last, + std::vector& segs) { + const int n_mels = s->sd->n_mels(); + int nf = 0; + if (n > 0) { + std::vector fm = s->mel->feed(pcm, n, nf); + push_frames(s->pending, fm, n_mels, nf); + s->samples_in += n; + } + if (is_last) { + std::vector fm = s->mel->finalize(nf); + push_frames(s->pending, fm, n_mels, nf); + // NeMo keeps floor(S / hop) frames; the centered STFT emits one more. + const long long valid = s->samples_in / (long long)s->ctx->diar->config().hop_length; + const long long have = s->sd->frames_done() + (long long)(s->pending.size() / n_mels); + if (have > valid) s->pending.resize(s->pending.size() - (size_t)(have - valid) * n_mels); + } + const int cm = s->sd->chunk_mel_frames(); + int chunks = 0; + for (;;) { + const int avail = (int)(s->pending.size() / n_mels); + const bool last = is_last && avail <= cm; + if (avail < cm && !(last && avail > 0)) break; + const int take = std::min(avail, cm); + std::vector chunk((size_t)n_mels * take); + for (int t = 0; t < take; ++t) + for (int m = 0; m < n_mels; ++m) + chunk[(size_t)m * take + t] = s->pending[(size_t)t * n_mels + m]; + s->pending.erase(s->pending.begin(), s->pending.begin() + (size_t)take * n_mels); + auto closed = s->sd->feed_mel_chunk(chunk, n_mels, take, last); + segs.insert(segs.end(), closed.begin(), closed.end()); + ++chunks; + if (last) break; + } + if (is_last && chunks == 0 && s->sd->frames_done() > 0) { + // Stream length was an exact multiple of the chunk: close open segments. + auto open = s->sd->open_segments(); + segs.insert(segs.end(), open.begin(), open.end()); + } + if (is_last) s->finished = true; + return chunks; +} + +} // namespace + +extern "C" parakeet_diar_stream* parakeet_capi_diarize_stream_begin(parakeet_ctx* diar_ctx) { + if (!require_diar(diar_ctx)) return nullptr; + try { + auto* s = new parakeet_diar_stream(); + s->ctx = diar_ctx; + s->sd = std::make_unique(diar_ctx->diar->loader()); + s->mel = std::make_unique(diar_ctx->diar->loader()); + diar_ctx->last_error.clear(); + return s; + } catch (const std::exception& e) { + diar_ctx->last_error = e.what(); + } catch (...) { + diar_ctx->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" int parakeet_capi_diarize_stream_chunk_samples(parakeet_diar_stream* s) { + if (!s) return 0; + return s->sd->chunk_mel_frames() * (int)s->ctx->diar->config().hop_length; +} + +extern "C" int parakeet_capi_diarize_stream_feed(parakeet_diar_stream* s, const float* pcm, + int n_samples, int is_last, + parakeet_diar_segment** out, int* n_out) { + if (!s || !out || !n_out) return 1; + *out = nullptr; + *n_out = 0; + if ((!pcm && n_samples > 0) || n_samples < 0) { s->ctx->last_error = "invalid samples buffer"; return 1; } + if (s->finished) { s->ctx->last_error = "stream already finished"; return 1; } + try { + std::vector segs; + diar_stream_advance(s, pcm, n_samples, is_last != 0, segs); + if (!segs.empty()) { + auto* r = static_cast(std::malloc(segs.size() * sizeof(parakeet_diar_segment))); + if (!r) { s->ctx->last_error = "out of memory"; return 1; } + for (size_t i = 0; i < segs.size(); ++i) r[i] = {segs[i].speaker, segs[i].start, segs[i].end}; + *out = r; + *n_out = (int)segs.size(); + } + s->ctx->last_error.clear(); + return 0; + } catch (const std::exception& e) { + s->ctx->last_error = e.what(); + } catch (...) { + s->ctx->last_error = "unknown error"; + } + return 1; +} + +extern "C" void parakeet_capi_free_diar_segments(parakeet_diar_segment* segs) { + std::free(segs); +} + +extern "C" void parakeet_capi_diarize_stream_free(parakeet_diar_stream* s) { + delete s; +} + +// --- Streaming speaker-attributed ASR --------------------------------------- + +struct parakeet_sas_stream { + parakeet_ctx* asr = nullptr; + parakeet_diar_stream* diar = nullptr; + std::vector audio; // uncommitted PCM, starting at commit_sec + double commit_sec = 0.0; // stream time of audio[0] + std::vector segs; // closed diarization segments + pk::Word last_word; // last committed word (absolute times) + bool have_last_word = false; +}; + +namespace { + +// Lowercase letters and digits only, for comparing a word heard twice. +std::string word_key(const std::string& w) { + std::string k; + for (unsigned char c : w) + if (std::isalnum(c) || c >= 0x80) k += (char)std::tolower(c); + return k; +} + +} // namespace + +extern "C" parakeet_sas_stream* parakeet_capi_sas_stream_begin(parakeet_ctx* asr_ctx, + parakeet_ctx* diar_ctx) { + if (!require_asr(asr_ctx) || !require_diar(diar_ctx)) return nullptr; + parakeet_diar_stream* d = parakeet_capi_diarize_stream_begin(diar_ctx); + if (!d) return nullptr; + auto* s = new (std::nothrow) parakeet_sas_stream(); + if (!s) { parakeet_capi_diarize_stream_free(d); return nullptr; } + s->asr = asr_ctx; + s->diar = d; + return s; +} + +extern "C" int parakeet_capi_sas_stream_feed(parakeet_sas_stream* s, const float* pcm, + int n_samples, int is_last, + parakeet_sas_result** out, int* n_out) { + if (!s || !out || !n_out) return 1; + *out = nullptr; + *n_out = 0; + if ((!pcm && n_samples > 0) || n_samples < 0) { s->asr->last_error = "invalid samples buffer"; return 1; } + if (s->diar->finished) { s->asr->last_error = "stream already finished"; return 1; } + parakeet_ctx* failed = s->diar->ctx; + try { + std::vector closed; + const int chunks = diar_stream_advance(s->diar, pcm, n_samples, is_last != 0, closed); + for (const auto& c : closed) s->segs.push_back({c.speaker, c.start, c.end}); + if (n_samples > 0) s->audio.insert(s->audio.end(), pcm, pcm + n_samples); + if (chunks == 0 && !is_last) return 0; + + // Diarized audio ends at frames_done; transcribe the uncommitted span. + const double hop_sec = (double)s->diar->ctx->diar->config().hop_length / 16000.0; + const double diar_end = s->diar->sd->frames_done() * hop_sec; + size_t span = is_last ? s->audio.size() + : std::min(s->audio.size(), + (size_t)std::max(0.0, (diar_end - s->commit_sec) * 16000.0)); + failed = s->asr; + std::vector words; + if (span > 0) { + const std::vector seg(s->audio.begin(), s->audio.begin() + span); + words = s->asr->model->transcribe_with_timestamps(seg, 16000).words; + } + // Commit only words that end kSasRightContextSec before the cut: the + // ASR needs right context, and a word at the edge may be cut in half. + // The rest is transcribed again with the next chunk. + constexpr double kSasRightContextSec = 1.0; + size_t keep = words.size(); + double next_commit = s->commit_sec + (double)span / 16000.0; + if (!is_last) { + const double limit = (double)span / 16000.0 - kSasRightContextSec; + keep = 0; + while (keep < words.size() && words[keep].end <= limit) ++keep; + next_commit = s->commit_sec + (keep < words.size() ? words[keep].start + : std::max(0.0, limit)); + } + std::vector committed(words.begin(), words.begin() + keep); + for (auto& w : committed) { w.start += (float)s->commit_sec; w.end += (float)s->commit_sec; } + // ASR timestamps are only accurate to a frame or two, so the tail of the + // previously committed word can be heard again at the new start. + if (s->have_last_word && !committed.empty() && + word_key(committed.front().text) == word_key(s->last_word.text) && + committed.front().start - s->last_word.start < 0.5f) { + committed.erase(committed.begin()); + } + if (!committed.empty()) { + s->last_word = committed.back(); + s->have_last_word = true; + } + + // Speaker segments known so far: closed ones plus those still open. + std::vector segs = s->segs; + for (const auto& o : s->diar->sd->open_segments()) segs.push_back({o.speaker, o.start, o.end}); + const auto utts = pk::group_speaker_words(pk::merge_asr_diarization(committed, segs)); + + const size_t drop = std::min(s->audio.size(), + (size_t)std::llround((next_commit - s->commit_sec) * 16000.0)); + s->audio.erase(s->audio.begin(), s->audio.begin() + drop); + s->commit_sec += (double)drop / 16000.0; + // Segments that ended before the commit point can no longer match a word. + s->segs.erase(std::remove_if(s->segs.begin(), s->segs.end(), + [&](const pk::SpeakerSegment& g) { return g.end < s->commit_sec; }), + s->segs.end()); + + if (!to_c_results(utts, out, n_out)) { s->asr->last_error = "out of memory"; return 1; } + s->asr->last_error.clear(); + return 0; + } catch (const std::exception& e) { + failed->last_error = e.what(); + } catch (...) { + failed->last_error = "unknown error"; + } + return 1; +} + +extern "C" void parakeet_capi_sas_stream_free(parakeet_sas_stream* s) { + if (!s) return; + parakeet_capi_diarize_stream_free(s->diar); + delete s; +} diff --git a/src/sas_merge.cpp b/src/sas_merge.cpp new file mode 100644 index 0000000..82c13b2 --- /dev/null +++ b/src/sas_merge.cpp @@ -0,0 +1,104 @@ +#include "sas_merge.hpp" + +#include +#include + +namespace pk { + +std::vector merge_asr_diarization( + const std::vector& words, + const std::vector& segs, + float max_snap_sec) +{ + std::vector result; + result.reserve(words.size()); + + // Sort segments by start time so we can advance a cursor. + // (They typically arrive already sorted, but don't assume it.) + std::vector sorted_segs = segs; + std::sort(sorted_segs.begin(), sorted_segs.end(), + [](const SpeakerSegment& a, const SpeakerSegment& b) { + if (a.start != b.start) return a.start < b.start; + return a.speaker < b.speaker; + }); + + for (const auto& w : words) { + // Find the dominant speaker: the one with the largest overlap + // between [w.start, w.end] and [seg.start, seg.end]. + int best_speaker = -1; + float best_overlap = 0.0f; + + for (const auto& seg : sorted_segs) { + if (seg.end <= w.start) continue; // segment ends before word + if (seg.start >= w.end) break; // segment starts after word + + float overlap = std::min(w.end, seg.end) - std::max(w.start, seg.start); + if (overlap > best_overlap) { + best_overlap = overlap; + best_speaker = seg.speaker; + } + } + + if (best_speaker < 0) { + // No overlap: snap to the nearest segment within max_snap_sec. + float best_dist = max_snap_sec; + for (const auto& seg : sorted_segs) { + const float dist = seg.end <= w.start ? w.start - seg.end : seg.start - w.end; + if (dist <= best_dist) { + best_dist = dist; + best_speaker = seg.speaker; + } + } + } + + SpeakerWord sw; + sw.speaker = best_speaker; + sw.text = w.text; + sw.start = w.start; + sw.end = w.end; + sw.conf = w.conf; + result.push_back(sw); + } + + return result; +} + +std::vector group_speaker_words( + const std::vector& swords, + float max_gap_sec) +{ + std::vector result; + if (swords.empty()) return result; + + SpeakerUtterance cur; + cur.speaker = swords[0].speaker; + cur.text = swords[0].text; + cur.start = swords[0].start; + cur.end = swords[0].end; + cur.conf = swords[0].conf; + + for (size_t i = 1; i < swords.size(); ++i) { + const auto& w = swords[i]; + float gap = w.start - cur.end; + + if (w.speaker == cur.speaker && gap <= max_gap_sec) { + // Extend current utterance + cur.text += " " + w.text; + cur.end = w.end; + cur.conf = std::min(cur.conf, w.conf); + } else { + // Flush and start new utterance + result.push_back(cur); + cur.speaker = w.speaker; + cur.text = w.text; + cur.start = w.start; + cur.end = w.end; + cur.conf = w.conf; + } + } + result.push_back(cur); + + return result; +} + +} // namespace pk diff --git a/src/sas_merge.hpp b/src/sas_merge.hpp new file mode 100644 index 0000000..2030c86 --- /dev/null +++ b/src/sas_merge.hpp @@ -0,0 +1,50 @@ +#pragma once +#include "transcription.hpp" +#include "diarization.hpp" + +#include +#include + +namespace pk { + +// A word attributed to a speaker — the output of merging ASR +// transcription with diarization segments. +struct SpeakerWord { + int speaker; // from diarization (0-based, -1 = no speaker) + std::string text; // from ASR + float start; // from ASR word (seconds) + float end; // from ASR word (seconds) + float conf; // from ASR word +}; + +// A speaker-attributed utterance: consecutive words from the same speaker +// that form a phrase. Groups SpeakerWords where the speaker doesn't change +// and the gap between words is small. +struct SpeakerUtterance { + int speaker; + std::string text; // space-joined words + float start; // first word start + float end; // last word end + float conf; // min word confidence +}; + +// Merge ASR word timestamps with diarization speaker segments. +// +// For each word, the dominant active speaker is the one whose diarization +// segment overlaps the word's [start, end] interval by the largest amount. +// A word that overlaps no segment (ASR and diarization boundaries can disagree +// by a frame or two) takes the nearest segment's speaker when that segment is +// within `max_snap_sec`; otherwise speaker = -1. +std::vector merge_asr_diarization( + const std::vector& words, + const std::vector& segs, + float max_snap_sec = 0.5f); + +// Group speaker-attributed words into utterances. +// Consecutive words with the same speaker and gap <= max_gap_sec are joined. +// speaker == -1 words are grouped together as "unknown" utterances. +std::vector group_speaker_words( + const std::vector& swords, + float max_gap_sec = 0.5f); + +} // namespace pk diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index d828c6e..96c92e9 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -70,6 +70,11 @@ pk_add_test(test_capi_stream_json) pk_add_test(test_capi_timestamps) pk_add_test(test_capi_batch_json) pk_add_test(test_capi_ctc_logits) +pk_add_test(test_diarization) +pk_add_test(test_diarization_accuracy) +pk_add_test(test_sas_merge) +pk_add_test(test_combined_offline) +pk_add_test(test_streaming_diarization) if(TARGET parakeet-cli) add_test(NAME cli_version_long COMMAND $ --version) @@ -127,6 +132,7 @@ set_tests_properties(test_model_loader test_mel test_mel_gpu test_subsampling te test_transcribe_ctc test_transcribe_rnnt test_transcribe_eou test_transcribe_nemotron test_streaming_decode test_streaming_eou_reset test_streaming_nemotron test_streaming_mel test_capi test_capi_batch test_capi_stream test_capi_stream_json test_capi_timestamps test_capi_batch_json test_capi_ctc_logits + test_combined_offline test_streaming_diarization test_diarization_accuracy test_diarization PROPERTIES LABELS "model") # These tests read fixtures/baselines via paths relative to the project root. set_tests_properties(test_mel test_mel_gpu test_subsampling test_subsampling_batch test_subsampling_batch_causal test_relpos_attention test_relpos_attention_batch test_conformer test_conformer_batch @@ -143,6 +149,7 @@ set_tests_properties(test_mel test_mel_gpu test_subsampling test_subsampling_bat test_transcribe_ctc test_transcribe_rnnt test_transcribe_eou test_transcribe_nemotron test_streaming_decode test_streaming_eou_reset test_streaming_nemotron test_streaming_mel test_capi test_capi_batch test_capi_stream test_capi_stream_json test_capi_timestamps test_capi_batch_json test_capi_ctc_logits + test_combined_offline test_streaming_diarization test_diarization_accuracy PROPERTIES WORKING_DIRECTORY ${CMAKE_SOURCE_DIR}) # Python converter check (skips with exit 77 when the venv/model are absent). diff --git a/tests/fixtures/two_speakers.wav b/tests/fixtures/two_speakers.wav new file mode 100644 index 0000000..f5b37de Binary files /dev/null and b/tests/fixtures/two_speakers.wav differ diff --git a/tests/test_combined_offline.cpp b/tests/test_combined_offline.cpp new file mode 100644 index 0000000..0e7fd54 --- /dev/null +++ b/tests/test_combined_offline.cpp @@ -0,0 +1,331 @@ +// Speaker-attributed ASR (SAS) and streaming diarization through the C-API. +// +// Runs on tests/fixtures/two_speakers.wav (LibriSpeech speakers 1272 and 2086 +// alternating A-B-A-B) and checks: +// 1. transcribe_and_diarize_json: valid document, every word attributed to a +// speaker, speaker turns follow A-B-A-B +// 2. transcribe_and_diarize (struct): same utterances as the JSON variant +// 3. diarize_stream_*: fed live in 0.5 s pieces, the segments match the +// offline diarize_pcm segments (speaker, boundaries within 0.1 s) +// 4. sas_stream_*: fed live in 0.5 s pieces, same turn pattern and about the +// same words as the offline SAS +// +// Env: PARAKEET_TEST_GGUF (ASR model) + PARAKEET_TEST_DIAR_GGUF; skips (77) +// when either is unset. WORKING_DIRECTORY is the repo root. + +#include "parakeet_capi.h" +#include "audio_io.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include + +// Reuse the tiny JSON scanner pattern from test_capi_timestamps.cpp. +namespace { + +struct Scan { + const std::string& s; + size_t i = 0; + explicit Scan(const std::string& str) : s(str) {} + void ws() { while (i < s.size() && (s[i]==' '||s[i]=='\t'||s[i]=='\n'||s[i]=='\r')) ++i; } + bool eat(char c) { ws(); if (i < s.size() && s[i]==c) { ++i; return true; } return false; } + bool str(std::string& out) { + ws(); + if (i >= s.size() || s[i] != '"') return false; + ++i; out.clear(); + while (i < s.size() && s[i] != '"') { + if (s[i] == '\\' && i + 1 < s.size()) { + char c = s[i+1]; + switch (c) { + case 'n': out += '\n'; break; case 't': out += '\t'; break; + case 'r': out += '\r'; break; case 'b': out += '\b'; break; + case 'f': out += '\f'; break; case '"': out += '"'; break; + case '\\': out += '\\'; break; case '/': out += '/'; break; + default: out += c; break; + } + i += 2; + } else { out += s[i++]; } + } + if (i >= s.size()) return false; + ++i; return true; + } + bool num(double& out) { + ws(); + size_t st = i; + while (i < s.size() && std::strchr("+-0123456789.eE", s[i])) ++i; + if (i == st) return false; + out = std::strtod(s.substr(st, i - st).c_str(), nullptr); + return true; + } + bool seek_key(const char* key) { + std::string pat = std::string("\"") + key + "\""; + size_t p = s.find(pat, i); + if (p == std::string::npos) return false; + i = p + pat.size(); + return eat(':'); + } +}; + +// Check that a JSON array key exists (e.g. "utterances":[...]). +bool has_array(const std::string& s, const char* key) { + Scan sc(s); + return sc.seek_key(key) && sc.eat('['); +} + +// Count elements in an array (rough: count top-level '{' or ',' at depth 1). +int count_array_elements(const std::string& s, const char* key) { + Scan sc(s); + if (!sc.seek_key(key)) return -1; + if (!sc.eat('[')) return -1; + sc.ws(); + if (sc.i < s.size() && s[sc.i] == ']') return 0; + int count = 0; + int depth = 0; + while (sc.i < s.size()) { + char c = s[sc.i]; + if (c == '{') { if (depth == 0) ++count; ++depth; } + else if (c == '}') { --depth; } + else if (c == ']' && depth == 0) break; + ++sc.i; + } + return count; +} + +// Parse all "speaker" integer values from the words array. +bool parse_word_speakers(const std::string& s, std::vector& speakers) { + Scan sc(s); + if (!sc.seek_key("words")) return false; + if (!sc.eat('[')) return false; + sc.ws(); + if (sc.i < s.size() && s[sc.i] == ']') return true; // empty + + while (true) { + if (!sc.eat('{')) return false; + // Parse fields until '}' + while (true) { + std::string key; + if (!sc.str(key)) return false; + if (!sc.eat(':')) return false; + if (key == "speaker") { + double v; + if (!sc.num(v)) return false; + speakers.push_back((int)v); + } else { + // Skip value: string or number + std::string tmp; + double d; + if (!sc.str(tmp) && !sc.num(d)) return false; + } + if (sc.eat(',')) continue; + break; + } + if (!sc.eat('}')) return false; + if (sc.eat(',')) continue; + break; + } + return sc.eat(']'); +} + +// Collapse consecutive repeats: [0,0,1,0,0,1] -> [0,1,0,1]. +std::vector turns(const std::vector& spk) { + std::vector t; + for (int s : spk) if (t.empty() || t.back() != s) t.push_back(s); + return t; +} + +std::string show(const std::vector& v) { + std::string s; + for (int x : v) s += (s.empty() ? "" : ",") + std::to_string(x); + return "[" + s + "]"; +} + +int word_count(const char* text) { + int n = 0; + bool in = false; + for (const char* p = text; *p; ++p) { + const bool sp = *p == ' '; + if (!sp && !in) ++n; + in = !sp; + } + return n; +} + +// Offline diarize_pcm segments as (speaker, start, end) triples. +bool parse_segments(const std::string& doc, std::vector>& out) { + Scan sc(doc); + if (!sc.seek_key("segments") || !sc.eat('[')) return false; + if (sc.eat(']')) return true; + do { + std::array seg{}; + if (!sc.eat('{')) return false; + for (int k = 0; k < 3; ++k) { + std::string key; + if (!sc.str(key) || !sc.eat(':') || !sc.num(seg[k])) return false; + if (k < 2 && !sc.eat(',')) return false; + } + if (!sc.eat('}')) return false; + out.push_back(seg); + } while (sc.eat(',')); + return sc.eat(']'); +} + +} // namespace + +#define CHECK(cond, ...) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL: " __VA_ARGS__); \ + std::fprintf(stderr, "\n"); \ + ok = false; \ + } \ + } while (0) + +int main() { + const char* asr_gguf = std::getenv("PARAKEET_TEST_GGUF"); + const char* diar_gguf = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + if (!asr_gguf || !diar_gguf) { + std::fprintf(stderr, "test_combined_offline: PARAKEET_TEST_GGUF and/or " + "PARAKEET_TEST_DIAR_GGUF not set; skip\n"); + return 77; + } + if (parakeet_capi_abi_version() < 7) { + std::fprintf(stderr, "test_combined_offline: ABI < 7\n"); + return 1; + } + parakeet_ctx* asr = parakeet_capi_load(asr_gguf); + parakeet_ctx* diar = parakeet_capi_load(diar_gguf); + if (!asr || !diar) { + std::fprintf(stderr, "test_combined_offline: load failed\n"); + parakeet_capi_free(asr); + parakeet_capi_free(diar); + return 1; + } + pk::Audio audio; + if (!pk::load_audio_16k_mono("tests/fixtures/two_speakers.wav", audio)) { + std::fprintf(stderr, "test_combined_offline: cannot read the fixture\n"); + return 1; + } + const std::vector& pcm = audio.samples; + const int n = (int)pcm.size(); + const std::vector expected_turns = {0, 1, 0, 1}; + bool ok = true; + + // Wrong-model guards. + CHECK(parakeet_capi_diarize_pcm(asr, pcm.data(), n, 16000) == nullptr, + "diarize_pcm accepted an ASR context"); + CHECK(parakeet_capi_transcribe_pcm(diar, pcm.data(), n, 16000, 0) == nullptr, + "transcribe_pcm accepted a diarization context"); + + // 1. JSON variant. + int n_words_offline = 0, n_utts_json = -1; + { + char* json = parakeet_capi_transcribe_and_diarize_json(asr, diar, pcm.data(), n, 16000); + CHECK(json != nullptr, "transcribe_and_diarize_json: %s", parakeet_capi_last_error(asr)); + if (json) { + const std::string doc(json); + parakeet_capi_free_string(json); + std::vector spk; + CHECK(parse_word_speakers(doc, spk), "cannot parse the words array"); + n_words_offline = (int)spk.size(); + n_utts_json = count_array_elements(doc, "utterances"); + int unassigned = 0; + for (int s : spk) unassigned += s < 0; + std::printf("offline SAS: %d words, %d utterances, turns %s, %d unassigned\n", + n_words_offline, n_utts_json, show(turns(spk)).c_str(), unassigned); + CHECK(n_words_offline > 40, "too few words (%d)", n_words_offline); + CHECK(unassigned == 0, "%d words without a speaker", unassigned); + CHECK(turns(spk) == expected_turns, "turns %s, expected [0,1,0,1]", + show(turns(spk)).c_str()); + } + } + + // 2. Struct variant. + { + parakeet_sas_result* r = nullptr; + int nr = 0; + const int rc = parakeet_capi_transcribe_and_diarize(asr, diar, pcm.data(), n, 16000, &r, &nr); + CHECK(rc == 0, "transcribe_and_diarize: %s", parakeet_capi_last_error(asr)); + CHECK(nr == n_utts_json, "struct count %d != JSON count %d", nr, n_utts_json); + for (int i = 0; i < nr; ++i) + CHECK(r[i].text && r[i].start <= r[i].end && r[i].speaker >= 0, + "bad result %d", i); + parakeet_capi_free_sas_results(r, nr); + } + + // 3. Streaming diarization vs offline diarization. + { + char* json = parakeet_capi_diarize_pcm(diar, pcm.data(), n, 16000); + std::vector> offline; + CHECK(json && parse_segments(json, offline), "diarize_pcm"); + parakeet_capi_free_string(json); + + parakeet_diar_stream* ds = parakeet_capi_diarize_stream_begin(diar); + CHECK(ds != nullptr, "diarize_stream_begin: %s", parakeet_capi_last_error(diar)); + std::vector streamed; + for (int lo = 0; ds && lo < n; lo += 8000) { + const int len = std::min(8000, n - lo); + parakeet_diar_segment* segs = nullptr; + int ns = 0; + const int rc = parakeet_capi_diarize_stream_feed(ds, pcm.data() + lo, len, + lo + len >= n, &segs, &ns); + CHECK(rc == 0, "diarize_stream_feed: %s", parakeet_capi_last_error(diar)); + streamed.insert(streamed.end(), segs, segs + ns); + parakeet_capi_free_diar_segments(segs); + } + parakeet_capi_diarize_stream_free(ds); + std::sort(streamed.begin(), streamed.end(), [](const auto& a, const auto& b) { + return a.start != b.start ? a.start < b.start : a.speaker < b.speaker; + }); + std::printf("streaming diarization: %zu segments (offline %zu)\n", + streamed.size(), offline.size()); + CHECK(streamed.size() == offline.size(), "segment count differs"); + for (size_t i = 0; i < std::min(streamed.size(), offline.size()); ++i) { + std::printf(" spk%d %6.2f-%6.2f offline spk%d %6.2f-%6.2f\n", + streamed[i].speaker, streamed[i].start, streamed[i].end, + (int)offline[i][0], offline[i][1], offline[i][2]); + CHECK(streamed[i].speaker == (int)offline[i][0] && + std::fabs(streamed[i].start - offline[i][1]) <= 0.1 && + std::fabs(streamed[i].end - offline[i][2]) <= 0.1, + "segment %zu differs", i); + } + } + + // 4. Streaming SAS. + { + parakeet_sas_stream* ss = parakeet_capi_sas_stream_begin(asr, diar); + CHECK(ss != nullptr, "sas_stream_begin"); + std::vector spk; + int words = 0; + std::string text; + for (int lo = 0; ss && lo < n; lo += 8000) { + const int len = std::min(8000, n - lo); + parakeet_sas_result* r = nullptr; + int nr = 0; + const int rc = parakeet_capi_sas_stream_feed(ss, pcm.data() + lo, len, + lo + len >= n, &r, &nr); + CHECK(rc == 0, "sas_stream_feed: %s", parakeet_capi_last_error(asr)); + for (int i = 0; i < nr; ++i) { + spk.push_back(r[i].speaker); + words += word_count(r[i].text); + text += std::string(text.empty() ? "" : " ") + r[i].text; + } + parakeet_capi_free_sas_results(r, nr); + } + parakeet_capi_sas_stream_free(ss); + std::printf("streaming SAS: %d words, turns %s\n %s\n", words, + show(turns(spk)).c_str(), text.c_str()); + CHECK(turns(spk) == expected_turns, "streaming turns %s", show(turns(spk)).c_str()); + CHECK(std::abs(words - n_words_offline) <= 3, "streaming words %d vs offline %d", + words, n_words_offline); + } + + parakeet_capi_free(asr); + parakeet_capi_free(diar); + std::printf(ok ? "test_combined_offline: PASS\n" : "test_combined_offline: FAIL\n"); + return ok ? 0 : 1; +} diff --git a/tests/test_diarization.cpp b/tests/test_diarization.cpp new file mode 100644 index 0000000..7a3c51a --- /dev/null +++ b/tests/test_diarization.cpp @@ -0,0 +1,86 @@ +#include "diarization.hpp" +#include "model_loader.hpp" +#include +#include +#include + +int main() { + const char* path = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + if (!path) { + std::fprintf(stderr, + "PARAKEET_TEST_DIAR_GGUF not set; skipping diarization test\n"); + return 77; + } + + // Load via DiarizationModel::load (exercises loader, mel, encoder, head). + std::unique_ptr m = pk::DiarizationModel::load(path); + if (!m) { + std::fprintf(stderr, "DiarizationModel::load failed for %s\n", path); + return 1; + } + + // Config sanity: arch must be "diarization" and diarization.present true. + pk::ModelLoader ml; + if (!ml.load(path)) { + std::fprintf(stderr, "ModelLoader::load failed\n"); + return 1; + } + const pk::ParakeetConfig& c = ml.config(); + if (c.arch != "diarization") { + std::fprintf(stderr, "arch != diarization (got %s)\n", c.arch.c_str()); + return 1; + } + if (!c.diarization.present) { + std::fprintf(stderr, "diarization.present is false\n"); + return 1; + } + if (c.diarization.n_speakers == 0) { + std::fprintf(stderr, "n_speakers == 0\n"); + return 1; + } + std::printf("diarization config OK: arch=%s n_spk=%u tf_d_model=%u " + "upsample=%u frame_sec=%.4f onset=%.2f offset=%.2f\n", + c.arch.c_str(), c.diarization.n_speakers, + c.diarization.tf_d_model, c.diarization.upsample_factor, + c.diarization.frame_resolution_sec, + c.diarization.onset_threshold, + c.diarization.offset_threshold); + + // Verify key sortformer tensors are present. + const char* required[] = { + "sortformer_modules.encoder_proj.weight", + "sortformer_modules.subpixel_upsample.weight", + "sortformer_modules.first_hidden_to_hidden.weight", + "sortformer_modules.single_hidden_to_spks.weight", + nullptr, + }; + for (size_t i = 0; required[i]; ++i) { + if (ml.tensor(required[i]) == nullptr) { + std::fprintf(stderr, "missing tensor: %s\n", required[i]); + return 1; + } + } + std::printf("all required sortformer tensors present\n"); + + // If a test audio file is provided, run end-to-end diarization. + const char* wav = std::getenv("PARAKEET_TEST_DIAR_WAV"); + if (wav) { + try { + pk::DiarizationResult r = m->diarize_path(wav); + std::printf("diarized %s -> %zu segments, %d speakers\n", + wav, r.segments.size(), r.n_speakers); + for (size_t i = 0; i < r.segments.size() && i < 20; ++i) { + std::printf(" spk %d: %.2f - %.2f\n", + r.segments[i].speaker, + r.segments[i].start, r.segments[i].end); + } + } catch (const std::exception& e) { + std::fprintf(stderr, "diarize_path threw: %s\n", e.what()); + return 1; + } + } else { + std::printf("PARAKEET_TEST_DIAR_WAV not set; skipping inference test\n"); + } + + return 0; +} diff --git a/tests/test_diarization_accuracy.cpp b/tests/test_diarization_accuracy.cpp new file mode 100644 index 0000000..e5879c2 --- /dev/null +++ b/tests/test_diarization_accuracy.cpp @@ -0,0 +1,158 @@ +// Diarization accuracy vs NeMo (nvidia/Nemotron-3-Diarization). +// +// Runs pk::DiarizationModel on the audio stored in a NeMo baseline +// (scripts/gen_diar_baseline.py) and checks it against NeMo's own output: +// +// 1. offline speaker probabilities: same shape, max/mean abs diff in bounds +// 2. offline segments: same count, same speakers, boundaries within 20 ms, +// and frame-level speaker activity agreement (10 ms grid) >= 99.5% +// 3. the same for diarize_pcm, which follows the model's streaming_mode +// like NeMo's diarize() (streaming for Nemotron-3-Diarization) +// +// The default fixture is tests/fixtures/two_speakers.wav (LibriSpeech 1272 and +// 2086 alternating, A-B-A-B), where NeMo finds 5 segments across 2 speakers. +// +// Env: +// PARAKEET_TEST_DIAR_GGUF diarization GGUF (required) +// PARAKEET_TEST_BASELINE_DIAR baseline GGUF from gen_diar_baseline.py (required) +// PARAKEET_TEST_DIAR_PROB_TOL max abs prob diff (default 0.02; F32 measures +// ~5e-3, quantized models need more headroom) +// Skips (77) when either required variable is unset. +#include "diarization.hpp" +#include "parity.hpp" + +#include +#include +#include +#include +#include + +namespace { + +struct Seg { int spk; float start, end; }; + +std::vector to_segs(const std::vector& flat) { + std::vector out; + for (size_t i = 0; i + 2 < flat.size(); i += 3) + out.push_back({(int)flat[i], flat[i + 1], flat[i + 2]}); + std::sort(out.begin(), out.end(), [](const Seg& a, const Seg& b) { + return a.start != b.start ? a.start < b.start : a.spk < b.spk; + }); + return out; +} + +// Speaker activity on a 10 ms grid: grid[s * T + t]. +std::vector to_grid(const std::vector& segs, int n_spk, int T) { + std::vector g((size_t)n_spk * T, 0); + for (const Seg& s : segs) { + if (s.spk < 0 || s.spk >= n_spk) continue; + const int a = std::max(0, (int)std::lround(s.start * 100.0f)); + const int b = std::min(T, (int)std::lround(s.end * 100.0f)); + for (int t = a; t < b; ++t) g[(size_t)s.spk * T + t] = 1; + } + return g; +} + +int check_segments(const char* label, const std::vector& got, + const std::vector& ref, int n_spk, int T) { + int fails = 0; + std::vector ours; + for (const auto& g : got) ours.push_back({g.speaker, g.start, g.end}); + auto same = [](const Seg& a, const Seg& b) { + return a.spk == b.spk && std::fabs(a.start - b.start) <= 0.02f && + std::fabs(a.end - b.end) <= 0.02f; + }; + std::printf("[%s] segments: ours %zu, NeMo %zu\n", label, ours.size(), ref.size()); + for (size_t i = 0; i < std::max(ours.size(), ref.size()); ++i) { + const bool ho = i < ours.size(), hr = i < ref.size(); + std::printf(" %s spk%d %6.2f-%6.2f NeMo spk%d %6.2f-%6.2f\n", + ho && hr && same(ours[i], ref[i]) ? "ok " : "DIFF", + ho ? ours[i].spk : -1, ho ? ours[i].start : 0.f, ho ? ours[i].end : 0.f, + hr ? ref[i].spk : -1, hr ? ref[i].start : 0.f, hr ? ref[i].end : 0.f); + if (!(ho && hr && same(ours[i], ref[i]))) ++fails; + } + + // Frame-level agreement over frames where either side has speech. + const std::vector go = to_grid(ours, n_spk, T), gr = to_grid(ref, n_spk, T); + int active = 0, agree = 0; + for (int t = 0; t < T; ++t) { + bool any = false, eq = true; + for (int s = 0; s < n_spk; ++s) { + const char a = go[(size_t)s * T + t], b = gr[(size_t)s * T + t]; + any = any || a || b; + eq = eq && a == b; + } + if (any) { ++active; agree += eq ? 1 : 0; } + } + const double agreement = active ? (double)agree / active : 1.0; + std::printf("[%s] frame agreement: %.2f%% of %d active frames\n", label, 100.0 * agreement, active); + if (agreement < 0.995) { + std::fprintf(stderr, "[%s] FAIL: frame agreement below 99.5%%\n", label); + ++fails; + } + if (fails) std::fprintf(stderr, "[%s] FAIL: segments differ from NeMo\n", label); + return fails ? 1 : 0; +} + +} // namespace + +int main() { + const char* gguf = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + const char* base = std::getenv("PARAKEET_TEST_BASELINE_DIAR"); + if (!gguf || !base) { + std::fprintf(stderr, "test_diarization_accuracy: PARAKEET_TEST_DIAR_GGUF and/or " + "PARAKEET_TEST_BASELINE_DIAR not set; skip\n"); + return 77; + } + const char* tol_env = std::getenv("PARAKEET_TEST_DIAR_PROB_TOL"); + const float prob_tol = tol_env ? (float)std::atof(tol_env) : 0.02f; + + auto m = pk::DiarizationModel::load(gguf); + if (!m) { std::fprintf(stderr, "load failed: %s\n", gguf); return 1; } + + std::vector audio, ref_probs, ref_segs_flat; + std::vector shape; + if (!pktest::load_baseline(base, "audio", audio, shape)) return 1; + if (!pktest::load_baseline(base, "offline_probs", ref_probs, shape)) return 1; + const int ref_spk = (int)shape[0], ref_T = (int)shape[1]; + if (!pktest::load_baseline(base, "offline_segs", ref_segs_flat, shape)) return 1; + + int fails = 0; + + // 1. Frame probabilities. + std::vector probs; + int n_spk = 0, T = 0; + m->speaker_probs(audio, probs, n_spk, T); + std::printf("probs: ours [%d, %d], NeMo [%d, %d]\n", n_spk, T, ref_spk, ref_T); + if (n_spk != ref_spk || T != ref_T) { + std::fprintf(stderr, "FAIL: probability shape mismatch\n"); + return 1; + } + double max_diff = 0.0, sum_diff = 0.0; + for (size_t i = 0; i < probs.size(); ++i) { + const double d = std::fabs((double)probs[i] - ref_probs[i]); + max_diff = std::max(max_diff, d); + sum_diff += d; + } + const double mean_diff = sum_diff / probs.size(); + std::printf("probs: max_diff=%.5f mean_diff=%.6f (tol max %.3f, mean 0.002)\n", + max_diff, mean_diff, prob_tol); + if (max_diff > prob_tol || mean_diff > 2e-3) { + std::fprintf(stderr, "FAIL: probabilities diverge from NeMo\n"); + ++fails; + } + + // 2. Offline segments. + fails += check_segments("offline", m->segments_from_probs(probs, n_spk, T), + to_segs(ref_segs_flat), n_spk, T); + + // 3. Default diarize_pcm (the model's streaming_mode, as NeMo diarize()). + const bool streaming = m->config().diarization.streaming_mode; + std::vector ref_default = ref_segs_flat; + if (streaming && !pktest::load_baseline(base, "stream_segs", ref_default, shape)) return 1; + fails += check_segments(streaming ? "diarize_pcm (streaming)" : "diarize_pcm (offline)", + m->diarize_pcm(audio, 16000).segments, to_segs(ref_default), n_spk, T); + + std::printf(fails ? "test_diarization_accuracy: FAIL\n" : "test_diarization_accuracy: PASS\n"); + return fails ? 1 : 0; +} diff --git a/tests/test_model_loader.cpp b/tests/test_model_loader.cpp index 04df63b..d6c5f48 100644 --- a/tests/test_model_loader.cpp +++ b/tests/test_model_loader.cpp @@ -18,8 +18,11 @@ int main() { const pk::ParakeetConfig& c = ml.config(); if (c.arch.empty()) { std::fprintf(stderr, "empty arch\n"); return 1; } if (c.d_model == 0 || c.n_layers == 0 || c.n_heads == 0) { std::fprintf(stderr, "bad encoder dims\n"); return 1; } - if (c.vocab_size == 0) { std::fprintf(stderr, "bad vocab\n"); return 1; } - if (c.blank_id != c.vocab_size) { std::fprintf(stderr, "blank!=vocab\n"); return 1; } + // Diarization models have no vocab (no text output); skip vocab checks. + if (c.arch != "diarization") { + if (c.vocab_size == 0) { std::fprintf(stderr, "bad vocab\n"); return 1; } + if (c.blank_id != c.vocab_size) { std::fprintf(stderr, "blank!=vocab\n"); return 1; } + } // mel filterbank tensor must be present if (ml.tensor("preprocessor.featurizer.fb") == nullptr) { std::fprintf(stderr, "no fb\n"); return 1; } // first conformer layer norm must be present (verbatim name) diff --git a/tests/test_sas_merge.cpp b/tests/test_sas_merge.cpp new file mode 100644 index 0000000..e409593 --- /dev/null +++ b/tests/test_sas_merge.cpp @@ -0,0 +1,248 @@ +// Unit test for the SAS merge layer (merge_asr_diarization + group_speaker_words). +// +// No model or audio needed — constructs synthetic Word and SpeakerSegment vectors +// and verifies the merge + grouping logic directly. + +#include "sas_merge.hpp" +#include "transcription.hpp" +#include "diarization.hpp" + +#include +#include + +using namespace pk; + +static int failures = 0; + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL: %s (line %d)\n", #cond, __LINE__); \ + ++failures; \ + } \ + } while (0) + +// ── Test 1: basic overlap assignment ────────────────────────────────────── +// Two speakers, non-overlapping segments. Words fall clearly within each +// speaker's segment. +static void test_basic_assignment() { + std::vector words = { + {"hello", 0.10f, 0.30f, 0.95f}, + {"world", 0.35f, 0.55f, 0.90f}, + {"foo", 1.10f, 1.30f, 0.80f}, + {"bar", 1.35f, 1.55f, 0.85f}, + }; + std::vector segs = { + {0, 0.00f, 1.00f}, // speaker 0: 0–1s + {1, 1.00f, 2.00f}, // speaker 1: 1–2s + }; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.size() == 4); + CHECK(swords[0].speaker == 0); + CHECK(swords[1].speaker == 0); + CHECK(swords[2].speaker == 1); + CHECK(swords[3].speaker == 1); + + // Verify text and timestamps are preserved + CHECK(swords[0].text == "hello"); + CHECK(swords[0].start == 0.10f); + CHECK(swords[0].end == 0.30f); + CHECK(swords[0].conf == 0.95f); +} + +// ── Test 2: dominant speaker (overlap) ──────────────────────────────────── +// Two overlapping segments: the one with the larger overlap wins. +static void test_dominant_speaker() { + std::vector words = { + {"x", 0.50f, 0.80f, 0.9f}, // overlaps both, but more with spk 0 + }; + std::vector segs = { + {0, 0.00f, 0.70f}, // overlap = 0.20s + {1, 0.60f, 1.00f}, // overlap = 0.20s — equal! first one (sorted) wins + }; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.size() == 1); + // Equal overlap → the first segment (sorted by start, then speaker) wins + // because we use strict > (not >=). + CHECK(swords[0].speaker == 0); + + // Now make speaker 1 overlap more + segs[1].start = 0.40f; // overlap with [0.5, 0.8] = 0.40s > 0.20s + swords = merge_asr_diarization(words, segs); + CHECK(swords[0].speaker == 1); +} + +// ── Test 3: no overlapping segment ──────────────────────────────────────── +// Word falls outside all segments → speaker = -1. +static void test_no_speaker() { + std::vector words = { + {"silence", 5.00f, 5.50f, 0.5f}, + }; + std::vector segs = { + {0, 0.00f, 1.00f}, + {1, 2.00f, 3.00f}, + }; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.size() == 1); + CHECK(swords[0].speaker == -1); +} + +// ── Test 3b: no overlap, but a segment within the snap distance ───────── +// ASR and diarization boundaries can disagree slightly: a word that just +// misses a segment takes the nearest segment's speaker. +static void test_snap_to_nearest() { + std::vector words = { + {"well", 19.92f, 20.00f, 0.7f}, // 0.10 s before spk 1 starts + {"far", 25.00f, 25.20f, 0.7f}, // 1.4 s after spk 1 ends + }; + std::vector segs = { + {0, 14.78f, 18.75f}, // 1.17 s away from "well" + {1, 20.10f, 23.60f}, + }; + auto swords = merge_asr_diarization(words, segs); + CHECK(swords[0].speaker == 1); + CHECK(swords[1].speaker == -1); + // Snapping can be disabled. + swords = merge_asr_diarization(words, segs, 0.0f); + CHECK(swords[0].speaker == -1); +} + +// ── Test 4: utterance grouping — same speaker, small gap ────────────────── +static void test_grouping_same_speaker() { + std::vector swords = { + {0, "hello", 0.10f, 0.30f, 0.95f}, + {0, "world", 0.40f, 0.60f, 0.90f}, // gap = 0.10s ≤ 0.5s + }; + + auto utts = group_speaker_words(swords); + CHECK(utts.size() == 1); + CHECK(utts[0].speaker == 0); + CHECK(utts[0].text == "hello world"); + CHECK(utts[0].start == 0.10f); + CHECK(utts[0].end == 0.60f); + CHECK(utts[0].conf == 0.90f); // min(0.95, 0.90) +} + +// ── Test 5: utterance grouping — speaker change ─────────────────────────── +static void test_grouping_speaker_change() { + std::vector swords = { + {0, "hello", 0.10f, 0.30f, 0.95f}, + {1, "world", 0.40f, 0.60f, 0.90f}, + }; + + auto utts = group_speaker_words(swords); + CHECK(utts.size() == 2); + CHECK(utts[0].speaker == 0); + CHECK(utts[0].text == "hello"); + CHECK(utts[1].speaker == 1); + CHECK(utts[1].text == "world"); +} + +// ── Test 6: utterance grouping — large gap splits ───────────────────────── +static void test_grouping_large_gap() { + std::vector swords = { + {0, "hello", 0.10f, 0.30f, 0.95f}, + {0, "world", 1.00f, 1.20f, 0.90f}, // gap = 0.70s > 0.5s + }; + + auto utts = group_speaker_words(swords); + CHECK(utts.size() == 2); + CHECK(utts[0].text == "hello"); + CHECK(utts[1].text == "world"); +} + +// ── Test 7: utterance grouping — unknown speaker (-1) ───────────────────── +static void test_grouping_unknown_speaker() { + std::vector swords = { + {0, "a", 0.10f, 0.20f, 0.9f}, + {-1, "b", 0.30f, 0.40f, 0.8f}, + {-1, "c", 0.45f, 0.55f, 0.7f}, + {1, "d", 0.60f, 0.70f, 0.85f}, + }; + + auto utts = group_speaker_words(swords); + CHECK(utts.size() == 3); + CHECK(utts[0].speaker == 0); + CHECK(utts[0].text == "a"); + CHECK(utts[1].speaker == -1); + CHECK(utts[1].text == "b c"); + CHECK(utts[2].speaker == 1); + CHECK(utts[2].text == "d"); +} + +// ── Test 8: empty inputs ────────────────────────────────────────────────── +static void test_empty() { + std::vector words; + std::vector segs; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.empty()); + + auto utts = group_speaker_words(swords); + CHECK(utts.empty()); +} + +// ── Test 9: word exactly at segment boundary ─────────────────────────────── +static void test_boundary() { + std::vector words = { + {"x", 1.00f, 1.10f, 0.9f}, // starts exactly at seg0 end / seg1 start + }; + std::vector segs = { + {0, 0.00f, 1.00f}, // seg.end <= w.start → skipped (<=) + {1, 1.00f, 2.00f}, // overlap = 0.10s + }; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.size() == 1); + CHECK(swords[0].speaker == 1); +} + +// ── Test 10: multiple speakers on same segment ──────────────────────────── +// Same segment list, multiple words — all should get the same speaker. +static void test_multiple_words_same_speaker() { + std::vector words = { + {"a", 0.10f, 0.20f, 0.9f}, + {"b", 0.25f, 0.35f, 0.8f}, + {"c", 0.40f, 0.50f, 0.7f}, + }; + std::vector segs = { + {2, 0.00f, 1.00f}, + }; + + auto swords = merge_asr_diarization(words, segs); + CHECK(swords.size() == 3); + for (int i = 0; i < 3; ++i) { + CHECK(swords[i].speaker == 2); + } + + // Group should merge all into one utterance + auto utts = group_speaker_words(swords); + CHECK(utts.size() == 1); + CHECK(utts[0].text == "a b c"); + CHECK(utts[0].speaker == 2); + CHECK(utts[0].conf == 0.7f); // min +} + +int main() { + test_basic_assignment(); + test_dominant_speaker(); + test_no_speaker(); + test_snap_to_nearest(); + test_grouping_same_speaker(); + test_grouping_speaker_change(); + test_grouping_large_gap(); + test_grouping_unknown_speaker(); + test_empty(); + test_boundary(); + test_multiple_words_same_speaker(); + + if (failures == 0) { + std::printf("All SAS merge tests passed.\n"); + return 0; + } + std::printf("%d assertion(s) failed.\n", failures); + return 1; +} diff --git a/tests/test_streaming_diarization.cpp b/tests/test_streaming_diarization.cpp new file mode 100644 index 0000000..bff39c9 --- /dev/null +++ b/tests/test_streaming_diarization.cpp @@ -0,0 +1,175 @@ +// Streaming diarization accuracy vs NeMo cache-aware streaming +// (SortformerEncLabelModel with streaming_mode=True, the model's own chunk and +// speaker-cache config). +// +// Feeds the baseline clip through pk::StreamingDiarization two ways and checks +// both against NeMo's streaming output: +// A. mel chunks cut from the whole-clip mel (exactly what NeMo does) +// B. pk::StreamingMel fed with small PCM pieces (the live-audio path the +// C-API uses), re-chunked to the model's chunk size +// Checks: per-frame probabilities (max/mean abs diff) and the segments +// (count, speaker, boundaries within 20 ms). +// +// Env: PARAKEET_TEST_DIAR_GGUF + PARAKEET_TEST_BASELINE_DIAR +// (scripts/gen_diar_baseline.py); PARAKEET_TEST_DIAR_PROB_TOL as in +// test_diarization_accuracy. Skips (77) when unset. +#include "diarization.hpp" +#include "diarization_streaming.hpp" +#include "mel.hpp" +#include "parity.hpp" + +#include +#include +#include +#include +#include + +namespace { + +struct Run { + std::vector probs; // [n_spk, T] + std::vector segs; +}; + +// Feed [n_mels, T] mel through the streaming diarizer in model-sized chunks. +Run stream_mel(pk::StreamingDiarization& sd, const std::vector& mel, int n_mels, int T) { + Run r; + const int ns = sd.n_speakers(), cm = sd.chunk_mel_frames(); + r.probs.assign((size_t)ns * T, 0.0f); + sd.reset(); + for (int lo = 0; lo < T; lo += cm) { + const int n = std::min(cm, T - lo); + std::vector chunk((size_t)n_mels * n); + for (int m = 0; m < n_mels; ++m) + std::copy_n(mel.begin() + (size_t)m * T + lo, n, chunk.begin() + (size_t)m * n); + auto segs = sd.feed_mel_chunk(chunk, n_mels, n, lo + n >= T); + r.segs.insert(r.segs.end(), segs.begin(), segs.end()); + for (int s = 0; s < ns; ++s) + std::copy_n(sd.last_chunk_probs().begin() + (size_t)s * n, n, + r.probs.begin() + (size_t)s * T + lo); + } + std::sort(r.segs.begin(), r.segs.end(), [](const auto& a, const auto& b) { + return a.start != b.start ? a.start < b.start : a.speaker < b.speaker; + }); + return r; +} + +int check(const char* label, const Run& r, const std::vector& ref_probs, + const std::vector& ref_segs, float tol) { + int fails = 0; + double max_d = 0.0, sum_d = 0.0; + for (size_t i = 0; i < ref_probs.size(); ++i) { + const double d = std::fabs((double)r.probs[i] - ref_probs[i]); + max_d = std::max(max_d, d); + sum_d += d; + } + const double mean_d = sum_d / ref_probs.size(); + if (std::getenv("PARAKEET_TEST_DIAR_VERBOSE")) { + // Per-10 s max diff, to localize divergence (e.g. after cache compression). + const size_t T = ref_probs.size() / 8; + for (size_t lo = 0; lo < T; lo += 1000) { + double m = 0.0; + for (size_t sp = 0; sp < 8; ++sp) + for (size_t t = lo; t < std::min(T, lo + 1000); ++t) + m = std::max(m, std::fabs((double)r.probs[sp * T + t] - ref_probs[sp * T + t])); + std::printf(" frames %zu..%zu max_diff %.4f\n", lo, std::min(T, lo + 1000), m); + } + } + std::printf("[%s] probs: max_diff=%.5f mean_diff=%.6f\n", label, max_d, mean_d); + if (max_d > tol || mean_d > 2e-3) { std::fprintf(stderr, "[%s] FAIL: probabilities\n", label); ++fails; } + + const size_t n_ref = ref_segs.size() / 3; + std::printf("[%s] segments: ours %zu, NeMo %zu\n", label, r.segs.size(), n_ref); + if (r.segs.size() != n_ref) { + std::fprintf(stderr, "[%s] FAIL: segment count\n", label); + return fails + 1; + } + // NeMo's rows are sorted by (start, speaker) the same way. + std::vector order(n_ref); + for (size_t i = 0; i < n_ref; ++i) order[i] = i; + std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { + return ref_segs[a * 3 + 1] != ref_segs[b * 3 + 1] ? ref_segs[a * 3 + 1] < ref_segs[b * 3 + 1] + : ref_segs[a * 3] < ref_segs[b * 3]; + }); + for (size_t i = 0; i < n_ref; ++i) { + const float* g = &ref_segs[order[i] * 3]; + const auto& o = r.segs[i]; + const bool ok = o.speaker == (int)g[0] && std::fabs(o.start - g[1]) <= 0.02f && + std::fabs(o.end - g[2]) <= 0.02f; + std::printf(" %s spk%d %6.2f-%6.2f NeMo spk%d %6.2f-%6.2f\n", ok ? "ok " : "DIFF", + o.speaker, o.start, o.end, (int)g[0], g[1], g[2]); + if (!ok) ++fails; + } + return fails; +} + +} // namespace + +int main() { + const char* gguf = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + const char* base = std::getenv("PARAKEET_TEST_BASELINE_DIAR"); + if (!gguf || !base) { + std::fprintf(stderr, "test_streaming_diarization: PARAKEET_TEST_DIAR_GGUF and/or " + "PARAKEET_TEST_BASELINE_DIAR not set; skip\n"); + return 77; + } + const char* tol_env = std::getenv("PARAKEET_TEST_DIAR_PROB_TOL"); + const float tol = tol_env ? (float)std::atof(tol_env) : 0.02f; + + auto m = pk::DiarizationModel::load(gguf); + if (!m) { std::fprintf(stderr, "load failed: %s\n", gguf); return 1; } + + std::vector audio, ref_probs, ref_segs; + std::vector shape; + if (!pktest::load_baseline(base, "audio", audio, shape)) return 1; + if (!pktest::load_baseline(base, "stream_probs", ref_probs, shape)) return 1; + const int T = (int)shape[1]; + if (!pktest::load_baseline(base, "stream_segs", ref_segs, shape)) return 1; + + pk::StreamingDiarization sd(m->loader()); + const int n_mels = sd.n_mels(); + int fails = 0; + + // A. Whole-clip mel (no peak normalization in streaming mode), trimmed to + // floor(S / hop) frames like NeMo. + { + std::vector mel; + int nm = 0, Tm = 0; + m->mel().compute(audio, mel, nm, Tm); + if (Tm < T) { std::fprintf(stderr, "mel too short: %d < %d\n", Tm, T); return 1; } + std::vector trimmed((size_t)nm * T); + for (int i = 0; i < nm; ++i) + std::copy_n(mel.begin() + (size_t)i * Tm, T, trimmed.begin() + (size_t)i * T); + fails += check("whole-clip mel", stream_mel(sd, trimmed, nm, T), ref_probs, ref_segs, tol); + } + + // B. Incremental mel from 100 ms PCM pieces. + { + pk::StreamingMel sm(m->loader()); + std::vector> cols; // per-frame mel columns + auto take = [&](const std::vector& fm, int n) { + for (int t = 0; t < n; ++t) { + std::vector col(n_mels); + for (int i = 0; i < n_mels; ++i) col[i] = fm[(size_t)i * n + t]; + cols.push_back(std::move(col)); + } + }; + for (size_t lo = 0; lo < audio.size(); lo += 1600) { + const int n = (int)std::min(1600, audio.size() - lo); + int nf = 0; + auto fm = sm.feed(audio.data() + lo, n, nf); + take(fm, nf); + } + int nf = 0; + auto tail = sm.finalize(nf); + take(tail, nf); + if ((int)cols.size() < T) { std::fprintf(stderr, "stream mel too short\n"); return 1; } + std::vector mel((size_t)n_mels * T); + for (int t = 0; t < T; ++t) + for (int i = 0; i < n_mels; ++i) mel[(size_t)i * T + t] = cols[t][i]; + fails += check("StreamingMel", stream_mel(sd, mel, n_mels, T), ref_probs, ref_segs, tol); + } + + std::printf(fails ? "test_streaming_diarization: FAIL\n" : "test_streaming_diarization: PASS\n"); + return fails ? 1 : 0; +}