diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 6a5c8c9..3c9d1bd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -32,11 +32,12 @@ jobs: # job once a models bundle is published (Phase 4). run: ctest --test-dir build --output-on-failure -LE model - - name: build without ced (PARAKEET_WITH_CED=OFF) - # PARAKEET_WITH_CED is on by default, so this is the gate that catches - # anything that quietly starts depending on ced.cpp being present. + - name: build without ced and voice-detect (PARAKEET_WITH_CED=OFF, PARAKEET_WITH_VOICEDETECT=OFF) + # PARAKEET_WITH_CED and PARAKEET_WITH_VOICEDETECT are on by default, so + # this is the gate that catches anything that quietly starts depending + # on ced.cpp or voice-detect.cpp being present. run: | - cmake -B build-noced -DPARAKEET_BUILD_TESTS=ON -DGGML_NATIVE=OFF -DPARAKEET_WITH_CED=OFF + cmake -B build-noced -DPARAKEET_BUILD_TESTS=ON -DGGML_NATIVE=OFF -DPARAKEET_WITH_CED=OFF -DPARAKEET_WITH_VOICEDETECT=OFF cmake --build build-noced -j ctest --test-dir build-noced --output-on-failure -LE model # scene --sound must fail cleanly (exit 2) with a clear message. @@ -46,6 +47,12 @@ jobs: echo "$out" test "$rc" -eq 2 grep -q "built without sound tagging" <<< "$out" + # enroll must fail the same way without voice-detect.cpp. + rc=0 + out=$(build-noced/examples/cli/parakeet-cli enroll --model x.gguf --name a --input x.wav --registry /tmp/r.bin 2>&1) || rc=$? + echo "$out" + test "$rc" -eq 2 + grep -q "built without speaker identification" <<< "$out" # ------------------------------------------------------------------------- # server-e2e: drive the real parakeet-server over HTTP. diff --git a/.gitmodules b/.gitmodules index 07e9117..ee760f2 100644 --- a/.gitmodules +++ b/.gitmodules @@ -4,3 +4,6 @@ [submodule "third_party/ced.cpp"] path = third_party/ced.cpp url = https://github.com/localai-org/ced.cpp +[submodule "third_party/voice-detect.cpp"] + path = third_party/voice-detect.cpp + url = https://github.com/localai-org/voice-detect.cpp diff --git a/AGENTS.md b/AGENTS.md index c650ce3..b405dec 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -85,6 +85,9 @@ src/ libparakeet implementation sound_stream.hpp/cpp, pk::SoundStream: sliding-window sound-event detection over live PCM scene_stream.hpp/cpp, pk::SceneStream: combined ASR + diarization + sound-event stream scene_render.hpp/cpp, pk::SceneRenderer + format_span/is_speech_label: `parakeet-cli scene` text rendering + speaker_registry.hpp/cpp, pk::SpeakerRegistry: enrolled voices (centroid per name), match, binary save/load + speaker_identifier.hpp/cpp, pk::SpeakerIdentifier: names diarization slots from their clean audio; identify_offline + speaker_encoder.hpp/cpp, pk::SpeakerEncoder: the only code that talks to voice-detect.cpp (voicedetect_capi.h) examples/cli/ parakeet-cli binary subcommands: info, transcribe (+ --stream), quantize, scene (ASR + diar + sound, one time-ordered feed) sound-window-eval: measures CED short-window accuracy vs whole-clip top-1 @@ -123,6 +126,11 @@ tests/ ctest targets test_sound_capi.cpp , sound_stream_* C-API (PARAKEET_TEST_CED_GGUF) test_scene_stream.cpp , pk::SceneStream / scene_stream_* C-API, all three models together (PARAKEET_TEST_GGUF + PARAKEET_TEST_DIAR_GGUF + PARAKEET_TEST_CED_GGUF) test_scene_render.cpp , SceneRenderer / format_span / is_speech_label (model-independent) + test_speaker_registry.cpp, SpeakerRegistry enroll/match/serialize (model-independent) + test_speaker_identifier.cpp, SpeakerIdentifier with a fake embedder (model-independent) + test_speaker_encoder.cpp, SpeakerEncoder vs voice-detect reference embedding (PARAKEET_TEST_VD_GGUF + PARAKEET_TEST_VD_REF_WAV + PARAKEET_TEST_VD_REF_JSON) + test_speaker_identify.cpp, scene stream names both fixture voices (PARAKEET_TEST_DIAR_GGUF + PARAKEET_TEST_VD_GGUF; PARAKEET_TEST_GGUF adds the named-utterance block) + test_capi_speaker.cpp , speaker C-API v9 (PARAKEET_TEST_DIAR_GGUF + PARAKEET_TEST_VD_GGUF; PARAKEET_TEST_GGUF optional; PARAKEET_TEST_VD_GGUF_ALT for the size-mismatch check) 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 @@ -134,6 +142,9 @@ third_party/ vendored deps built as a static `ced` target linked into libparakeet, not a separate process; dr_wav is shared via CED_EXTERNAL_DR_WAV so there is one DR_WAV_IMPLEMENTATION in the whole build + voice-detect.cpp/, submodule, speaker encoders (PARAKEET_WITH_VOICEDETECT, on by default); + static `voicedetect` target linked into libparakeet, dr_wav shared via + VOICEDETECT_EXTERNAL_DR_WAV dr_wav.h , vendored single header models/ output dir for converted GGUFs (gitignored; MANIFEST.md tracks the expected published set) @@ -142,6 +153,8 @@ docs/ 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 + sound.md , sound-event detection (CED) and the combined scene stream + speaker.md , speaker identification: enroll, scene naming, C-API v9, measured numbers .github/workflows/ ci.yml , build job (per-push) + closed-loop job (pull_request + dispatch) ``` @@ -164,6 +177,7 @@ cmake -B build -DPARAKEET_BUILD_TESTS=ON -DGGML_NATIVE=ON && cmake --build build | `PARAKEET_GGML_VULKAN` | OFF | Forward GGML_VULKAN to the submodule | | `PARAKEET_GGML_HIPBLAS` | OFF | Forward GGML_HIPBLAS to the submodule | | `PARAKEET_WITH_CED` | ON | Sound-event detection through ced.cpp | +| `PARAKEET_WITH_VOICEDETECT` | ON | Speaker identification through voice-detect.cpp | Use `-DGGML_NATIVE=OFF` when building for CI or portable binaries. @@ -256,6 +270,8 @@ parakeet-cli info parakeet-cli transcribe --model --input [--decoder ctc|tdt] [--stream] [--timestamps] [--json] parakeet-cli quantize parakeet-cli scene [--model ] [--diar ] [--sound ] --input [--latency model|low|very_low|ultra_low] [--chunk-ms N] [--show-speech] [--json] +parakeet-cli scene ... --speakers --registry [--speaker-threshold F] # names diarized speakers +parakeet-cli enroll --model --name --input [--input ...] --registry ``` `--timestamps` prints one `- ()` line per word (also @@ -310,7 +326,22 @@ parakeet_capi_sound_stream_begin / _feed / _active / _drain_scores_json / _free parakeet_capi_free_sound_segments parakeet_capi_num_classes parakeet_capi_class_label -parakeet_capi_model_kind # which kind of ctx (NONE/ASR/DIARIZATION/SOUND) +parakeet_capi_model_kind # which kind of ctx (NONE/ASR/DIARIZATION/SOUND/SPEAKER) +``` + +Speaker identification (ABI v9, additive; not used by LocalAI yet). A +voice-detect.cpp speaker GGUF loads into a fourth `parakeet_ctx` +kind (`PARAKEET_MODEL_KIND_SPEAKER`, 4) through the same `parakeet_capi_load`; +see `docs/speaker.md`: + +``` +parakeet_capi_speaker_dim +parakeet_capi_speaker_registry_new / _free / _size / _last_error +parakeet_capi_speaker_enroll +parakeet_capi_speaker_registry_save / _load +parakeet_capi_speaker_identify_pcm_json +parakeet_capi_scene_stream_begin_speaker +parakeet_capi_transcribe_and_diarize_named_json ``` Combined scene stream (ABI v8, additive; not used by LocalAI yet). One stream @@ -454,7 +485,9 @@ See `docs/conversion.md` for the authoritative schema. Quick summary: ## ggml submodule -Pinned at v0.13.0 in `third_party/ggml`. No local patches. To bump: +Pinned at v0.13.0 in `third_party/ggml`. CMake applies the patches in +`third_party/ggml-patches` in-tree at configure time (`scripts/apply_ggml_patches.sh`), +so the submodule shows as modified. To bump: 1. Update the submodule SHA. 2. Run `ctest --test-dir build --output-on-failure`. 3. Fix any API breakage in `src/model_loader.cpp`. diff --git a/CMakeLists.txt b/CMakeLists.txt index cbaa2db..6d0e22c 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -18,6 +18,7 @@ option(PARAKEET_GGML_METAL "Forward GGML_METAL" OFF) option(PARAKEET_GGML_VULKAN "Forward GGML_VULKAN" OFF) option(PARAKEET_GGML_HIP "Forward GGML_HIP (ROCm)" OFF) option(PARAKEET_WITH_CED "Sound-event detection through ced.cpp" ON) +option(PARAKEET_WITH_VOICEDETECT "Speaker identification through voice-detect.cpp" ON) set(GGML_CUDA ${PARAKEET_GGML_CUDA} CACHE BOOL "" FORCE) set(GGML_METAL ${PARAKEET_GGML_METAL} CACHE BOOL "" FORCE) @@ -91,6 +92,19 @@ if(PARAKEET_WITH_CED) target_link_libraries(ced PRIVATE dr_wav_impl) endif() +if(PARAKEET_WITH_VOICEDETECT) + if(NOT EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/third_party/voice-detect.cpp/CMakeLists.txt") + message(FATAL_ERROR "third_party/voice-detect.cpp is missing: run `git submodule update --init third_party/voice-detect.cpp`, or configure with -DPARAKEET_WITH_VOICEDETECT=OFF") + endif() + set(VOICEDETECT_BUILD_CLI OFF CACHE BOOL "" FORCE) + set(VOICEDETECT_BUILD_TESTS OFF CACHE BOOL "" FORCE) + set(VOICEDETECT_SHARED OFF CACHE BOOL "" FORCE) + set(VOICEDETECT_EXTERNAL_DR_WAV ON CACHE BOOL "" FORCE) # dr_wav_impl provides it + add_subdirectory(third_party/voice-detect.cpp EXCLUDE_FROM_ALL) + set_target_properties(voicedetect PROPERTIES POSITION_INDEPENDENT_CODE ON) + target_link_libraries(voicedetect PRIVATE dr_wav_impl) +endif() + set(PARAKEET_SRC src/parakeet.cpp src/model.cpp @@ -131,6 +145,9 @@ set(PARAKEET_SRC src/diarization_streaming.cpp src/ced_tagger.cpp src/sound_stream.cpp + src/speaker_registry.cpp + src/speaker_identifier.cpp + src/speaker_encoder.cpp src/scene_render.cpp) if(PARAKEET_SHARED) @@ -154,6 +171,11 @@ if(PARAKEET_WITH_CED) target_compile_definitions(parakeet PRIVATE PARAKEET_WITH_CED=1) endif() +if(PARAKEET_WITH_VOICEDETECT) + target_link_libraries(parakeet PRIVATE voicedetect) + target_compile_definitions(parakeet PRIVATE PARAKEET_WITH_VOICEDETECT=1) +endif() + if(PARAKEET_BUILD_CLI) add_subdirectory(examples/cli) endif() diff --git a/README.md b/README.md index 0a3503e..46c881c 100644 --- a/README.md +++ b/README.md @@ -117,6 +117,7 @@ cmake --build build-shared -j | `PARAKEET_GGML_VULKAN` | OFF | Forward GGML_VULKAN to the submodule | | `PARAKEET_GGML_HIP` | OFF | Forward GGML_HIP (ROCm) to the submodule | | `PARAKEET_WITH_CED` | ON | Sound-event detection through ced.cpp | +| `PARAKEET_WITH_VOICEDETECT` | ON | Speaker identification through voice-detect.cpp | To build for a GPU backend, forward its flag, e.g. Apple Metal: @@ -337,6 +338,15 @@ parakeet-cli scene --model asr.gguf --diar diar.gguf --sound ced-base-q8_0.gguf See [`docs/sound.md`](docs/sound.md) for the CED GGUFs, the sound and scene stream C-API (ABI v8), and the `--sound-model` server option. +### Naming speakers + +With a voice-detect.cpp speaker encoder (`PARAKEET_WITH_VOICEDETECT`, on by +default) the scene stream can say who is talking instead of `Speaker 0`. +Enroll each person from a short clip with `parakeet-cli enroll`, then pass +`--speakers --registry ` to `scene`. Only one two-voice +fixture has been measured so far. See [`docs/speaker.md`](docs/speaker.md) for +the models, the commands, the C-API (ABI v9) and what is still untested. + --- ## C-API (`libparakeet.so`) diff --git a/docs/speaker.md b/docs/speaker.md new file mode 100644 index 0000000..566ae0e --- /dev/null +++ b/docs/speaker.md @@ -0,0 +1,272 @@ +# Speaker identification + +parakeet.cpp can put a name on a diarized speaker. You enroll a few people +from short clips, and the scene stream and the speaker-attributed ASR output +then say `Ada:` where they would otherwise say `Speaker 0:`. + +It runs a speaker-embedding model from +[voice-detect.cpp](https://github.com/localai-org/voice-detect.cpp), built in as a +static library (`PARAKEET_WITH_VOICEDETECT`, on by default, the same way +ced.cpp is built in for sound events). `pk::SpeakerEncoder` +(`src/speaker_encoder.hpp`) is the only parakeet code that talks to it. + +## What it does and does not do + +It does: + +- Turn a clip of speech into an L2-normalized embedding, and keep enrolled + voices in a registry (one centroid per name). +- Give each diarization slot a name by embedding the slot's clean audio and + matching it against the registry. +- Leave a slot unnamed when nobody in the registry matches well enough. + +It does not: + +- Find or diarize speakers by itself. It names the slots that the diarization + model (Sortformer) produces, so it needs `--diar`. +- Resolve overlapped speech. Time where two speakers overlap is skipped when a + slot's voice is built. It is not attributed to anyone. +- Learn new voices on the fly. The registry is only changed by enrolling. + +## Which GGUFs work + +Use a speaker-encoder GGUF from +[`mudler/voice-detect-gguf`](https://huggingface.co/mudler/voice-detect-gguf). +The four speaker encoders are: + +| Model | Embedding size | f32 GGUF size | Starting threshold | +| --- | --- | --- | --- | +| WeSpeaker ResNet34 | 256 | 26.5 MB | 0.5 | +| CAM++ (3D-Speaker, zh-cn) | 192 | 27.7 MB | 0.5 | +| ECAPA-TDNN (SpeechBrain, VoxCeleb) | 192 | 83.2 MB | 0.7 | +| ERes2Net (3D-Speaker, base) | 512 | 39.5 MB | not measured | + +Start with WeSpeaker ResNet34: it kept the two voices furthest apart in the +measurements below. The starting threshold is the `accept_threshold` to begin +with (`--speaker-threshold` on the command line). The default is 0.5, which is +right for WeSpeaker and CAM++ here, but ECAPA scored a voice that was not +enrolled at 0.566, so it needs about 0.7 (0.13 above that impostor and 0.26 +below the lowest genuine ECAPA score). These numbers come from one fixture, +where the enrollment clips and the test audio share a recording and genuine +scores were 0.92 to 0.98. Expect lower genuine scores when enrollment and test +audio come from different sessions or microphones, and check the threshold on +your own audio. + +The repository also holds age, gender and emotion models. They are not +speaker encoders and cannot be used here (`SpeakerEncoder::load` returns null +for a GGUF with no speaker embedding). f16 and q8_0 files are published as +well. Sizes are for the f32 files. ERes2Net has not been run through any of +the tests here. + +A registry belongs to the encoder that made it. The embedding sizes differ, and +even two encoders with the same size do not share a space, so enroll again if +you switch models. `scene` checks the size and stops if it does not match; it +cannot tell two encoders of the same size apart. + +The speaker-model weights have their own licences (WeSpeaker, 3D-Speaker and +SpeechBrain each publish theirs). voice-detect.cpp's own licence does +not cover the weights. Read the licence of the checkpoint you +ship. + +## Enroll + +``` +parakeet-cli enroll --model --name \ + --input [--input ...] --registry +``` + +Each `--input` is one clip. The registry file is created when it is missing +and added to when it exists. Nothing is written unless every clip embedded. +The number printed is the clips enrolled by that command, not the total for +that name. + +The example below uses three clips cut out of +`tests/fixtures/two_speakers.wav` (voice A at 0.6 to 4.6 s and 14.9 to 18.5 s, +voice B at 6.9 to 10.9 s), enrolled with WeSpeaker ResNet34. Real output: + +``` +$ parakeet-cli enroll --model wespeaker_resnet34_f32.gguf --name Ada --input a.wav --registry reg.bin +enrolled Ada (1 clip(s)), registry has 1 speaker(s) +$ parakeet-cli enroll --model wespeaker_resnet34_f32.gguf --name Ben --input b.wav --registry reg.bin +enrolled Ben (1 clip(s)), registry has 2 speaker(s) +$ parakeet-cli enroll --model wespeaker_resnet34_f32.gguf --name Ada --input a2.wav --input a.wav --registry reg.bin +enrolled Ada (2 clip(s)), registry has 2 speaker(s) +``` + +Enrolling a name again refines that voice (the centroid moves) and does not +add a second speaker. The same voice enrolled under two names comes out +unknown: both names match about equally well, so neither beats the other by +the margin. Names are compared exactly, so near-duplicate names (`Ada` and +`ada`, or a trailing space) count as two speakers. + +## Scene with names + +``` +parakeet-cli scene --model --diar \ + --speakers --registry [--speaker-threshold F] \ + --input +``` + +`--speakers` needs `--diar` and `--registry`. Real output on the same fixture +(110m TDT, Sortformer, WeSpeaker; the enrollment clips come from the same +recording): + +``` +[00:00.4 - 00:05.4] Ada: mister Quilter is the apostle of the middle classes, and we are glad to welcome his gospel. +[00:06.8 - 00:10.8] Ben: Well, I don't wish to see it any more, observed Phoebe, turning away her eyes. +[00:11.4 - 00:13.6] Ben: It is certainly very like the old portrait. +[00:14.8 - 00:18.5] Ada: Nor is mister Quilter's manner less interesting than his matter. +[00:19.9 - 00:20.0] Ben: Well, +[00:20.4 - 00:23.3] Ben: I don't wish to see it any more, observed Phoebe, turning away her +``` + +With `--json` each update carries a `"names"` map (empty, `{}`, until a slot +is seen), for example at the end of the file: + +``` +"names":{"0":{"name":"Ada","score":0.9752},"1":{"name":"Ben","score":0.9681}} +``` + +The scores in this sample come from enrolling with the whole clips used in the +example above (Ada from `a.wav` and `a2.wav`, Ben from `b.wav`), so they differ +a little from the numbers in the measured section, where each voice is enrolled +from one clip. + +An unnamed slot still renders as `Speaker N:`. + +Errors exit with a one-line message: 2 for a usage problem (`--speakers` +without `--diar` or without `--registry`, a bad `--speaker-threshold`), 1 for a +runtime problem (missing or invalid registry file, registry from a model with a +different embedding size, model that fails to load). + +## The timing rule + +A slot needs some clean audio before it can be named (2 s by default). That +audio is collected while the slot is still talking, not only after it pauses, +so a speaker who talks without a break is named during that first turn. Still, +a word can be committed before its slot is identified. Such a word keeps the +label it had when it was committed (empty name, rendered as `Speaker N`), and +it is not rewritten later. The `names` map in each update, and `active`, carry +the current identity of each slot. In the run above every utterance was named +from its first word, but that is one fixture and it depends on how early the +speakers start talking. + +## Defaults + +| Option | Default | Meaning | +| --- | --- | --- | +| `min_voice_sec` | 2.0 | clean audio a slot needs before it is embedded | +| `refresh_sec` | 3.0 | new clean audio that triggers another embedding | +| `max_voice_sec` | 10.0 | the newest audio kept per slot for embedding | +| `accept_threshold` | 0.5 | minimum cosine to take a name | +| `margin` | 0.05 | best match must beat the runner-up by this much | + +Change them in C++ through `SceneParts::speaker_opts` (`pk::SpeakerIdOpts`), in +the C-API through the `speaker_*` fields of `parakeet_scene_opts`, and on the +command line with `--speaker-threshold` (only `accept_threshold`). + +`accept_threshold` is a starting point, not a tuned value. It depends on the +encoder: see the starting threshold column in "Which GGUFs work" and the +numbers below. + +## C-API (ABI v9) + +Additive: no earlier signature changed, LocalAI does not use these yet. A +speaker GGUF loads through `parakeet_capi_load` into a context of kind +`PARAKEET_MODEL_KIND_SPEAKER` (4, from `parakeet_capi_model_kind`). + +``` +parakeet_capi_speaker_dim # embedding size, -1 if not a speaker ctx +parakeet_capi_speaker_registry_new / _free / _size / _last_error +parakeet_capi_speaker_enroll # embed PCM and add it under a name +parakeet_capi_speaker_registry_save / _load # binary file +parakeet_capi_speaker_identify_pcm_json # {"name":"alice","score":0.71} +parakeet_capi_scene_stream_begin_speaker # scene stream with a speaker ctx + registry +parakeet_capi_transcribe_and_diarize_named_json # offline speaker-attributed ASR with names +``` + +`parakeet_capi_scene_stream_begin_speaker` takes the same arguments as +`parakeet_capi_scene_stream_begin` plus a speaker ctx and a registry. The +registry is borrowed: keep it alive and unchanged while the stream runs. + +JSON fields: each utterance, word and speaker segment gets `"name"` and +`"name_score"` (empty name and 0.0000 mean unknown), and the top level gets +`"names"`, a map from slot to `{"name","score"}`. In a scene stream with a +speaker part these fields are there from the first document on: `"names"` is +`{}` until diarization has seen a slot, so the shape of the document does not +change during the stream. Without a speaker part none of them appear. The +offline named document is the SAS document plus those fields. + +## Devices and threads + +voice-detect keeps its own backend and device selection, like ced.cpp does with +`CED_DEVICE`: `VOICEDETECT_DEVICE` picks the device and `VOICEDETECT_THREADS` +the CPU thread count. They are read separately from `PARAKEET_DEVICE`. An +embedded voice-detect build does not apply voice-detect's own CUDA and cuDNN +ggml patch. That does not matter on CPU. + +## The registry file + +A small binary blob (version 1) written by `enroll` and by +`parakeet_capi_speaker_registry_save`. It is not a stable interchange format +yet: do not depend on it outside parakeet.cpp. + +`SpeakerIdentifier::update` has an internal contract about which segments are +listed as still open (see `src/speaker_identifier.hpp`). Callers of the scene +stream do not need to care about it. + +## What has been measured + +Everything here is one fixture: `tests/fixtures/two_speakers.wav`, two +read-speech LibriSpeech voices (1272 and 2086) alternating A-B-A-B. Nothing +else has been run. + +Clip-to-clip cosine between two clips of one voice, and between clips of two +different voices. Each clip is a whole turn, as in +`tests/test_speaker_encoder.cpp`: voice A 0.6 to 5.4 s and 14.9 to 18.7 s, +voice B 6.9 to 10.7 s and 20.2 to 23.5 s. "Same voice" averages the A pair and +the B pair, "different voices" averages the four A-B pairs. The design spike +measured 2 s windows of the same file instead, and shorter windows give lower +numbers (for example WeSpeaker 0.585 same voice), so the two sets differ but +do not disagree. + +| Encoder | same voice | different voices | +| --- | --- | --- | +| WeSpeaker ResNet34 | 0.869 | -0.012 | +| CAM++ | 0.894 | 0.383 | +| ECAPA-TDNN | 0.934 | 0.558 | + +In `tests/test_speaker_identify.cpp` the scene stream names both voices right +with all three encoders and the default thresholds, with the two voices +enrolled in the reverse of the order they speak, so slot numbers cannot be +matched to registry order. The enrollment clips are cut from the same +recording that is then streamed, so these scores are optimistic. Genuine slot +scores: + +| Encoder | genuine slot scores | +| --- | --- | +| WeSpeaker ResNet34 | 0.922 to 0.968 | +| CAM++ | 0.940 to 0.983 | +| ECAPA-TDNN | 0.959 to 0.985 | + +An impostor voice (the voice that is not in the registry) scored -0.019 with +WeSpeaker, 0.374 with CAM++ and 0.546 to 0.566 with ECAPA. The default +threshold of 0.5 is fine for WeSpeaker and CAM++ on this fixture, but ECAPA +would admit that impostor. That is why the test uses `accept_threshold` 0.7 for +its unenrolled-voice check, and why the right threshold is encoder specific. + +With ASR on as well (`PARAKEET_TEST_GGUF`), the same test streams the fixture +and checks that utterances for both slots carry the right name and that no +utterance ever carries the other voice's name. + +### Not measured yet (open work) + +The accuracy beyond this fixture is unknown. Still to do: + +- a third voice, and more than two speakers in one recording; +- noisy audio, and audio with overlapping speech; +- telephone-band or other far-from-read-speech audio; +- enrollment from a different session or microphone than the test audio + (here enrollment and test share a recording); +- a threshold sweep per encoder against a labelled set, so the defaults come + from data. diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index f2fa537..6ceadcf 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -18,10 +18,13 @@ #include "ggml.h" #include "gguf.h" #include "transcription_json.hpp" +#include "common.hpp" // pk::write_file_atomic #include "diarization.hpp" #include "ced_tagger.hpp" #include "scene_stream.hpp" #include "scene_render.hpp" +#include "speaker_encoder.hpp" +#include "speaker_registry.hpp" #include #include #include @@ -30,9 +33,11 @@ #include #include #include +#include #include #include #include +#include #include #include #include @@ -1339,21 +1344,131 @@ static int cmd_bench_decode(int argc, char** argv) { return 0; } +// Reads a whole file. Returns 0 on success, ENOENT when the path does not +// exist, EISDIR when it is a directory, else the errno of the failure. +static int read_file_bytes(const std::string& path, std::string& out) { + std::error_code ec; + if (std::filesystem::is_directory(path, ec)) return EISDIR; + errno = 0; + FILE* f = std::fopen(path.c_str(), "rb"); + if (!f) return errno ? errno : EIO; + out.clear(); + char buf[4096]; + size_t k; + while ((k = std::fread(buf, 1, sizeof(buf), f)) > 0) out.append(buf, k); + const int e = std::ferror(f) ? (errno ? errno : EIO) : 0; + std::fclose(f); + return e; +} + +// One line for a registry file that cannot be read. +static std::string registry_read_error(const std::string& path, int e) { + if (e == EISDIR) return path + " is a directory, not a speaker registry"; + return "cannot read registry " + path + ": " + std::strerror(e); +} + +static const char* kEnrollUsage = + "usage: parakeet-cli enroll --model --name " + "--input [--input ...] --registry \n"; + +// parakeet-cli enroll --model --name --input [--input ...] +// --registry +// Embeds each input as one clip of and adds it to the registry file +// (created when missing). The file is written only after every clip embedded. +static int cmd_enroll(int argc, char** argv) { + std::string model, name, registry_path; + std::vector inputs; + for (int i = 0; i < argc; ++i) { + if (std::strcmp(argv[i], "--model") == 0 && i + 1 < argc) model = argv[++i]; + else if (std::strcmp(argv[i], "--name") == 0 && i + 1 < argc) name = argv[++i]; + else if (std::strcmp(argv[i], "--input") == 0 && i + 1 < argc) inputs.push_back(argv[++i]); + else if (std::strcmp(argv[i], "--registry") == 0 && i + 1 < argc) registry_path = argv[++i]; + else { std::fprintf(stderr, "%s", kEnrollUsage); return 2; } + } + if (model.empty() || name.empty() || inputs.empty() || registry_path.empty()) { + std::fprintf(stderr, "%s", kEnrollUsage); + return 2; + } + if (!pk::SpeakerEncoder::available()) { + std::fprintf(stderr, "parakeet-cli: built without speaker identification (PARAKEET_WITH_VOICEDETECT=OFF)\n"); + return 2; + } + auto enc = pk::SpeakerEncoder::load(model); + if (!enc) { + std::fprintf(stderr, "parakeet-cli enroll: failed to load speaker model %s\n", model.c_str()); + return 1; + } + pk::SpeakerRegistry reg; + std::string blob; + const int rerr = read_file_bytes(registry_path, blob); + if (rerr != 0 && rerr != ENOENT) { // only a missing file means "start a new registry" + std::fprintf(stderr, "parakeet-cli enroll: %s\n", registry_read_error(registry_path, rerr).c_str()); + return 1; + } + if (rerr == 0) { // add to an existing registry + try { reg = pk::SpeakerRegistry::deserialize(blob); } + catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-cli enroll: %s is not a speaker registry: %s\n", + registry_path.c_str(), e.what()); + return 1; + } + } + int clips = 0; + for (const std::string& in : inputs) { + pk::Audio audio; + if (!load_audio_arg_16k_mono(in, audio)) { + std::fprintf(stderr, "parakeet-cli enroll: failed to load audio %s\n", + input_display_name(in).c_str()); + return 1; + } + std::vector emb; + if (!enc->embed(audio.samples.data(), (int)audio.samples.size(), emb)) { + std::fprintf(stderr, "parakeet-cli enroll: %s: %s\n", input_display_name(in).c_str(), + enc->last_error().c_str()); + return 1; + } + try { reg.enroll(name, emb); } + catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-cli enroll: %s\n", e.what()); + return 1; + } + ++clips; + } + // Written next to the target and moved over it, so a failed write never + // costs the user the registry they already had. + std::string werr; + if (!pk::write_file_atomic(registry_path, reg.serialize(), &werr)) { + std::fprintf(stderr, "parakeet-cli enroll: %s\n", werr.c_str()); + return 1; + } + std::printf("enrolled %s (%d clip(s)), registry has %zu speaker(s)\n", name.c_str(), clips, + reg.size()); + return 0; +} + static const char* kSceneUsage = "usage: parakeet-cli scene [--model ] [--diar ] " - "[--sound ] --input " + "[--sound ] [--speakers --registry " + "[--speaker-threshold F]] --input " "[--latency model|low|very_low|ultra_low] [--chunk-ms N] " - "[--show-speech] [--json]\n"; + "[--show-speech] [--json]\n" + " --speaker-threshold: default 0.5; ECAPA needs about 0.7, see docs/speaker.md\n"; // parakeet-cli scene [--model ] [--diar ] [--sound ] +// [--speakers --registry [--speaker-threshold F]] // --input [--latency model|low|very_low|ultra_low] // [--chunk-ms N] [--show-speech] [--json] +// --speakers names diarized speakers from the enrolled voices in --registry +// (made by `parakeet-cli enroll`); it needs --diar and --registry. // Streams the WAV through pk::SceneStream (ASR + diarization + sound events, // each optional -- at least one is required) and prints a time-ordered // transcript with sound annotations. --json prints scene_update_to_json per // update (one JSON document per line) instead of the rendered transcript. static int cmd_scene(int argc, char** argv) { std::string model, diar, sound, input, latency_str; + std::string speakers, registry_path; + bool have_threshold = false; + float speaker_threshold = 0.0f; bool json = false; bool show_speech = false; int chunk_ms = 200; @@ -1364,6 +1479,19 @@ static int cmd_scene(int argc, char** argv) { diar = argv[++i]; } else if (std::strcmp(argv[i], "--sound") == 0 && i + 1 < argc) { sound = argv[++i]; + } else if (std::strcmp(argv[i], "--speakers") == 0 && i + 1 < argc) { + speakers = argv[++i]; + } else if (std::strcmp(argv[i], "--registry") == 0 && i + 1 < argc) { + registry_path = argv[++i]; + } else if (std::strcmp(argv[i], "--speaker-threshold") == 0 && i + 1 < argc) { + char* end = nullptr; + const char* txt = argv[++i]; + speaker_threshold = std::strtof(txt, &end); + if (end == txt || *end != '\0') { + std::fprintf(stderr, "parakeet-cli scene: --speaker-threshold needs a number, got '%s'\n", txt); + return 2; + } + have_threshold = true; } else if (std::strcmp(argv[i], "--input") == 0 && i + 1 < argc) { input = argv[++i]; } else if (std::strcmp(argv[i], "--latency") == 0 && i + 1 < argc) { @@ -1415,6 +1543,31 @@ static int cmd_scene(int argc, char** argv) { std::fprintf(stderr, "parakeet-cli: built without sound tagging (PARAKEET_WITH_CED=OFF)\n"); return 2; } + if (speakers.empty() && (!registry_path.empty() || have_threshold)) { + std::fprintf(stderr, "parakeet-cli scene: --registry and --speaker-threshold need --speakers\n"); + return 2; + } + pk::SpeakerIdOpts speaker_opts; + if (!speakers.empty()) { + if (diar.empty()) { + std::fprintf(stderr, "parakeet-cli scene: --speakers needs --diar\n"); + return 2; + } + if (registry_path.empty()) { + std::fprintf(stderr, "parakeet-cli scene: --speakers needs --registry\n"); + return 2; + } + if (!pk::SpeakerEncoder::available()) { + std::fprintf(stderr, "parakeet-cli: built without speaker identification (PARAKEET_WITH_VOICEDETECT=OFF)\n"); + return 2; + } + if (have_threshold) speaker_opts.accept_threshold = speaker_threshold; + const std::string bad = pk::validate_speaker_opts(speaker_opts); + if (!bad.empty()) { + std::fprintf(stderr, "parakeet-cli scene: invalid speaker options: %s\n", bad.c_str()); + return 2; + } + } std::unique_ptr asr_model; if (!model.empty()) { @@ -1442,6 +1595,36 @@ static int cmd_scene(int argc, char** argv) { } } + std::unique_ptr speaker_enc; + pk::SpeakerRegistry registry; + if (!speakers.empty()) { + speaker_enc = pk::SpeakerEncoder::load(speakers); + if (!speaker_enc) { + std::fprintf(stderr, "parakeet-cli scene: failed to load speaker model %s\n", + speakers.c_str()); + return 1; + } + std::string blob; + const int rerr = read_file_bytes(registry_path, blob); + if (rerr != 0) { + std::fprintf(stderr, "parakeet-cli scene: %s\n", registry_read_error(registry_path, rerr).c_str()); + return 1; + } + try { registry = pk::SpeakerRegistry::deserialize(blob); } + catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-cli scene: %s is not a speaker registry: %s\n", + registry_path.c_str(), e.what()); + return 1; + } + if (registry.dim() != speaker_enc->dim()) { + std::fprintf(stderr, + "parakeet-cli scene: registry %s holds %d-dim voices but %s makes %d-dim embeddings " + "(enroll again with this model)\n", + registry_path.c_str(), registry.dim(), speakers.c_str(), speaker_enc->dim()); + return 1; + } + } + pk::Audio audio; if (!load_audio_arg_16k_mono(input, audio)) { std::string display = input_display_name(input); @@ -1454,6 +1637,11 @@ static int cmd_scene(int argc, char** argv) { parts.diar = diar_model.get(); parts.diar_latency = latency; parts.tagger = tagger.get(); + if (speaker_enc) { + parts.speaker_embed = speaker_enc->embedder(); + parts.registry = ®istry; + parts.speaker_opts = speaker_opts; + } // scene_update_to_json's label(i) may return nullptr (emitted as ""); the // same lambda drives the renderer's --json-less line formatting. @@ -1524,6 +1712,8 @@ int main(int argc, char** argv) { return run_and_shutdown(cmd_bench_decode, argc - 2, argv + 2); if (argc >= 2 && std::strcmp(argv[1], "bench") == 0) return run_and_shutdown(cmd_bench, argc - 2, argv + 2); + if (argc >= 2 && std::strcmp(argv[1], "enroll") == 0) + return run_and_shutdown(cmd_enroll, argc - 2, argv + 2); if (argc >= 2 && std::strcmp(argv[1], "scene") == 0) return run_and_shutdown(cmd_scene, argc - 2, argv + 2); std::fprintf(stderr, @@ -1542,8 +1732,12 @@ int main(int argc, char** argv) { " parakeet-cli bench-decode --model --audio " "[--batch-sizes 1,4,8,16] [--threads N] [--reps R] [--json ]\n" " parakeet-cli scene [--model ] [--diar ] " - "[--sound ] --input " + "[--sound ] [--speakers --registry " + "[--speaker-threshold F]] --input " "[--latency model|low|very_low|ultra_low] [--chunk-ms N] " - "[--show-speech] [--json]\n"); + "[--show-speech] [--json]\n" + " --speaker-threshold: default 0.5; ECAPA needs about 0.7, see docs/speaker.md\n" + " parakeet-cli enroll --model --name " + "--input [--input ...] --registry \n"); return 2; } diff --git a/include/parakeet_capi.h b/include/parakeet_capi.h index 05a27a3..04e4502 100644 --- a/include/parakeet_capi.h +++ b/include/parakeet_capi.h @@ -58,6 +58,11 @@ typedef struct parakeet_ctx parakeet_ctx; // v8: sound-event detection (CED), sound_stream_*, scene_stream_*; additive. // A CED GGUF loads into a third parakeet_ctx kind (a "tagger"); no // existing signatures changed. +// v9: speaker identification (voice-detect.cpp). A voice-detect GGUF loads +// into a fourth parakeet_ctx kind (a "speaker" encoder); a +// parakeet_speaker_registry holds enrolled voices; the scene stream and +// speaker-attributed ASR can name diarized speakers. Additive: no +// existing signature changed. int parakeet_capi_abi_version(void); // Load a GGUF model. Returns an owning context, or NULL on failure. @@ -528,6 +533,7 @@ const char* parakeet_capi_class_label(const parakeet_ctx* ctx, int index); #define PARAKEET_MODEL_KIND_ASR 1 #define PARAKEET_MODEL_KIND_DIARIZATION 2 #define PARAKEET_MODEL_KIND_SOUND 3 +#define PARAKEET_MODEL_KIND_SPEAKER 4 int parakeet_capi_model_kind(const parakeet_ctx* ctx); // --- Combined scene stream (ABI v8) ----------------------------------------- @@ -541,6 +547,13 @@ typedef struct { int diar_latency; // PARAKEET_DIAR_LATENCY_*, used only with a diar ctx parakeet_sound_opts sound; // used only with a tagger ctx int flags; // reserved, must be 0 + // Speaker identification (used only with a speaker ctx and a registry). + // Read only when `size` covers them; 0 keeps the default of that field. + float speaker_accept_threshold; // default 0.5 + float speaker_margin; // default 0.05 + float speaker_min_voice_sec; // default 2.0 + float speaker_refresh_sec; // default 3.0 + float speaker_max_voice_sec; // default 10.0 } parakeet_scene_opts; void parakeet_capi_scene_opts_default(parakeet_scene_opts* o); @@ -572,6 +585,65 @@ char* parakeet_capi_scene_stream_drain_scores_json(parakeet_scene_stream* s); const char* parakeet_capi_scene_stream_last_error(parakeet_scene_stream* s); void parakeet_capi_scene_stream_free(parakeet_scene_stream* s); +// --- Speaker identification (ABI v9) ---------------------------------------- +// A voice-detect.cpp GGUF (WeSpeaker, CAM++, ECAPA, ERes2Net) loads through +// parakeet_capi_load into a "speaker" context. A registry holds enrolled +// voices; it is model specific (the embedding size must match the model that +// enrolled the voices). Errors are reported on the speaker ctx +// (parakeet_capi_last_error) unless stated otherwise. + +// Embedding size of a speaker ctx; -1 for a context that is not a speaker model. +int parakeet_capi_speaker_dim(const parakeet_ctx* ctx); + +typedef struct parakeet_speaker_registry parakeet_speaker_registry; + +// New empty registry, or NULL on out of memory. Free with _free (safe on NULL). +parakeet_speaker_registry* parakeet_capi_speaker_registry_new(void); +void parakeet_capi_speaker_registry_free(parakeet_speaker_registry* reg); +// Number of enrolled speakers; 0 for NULL. +int parakeet_capi_speaker_registry_size(const parakeet_speaker_registry* reg); +// Last error of this registry (save/load/enroll bookkeeping), "" if none. Borrowed. +const char* parakeet_capi_speaker_registry_last_error(const parakeet_speaker_registry* reg); + +// Embeds 16 kHz (or resampled) mono PCM with `speaker` and adds it under `name`. +// Enrolling a name again refines that voice. 0 on success, nonzero on error +// (empty name, no audio, ctx not a speaker model, embedding size differs from +// the registry's); the message is on the speaker ctx. +int parakeet_capi_speaker_enroll(parakeet_speaker_registry* reg, parakeet_ctx* speaker, + const char* name, const float* pcm, int n, int sample_rate); + +// Binary file. 0 on success; nonzero on error (message on the registry). +int parakeet_capi_speaker_registry_save(const parakeet_speaker_registry* reg, const char* path); +// NULL when the file is missing or is not a valid registry. Free with _free. +parakeet_speaker_registry* parakeet_capi_speaker_registry_load(const char* path); + +// One-shot identification of a clip: {"name":"alice","score":0.71}, with +// "name":"" when unknown (score is then the best cosine). NULL on error +// (message on the speaker ctx). Free with parakeet_capi_free_string. +char* parakeet_capi_speaker_identify_pcm_json(parakeet_speaker_registry* reg, parakeet_ctx* speaker, + const float* pcm, int n, int sample_rate); + +// Like parakeet_capi_scene_stream_begin, plus speaker naming. A speaker needs a registry and a diarization +// ctx (a registry without a speaker is ignored). The registry is borrowed: +// keep it alive and unchanged while the stream runs. Speaker option fields of `o` are honoured only when o->size covers +// them. With NULL speaker and registry this is exactly +// parakeet_capi_scene_stream_begin. +parakeet_scene_stream* parakeet_capi_scene_stream_begin_speaker(parakeet_ctx* asr, parakeet_ctx* diar, + parakeet_ctx* tagger, + parakeet_ctx* speaker, + parakeet_speaker_registry* registry, + const parakeet_scene_opts* o); + +// Same document as parakeet_capi_transcribe_and_diarize_json, plus "name" and +// "name_score" on each utterance and word (empty name = unknown) and a +// top-level "names" map from slot to {"name","score"}. NULL on error. Free with +// parakeet_capi_free_string. +char* parakeet_capi_transcribe_and_diarize_named_json(parakeet_ctx* asr, parakeet_ctx* diar, + parakeet_ctx* speaker, + parakeet_speaker_registry* registry, + const float* samples, int n_samples, + int sample_rate); + #ifdef __cplusplus } // extern "C" #endif diff --git a/src/common.cpp b/src/common.cpp index 7a1adb1..8b0d84e 100644 --- a/src/common.cpp +++ b/src/common.cpp @@ -1 +1,54 @@ #include "common.hpp" + +#include +#include + +#ifdef _WIN32 +#ifndef WIN32_LEAN_AND_MEAN +#define WIN32_LEAN_AND_MEAN +#endif +#ifndef NOMINMAX +#define NOMINMAX +#endif +#include +#endif + +namespace pk { + +bool write_file_atomic(const std::string& path, const std::string& bytes, std::string* err) { + auto fail = [&](const std::string& why) { + if (err) *err = why; + return false; + }; + if (path.empty()) return fail("path is empty"); + const std::string tmp = path + ".tmp"; + std::FILE* f = std::fopen(tmp.c_str(), "wb"); + if (!f) return fail("cannot write " + tmp + ": " + std::strerror(errno)); + bool ok = bytes.empty() || std::fwrite(bytes.data(), 1, bytes.size(), f) == bytes.size(); + int e = ok ? 0 : errno; + if (std::fclose(f) != 0 && ok) { + ok = false; + e = errno; + } + if (!ok) { + std::remove(tmp.c_str()); + return fail("cannot write " + tmp + ": " + (e ? std::strerror(e) : "write failed")); + } +#ifdef _WIN32 + // rename() on the MSVC runtime fails when the target exists. + if (!MoveFileExA(tmp.c_str(), path.c_str(), MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH)) { + const unsigned long code = GetLastError(); + std::remove(tmp.c_str()); + return fail("cannot replace " + path + " (Windows error " + std::to_string(code) + ")"); + } +#else + if (std::rename(tmp.c_str(), path.c_str()) != 0) { + const int re = errno; + std::remove(tmp.c_str()); + return fail("cannot replace " + path + ": " + std::strerror(re)); + } +#endif + return true; +} + +} // namespace pk diff --git a/src/common.hpp b/src/common.hpp index b644916..37b9853 100644 --- a/src/common.hpp +++ b/src/common.hpp @@ -1,3 +1,15 @@ #pragma once #include +#include #define PK_LOG(...) do { std::fprintf(stderr, "[parakeet] " __VA_ARGS__); std::fprintf(stderr, "\n"); } while (0) + +namespace pk { + +// Writes `bytes` to `path` so a reader sees either the old file or the whole +// new one. It writes `.tmp` in the same directory, checks every write +// and the close, then replaces `path` with it (MoveFileExA on Windows, rename +// elsewhere). On any failure it removes the tmp file, leaves `path` as it was, +// sets `*err` (when not null) to a one-line reason and returns false. +bool write_file_atomic(const std::string& path, const std::string& bytes, std::string* err); + +} // namespace pk diff --git a/src/parakeet_capi.cpp b/src/parakeet_capi.cpp index 8998f4b..bb91682 100644 --- a/src/parakeet_capi.cpp +++ b/src/parakeet_capi.cpp @@ -10,6 +10,11 @@ #include "sas_merge.hpp" // pk::merge_asr_diarization, pk::group_speaker_words #include "diar_pcm_stream.hpp" // pk::DiarPcmStream #include "scene_stream.hpp" // pk::SceneStream +#include "speaker_encoder.hpp" // pk::SpeakerEncoder +#include "speaker_identifier.hpp" // pk::identify_offline +#include "speaker_registry.hpp" // pk::SpeakerRegistry +#include "audio_io.hpp" // pk::resample_linear +#include "common.hpp" // pk::write_file_atomic #include "transcription.hpp" // pk::Transcription, pk::Word #include "transcription_json.hpp" @@ -22,6 +27,7 @@ #include #include #include +#include #include #include #include @@ -50,16 +56,26 @@ // (transcribe_and_diarize*, sas_stream_*) and streaming diarization // (diarize_stream_*). A context holds either an ASR or a diarization model. // v8: sound-event detection (CED), sound_stream_*, scene_stream_*; additive. -#define PARAKEET_CAPI_ABI_VERSION 8 +// v9: speaker identification (voice-detect.cpp): a speaker ctx kind, a speaker +// registry, scene_stream_begin_speaker, transcribe_and_diarize_named_json; +// additive. +#define PARAKEET_CAPI_ABI_VERSION 9 // The opaque context: a loaded model plus a buffer for the last error message. -// Exactly one of `model` / `diar` / `tagger` is non-null: ASR models use -// `model`, diarization models (Sortformer) use `diar`, CED sound-event -// taggers use `tagger`. +// Exactly one of `model` / `diar` / `tagger` / `speaker` is non-null: ASR models +// use `model`, diarization models (Sortformer) use `diar`, CED sound-event +// taggers use `tagger`, voice-detect speaker encoders use `speaker`. struct parakeet_ctx { std::unique_ptr model; std::unique_ptr diar; std::unique_ptr tagger; + std::unique_ptr speaker; + std::string last_error; +}; + +// Enrolled voices plus a buffer for the last save/load error. +struct parakeet_speaker_registry { + pk::SpeakerRegistry reg; std::string last_error; }; @@ -204,6 +220,14 @@ extern "C" parakeet_ctx* parakeet_capi_load(const char* gguf_path) { return nullptr; } + // A voice-detect GGUF (architecture "voicedetect") is a speaker encoder. + if (pk::gguf_is_voicedetect(gguf_path)) { + ctx->speaker = pk::SpeakerEncoder::load(gguf_path); + if (ctx->speaker) return ctx; + delete ctx; + return nullptr; + } + // 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. @@ -940,6 +964,7 @@ bool require_diar(parakeet_ctx* ctx) { if (!ctx->diar) { ctx->last_error = ctx->model ? "context holds an ASR model; diarize_* needs a diarization model" : ctx->tagger ? "context holds a CED sound model; diarize_* needs a diarization model" + : ctx->speaker ? "context holds a speaker model; diarize_* needs a diarization model" : "context has no loaded model"; return false; } @@ -951,6 +976,7 @@ bool require_asr(parakeet_ctx* ctx) { if (!ctx->model) { ctx->last_error = ctx->diar ? "context holds a diarization model; an ASR model is needed here" : ctx->tagger ? "context holds a CED sound model; an ASR model is needed here" + : ctx->speaker ? "context holds a speaker model; an ASR model is needed here" : "context has no loaded model"; return false; } @@ -965,12 +991,28 @@ bool require_tagger(parakeet_ctx* ctx) { if (!ctx->tagger) { ctx->last_error = ctx->model ? "context holds an ASR model; a CED sound model is needed here" : ctx->diar ? "context holds a diarization model; a CED sound model is needed here" + : ctx->speaker ? "context holds a speaker model; a CED sound model is needed here" : "context has no loaded model"; return false; } return true; } +constexpr const char* kNoSpeaker = "built without speaker identification (PARAKEET_WITH_VOICEDETECT=OFF)"; + +bool require_speaker(parakeet_ctx* ctx) { + if (!ctx) return false; + if (!pk::SpeakerEncoder::available()) { ctx->last_error = kNoSpeaker; return false; } + if (!ctx->speaker) { + ctx->last_error = ctx->model ? "context holds an ASR model; a speaker model is needed here" + : ctx->diar ? "context holds a diarization model; a speaker model is needed here" + : ctx->tagger ? "context holds a CED sound model; a speaker 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); @@ -990,9 +1032,16 @@ char* diar_result_to_json(const pk::DiarizationResult& r) { } template -void append_speaker_item(std::string& s, const T& x, const char* time_fmt) { +void append_speaker_item(std::string& s, const T& x, const char* time_fmt, + const std::map* names = nullptr) { s += "{\"speaker\":"; pk::append_json_int(s, x.speaker); + if (names) { + s += ",\"name\":"; + pk::append_json_string(s, x.name); + s += ",\"name_score\":"; + pk::append_json_float(s, "%.4f", x.name_score); + } s += ",\"text\":"; pk::append_json_string(s, x.text); s += ",\"start\":"; @@ -1028,7 +1077,9 @@ bool to_c_results(const std::vector& utts, // 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) { + std::vector& words, int& n_speakers, + std::vector* segs_out = nullptr, + std::vector* pcm16k_out = nullptr) { if (!require_asr(asr_ctx) || !require_diar(diar_ctx)) return false; if (!samples || n_samples < 0) { asr_ctx->last_error = "invalid samples buffer"; @@ -1051,6 +1102,8 @@ bool run_sas(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, } n_speakers = dr.n_speakers; words = pk::merge_asr_diarization(tr.words, dr.segments); + if (segs_out) *segs_out = dr.segments; + if (pcm16k_out) *pcm16k_out = sample_rate == 16000 ? pcm : pk::resample_linear(pcm, sample_rate, 16000); asr_ctx->last_error.clear(); diar_ctx->last_error.clear(); return true; @@ -1485,6 +1538,7 @@ extern "C" int parakeet_capi_model_kind(const parakeet_ctx* ctx) { if (ctx->model) return PARAKEET_MODEL_KIND_ASR; if (ctx->diar) return PARAKEET_MODEL_KIND_DIARIZATION; if (ctx->tagger) return PARAKEET_MODEL_KIND_SOUND; + if (ctx->speaker) return PARAKEET_MODEL_KIND_SPEAKER; return PARAKEET_MODEL_KIND_NONE; } @@ -1497,6 +1551,7 @@ struct parakeet_scene_stream { parakeet_ctx* asr_ctx = nullptr; parakeet_ctx* diar_ctx = nullptr; parakeet_ctx* tagger_ctx = nullptr; + parakeet_ctx* speaker_ctx = nullptr; std::unique_ptr scene; std::string last_error; }; @@ -1510,6 +1565,7 @@ parakeet_ctx* scene_failed_ctx(parakeet_scene_stream* s) { case pk::ScenePart::Diarization: return s->diar_ctx ? s->diar_ctx : s->asr_ctx; case pk::ScenePart::Asr: return s->asr_ctx ? s->asr_ctx : s->diar_ctx; case pk::ScenePart::Sound: return s->tagger_ctx; + case pk::ScenePart::Speaker: return s->speaker_ctx ? s->speaker_ctx : s->diar_ctx; default: return s->asr_ctx ? s->asr_ctx : s->diar_ctx ? s->diar_ctx : s->tagger_ctx; } @@ -1523,11 +1579,43 @@ extern "C" void parakeet_capi_scene_opts_default(parakeet_scene_opts* o) { o->diar_latency = PARAKEET_DIAR_LATENCY_MODEL; parakeet_capi_sound_opts_default(&o->sound); o->flags = 0; + const pk::SpeakerIdOpts d; + o->speaker_accept_threshold = d.accept_threshold; + o->speaker_margin = d.margin; + o->speaker_min_voice_sec = d.min_voice_sec; + o->speaker_refresh_sec = d.refresh_sec; + o->speaker_max_voice_sec = d.max_voice_sec; } +namespace { + +// Reads the float at byte offset `off` of `o` only when the caller's `size` +// covers it (a caller built against an older header owns a shorter struct, so +// nothing past `size` may be touched). A zero or uncovered field returns +// `dflt`. +float scene_float_field(const parakeet_scene_opts* o, size_t off, float dflt) { + if (o->size < (int)(off + sizeof(float))) return dflt; + float v; + std::memcpy(&v, reinterpret_cast(o) + off, sizeof(v)); + return v != 0.0f ? v : dflt; +} + +} // namespace + extern "C" parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* asr, parakeet_ctx* diar, parakeet_ctx* tagger, const parakeet_scene_opts* o) { + return parakeet_capi_scene_stream_begin_speaker(asr, diar, tagger, nullptr, nullptr, o); +} + +extern "C" parakeet_scene_stream* parakeet_capi_scene_stream_begin_speaker( + parakeet_ctx* asr, parakeet_ctx* diar, parakeet_ctx* tagger, parakeet_ctx* speaker, + parakeet_speaker_registry* registry, const parakeet_scene_opts* o) { + if (speaker) { + if (!require_speaker(speaker)) return nullptr; + if (!diar) { speaker->last_error = "speaker identification needs a diarization model"; return nullptr; } + if (!registry) { speaker->last_error = "speaker identification needs a registry"; return nullptr; } + } if (!asr && !diar && !tagger) return nullptr; if ((asr && !require_asr(asr)) || (diar && !require_diar(diar)) || (tagger && !require_tagger(tagger))) return nullptr; @@ -1543,8 +1631,30 @@ extern "C" parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* diar->last_error = "unknown diarization latency mode"; return nullptr; } + pk::SpeakerIdOpts so; + if (speaker) { + so.accept_threshold = scene_float_field(o, offsetof(parakeet_scene_opts, speaker_accept_threshold), so.accept_threshold); + so.margin = scene_float_field(o, offsetof(parakeet_scene_opts, speaker_margin), so.margin); + so.min_voice_sec = scene_float_field(o, offsetof(parakeet_scene_opts, speaker_min_voice_sec), so.min_voice_sec); + so.refresh_sec = scene_float_field(o, offsetof(parakeet_scene_opts, speaker_refresh_sec), so.refresh_sec); + so.max_voice_sec = scene_float_field(o, offsetof(parakeet_scene_opts, speaker_max_voice_sec), so.max_voice_sec); + const std::string err = pk::validate_speaker_opts(so); + if (!err.empty()) { speaker->last_error = "invalid speaker options: " + err; return nullptr; } + const int rd = registry->reg.dim(); + if (rd != 0 && rd != speaker->speaker->dim()) { + speaker->last_error = "registry holds " + std::to_string(rd) + + "-value embeddings, this model produces " + + std::to_string(speaker->speaker->dim()); + return nullptr; + } + } try { pk::SceneParts p; + if (speaker) { + p.speaker_embed = speaker->speaker->embedder(); + p.registry = ®istry->reg; + p.speaker_opts = so; + } p.asr = asr ? asr->model.get() : nullptr; p.diar = diar ? diar->diar.get() : nullptr; p.diar_latency = latency_from_int(o->diar_latency); @@ -1554,19 +1664,22 @@ extern "C" parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* const std::string err = pk::validate_sound_opts(p.sound, tagger->tagger->n_classes()); if (!err.empty()) { tagger->last_error = "invalid sound options: " + err; return nullptr; } } + auto scene = std::make_unique(p); auto* s = new parakeet_scene_stream(); + s->scene = std::move(scene); s->asr_ctx = asr; s->diar_ctx = diar; s->tagger_ctx = tagger; - s->scene = std::make_unique(p); + s->speaker_ctx = speaker; if (asr) asr->last_error.clear(); if (diar) diar->last_error.clear(); if (tagger) tagger->last_error.clear(); + if (speaker) speaker->last_error.clear(); return s; } catch (const std::exception& e) { - (asr ? asr : diar ? diar : tagger)->last_error = e.what(); + (speaker ? speaker : asr ? asr : diar ? diar : tagger)->last_error = e.what(); } catch (...) { - (asr ? asr : diar ? diar : tagger)->last_error = "unknown error"; + (speaker ? speaker : asr ? asr : diar ? diar : tagger)->last_error = "unknown error"; } return nullptr; } @@ -1612,3 +1725,215 @@ extern "C" const char* parakeet_capi_scene_stream_last_error(parakeet_scene_stre } extern "C" void parakeet_capi_scene_stream_free(parakeet_scene_stream* s) { delete s; } + +// --------------------------------------------------------------------------- +// Speaker identification (ABI v9) +// --------------------------------------------------------------------------- + +namespace { + +// Embeds mono PCM at `sample_rate` with the speaker ctx. Errors go on the ctx. +bool speaker_embed_pcm(parakeet_ctx* speaker, const float* pcm, int n, int sample_rate, + std::vector& emb) { + if (!pcm || n <= 0) { speaker->last_error = "no audio"; return false; } + if (sample_rate <= 0) { speaker->last_error = "invalid sample rate"; return false; } + bool ok; + if (sample_rate == 16000) { + ok = speaker->speaker->embed(pcm, n, emb); + } else { + const std::vector in(pcm, pcm + n); + const std::vector r = pk::resample_linear(in, sample_rate, 16000); + if (r.empty()) { speaker->last_error = "no audio"; return false; } + ok = speaker->speaker->embed(r.data(), (int)r.size(), emb); + } + if (!ok) { + const std::string& e = speaker->speaker->last_error(); + speaker->last_error = e.empty() ? "speaker embedding failed" : e; + return false; + } + return true; +} + +} // namespace + +extern "C" int parakeet_capi_speaker_dim(const parakeet_ctx* ctx) { + return (ctx && ctx->speaker) ? ctx->speaker->dim() : -1; +} + +extern "C" parakeet_speaker_registry* parakeet_capi_speaker_registry_new(void) { + return new (std::nothrow) parakeet_speaker_registry(); +} + +extern "C" void parakeet_capi_speaker_registry_free(parakeet_speaker_registry* reg) { delete reg; } + +extern "C" int parakeet_capi_speaker_registry_size(const parakeet_speaker_registry* reg) { + return reg ? (int)reg->reg.size() : 0; +} + +extern "C" const char* parakeet_capi_speaker_registry_last_error(const parakeet_speaker_registry* reg) { + return reg ? reg->last_error.c_str() : ""; +} + +extern "C" int parakeet_capi_speaker_enroll(parakeet_speaker_registry* reg, parakeet_ctx* speaker, + const char* name, const float* pcm, int n, int sample_rate) { + if (!speaker) return 1; + try { + if (!require_speaker(speaker)) return 1; + if (!reg) { speaker->last_error = "registry is NULL"; return 1; } + if (!name || !*name) { speaker->last_error = "speaker name is empty"; return 1; } + std::vector emb; + if (!speaker_embed_pcm(speaker, pcm, n, sample_rate, emb)) return 1; + reg->reg.enroll(name, emb); + speaker->last_error.clear(); + return 0; + } catch (const std::exception& e) { + speaker->last_error = e.what(); + } catch (...) { + speaker->last_error = "unknown error"; + } + return 1; +} + +extern "C" int parakeet_capi_speaker_registry_save(const parakeet_speaker_registry* reg, const char* path) { + if (!reg) return 1; + auto* mreg = const_cast(reg); // only last_error is written + if (!path || !*path) { mreg->last_error = "path is empty"; return 1; } + try { + // Written to .tmp and moved over the target, so a failed save + // never costs the caller the registry file they already had. + std::string err; + if (!pk::write_file_atomic(path, reg->reg.serialize(), &err)) { + mreg->last_error = err; + return 1; + } + mreg->last_error.clear(); + return 0; + } catch (const std::exception& e) { + mreg->last_error = e.what(); + } catch (...) { + mreg->last_error = "unknown error"; + } + return 1; +} + +extern "C" parakeet_speaker_registry* parakeet_capi_speaker_registry_load(const char* path) { + if (!path) return nullptr; + try { + std::FILE* f = std::fopen(path, "rb"); + if (!f) return nullptr; + std::string blob; + char buf[4096]; + size_t got; + while ((got = std::fread(buf, 1, sizeof(buf), f)) > 0) blob.append(buf, got); + const bool err = std::ferror(f) != 0; + std::fclose(f); + if (err) return nullptr; + auto* r = new (std::nothrow) parakeet_speaker_registry(); + if (!r) return nullptr; + try { + r->reg = pk::SpeakerRegistry::deserialize(blob); + } catch (...) { + delete r; + return nullptr; + } + return r; + } catch (...) { + return nullptr; + } +} + +extern "C" char* parakeet_capi_speaker_identify_pcm_json(parakeet_speaker_registry* reg, parakeet_ctx* speaker, + const float* pcm, int n, int sample_rate) { + if (!speaker) return nullptr; + try { + if (!require_speaker(speaker)) return nullptr; + if (!reg) { speaker->last_error = "registry is NULL"; return nullptr; } + std::vector emb; + if (!speaker_embed_pcm(speaker, pcm, n, sample_rate, emb)) return nullptr; + const pk::SpeakerIdOpts d; + const pk::SpeakerMatch m = reg->reg.identify(emb, d.accept_threshold, d.margin); + std::string j = "{\"name\":"; + pk::append_json_string(j, m.name); + j += ",\"score\":"; + pk::append_json_float(j, "%.4f", m.score); + j += '}'; + speaker->last_error.clear(); + return dup_to_c(j); + } catch (const std::exception& e) { + speaker->last_error = e.what(); + } catch (...) { + speaker->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" char* parakeet_capi_transcribe_and_diarize_named_json( + parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, parakeet_ctx* speaker, + parakeet_speaker_registry* registry, const float* samples, int n_samples, int sample_rate) { + try { + if (!speaker) return nullptr; + if (!require_speaker(speaker)) return nullptr; + if (!registry) { speaker->last_error = "speaker identification needs a registry"; return nullptr; } + std::vector words; + std::vector segs; + std::vector pcm16k; + int n_speakers = 0; + if (!run_sas(asr_ctx, diar_ctx, samples, n_samples, sample_rate, words, n_speakers, &segs, &pcm16k)) + return nullptr; + const int rd = registry->reg.dim(); + if (rd != 0 && rd != speaker->speaker->dim()) { + speaker->last_error = "registry holds " + std::to_string(rd) + + "-value embeddings, this model produces " + + std::to_string(speaker->speaker->dim()); + return nullptr; + } + std::map names; + try { + names = pk::identify_offline(pcm16k, segs, speaker->speaker->embedder(), registry->reg, + pk::SpeakerIdOpts{}); + } catch (const std::exception& e) { + speaker->last_error = e.what(); + return nullptr; + } + for (pk::SpeakerWord& w : words) { + auto it = names.find(w.speaker); + if (w.speaker >= 0 && it != names.end()) { + w.name = it->second.name; + w.name_score = it->second.score; + } + } + const std::vector utts = pk::group_speaker_words(words); + std::string s = "{\"speakers\":"; + pk::append_json_int(s, n_speakers); + s += ",\"names\":{"; + bool first = true; + for (const auto& kv : names) { + if (!first) s += ','; + first = false; + s += "\"" + std::to_string(kv.first) + "\":{\"name\":"; + pk::append_json_string(s, kv.second.name); + s += ",\"score\":"; + pk::append_json_float(s, "%.4f", kv.second.score); + s += '}'; + } + s += "},\"utterances\":["; + for (size_t i = 0; i < utts.size(); ++i) { + if (i) s += ','; + append_speaker_item(s, utts[i], "%.2f", &names); + } + s += "],\"words\":["; + for (size_t i = 0; i < words.size(); ++i) { + if (i) s += ','; + append_speaker_item(s, words[i], "%.3f", &names); + } + s += "]}"; + speaker->last_error.clear(); + return dup_to_c(s); + } catch (const std::exception& e) { + if (speaker) speaker->last_error = e.what(); + return nullptr; + } catch (...) { + if (speaker) speaker->last_error = "unknown error"; + return nullptr; + } +} diff --git a/src/sas_merge.cpp b/src/sas_merge.cpp index 82c13b2..73ada61 100644 --- a/src/sas_merge.cpp +++ b/src/sas_merge.cpp @@ -76,6 +76,8 @@ std::vector group_speaker_words( cur.start = swords[0].start; cur.end = swords[0].end; cur.conf = swords[0].conf; + cur.name = swords[0].name; + cur.name_score = swords[0].name_score; for (size_t i = 1; i < swords.size(); ++i) { const auto& w = swords[i]; @@ -94,6 +96,8 @@ std::vector group_speaker_words( cur.start = w.start; cur.end = w.end; cur.conf = w.conf; + cur.name = w.name; + cur.name_score = w.name_score; } } result.push_back(cur); diff --git a/src/sas_merge.hpp b/src/sas_merge.hpp index 2030c86..7a9e0f0 100644 --- a/src/sas_merge.hpp +++ b/src/sas_merge.hpp @@ -15,6 +15,8 @@ struct SpeakerWord { float start; // from ASR word (seconds) float end; // from ASR word (seconds) float conf; // from ASR word + std::string name; // enrolled speaker name; empty = unknown or no speaker model + float name_score = 0.0f; }; // A speaker-attributed utterance: consecutive words from the same speaker @@ -26,6 +28,8 @@ struct SpeakerUtterance { float start; // first word start float end; // last word end float conf; // min word confidence + std::string name; // enrolled speaker name; empty = unknown or no speaker model + float name_score = 0.0f; }; // Merge ASR word timestamps with diarization speaker segments. diff --git a/src/scene_render.cpp b/src/scene_render.cpp index c3a1830..f2e2653 100644 --- a/src/scene_render.cpp +++ b/src/scene_render.cpp @@ -61,8 +61,11 @@ SceneRenderer::SceneRenderer(bool has_diar, bool show_speech, std::functionsecond.name.empty()) + ? nm->second.name + : "Speaker " + std::to_string(g.speaker); + pending_.push_back({(double)g.start, format_span(g.start, g.end) + " " + who}); diarized_ = std::max(diarized_, (double)g.end); } // A closed segment ends at or before the diarized time, and an open @@ -77,7 +80,9 @@ void SceneRenderer::add(const SceneUpdate& u) { for (const SpeakerUtterance& utt : u.utterances) { std::string line = format_span(utt.start, utt.end) + " "; if (has_diar_) { - if (utt.speaker >= 0) + if (!utt.name.empty()) + line += utt.name + ": "; + else if (utt.speaker >= 0) line += "Speaker " + std::to_string(utt.speaker) + ": "; else line += "Speaker ?: "; diff --git a/src/scene_stream.cpp b/src/scene_stream.cpp index 47028a8..06d28a3 100644 --- a/src/scene_stream.cpp +++ b/src/scene_stream.cpp @@ -10,8 +10,18 @@ namespace pk { SceneStream::SceneStream(const SceneParts& p) { + if (p.speaker_embed && !p.diar) + throw std::invalid_argument("speaker identification needs a diarization model"); + if (p.speaker_embed && !p.registry) + throw std::invalid_argument("speaker identification needs a registry"); + if (p.speaker_embed) { + const std::string err = validate_speaker_opts(p.speaker_opts); + if (!err.empty()) throw std::invalid_argument("invalid speaker options: " + err); + } if (!p.asr && !p.diar && !p.tagger) throw std::invalid_argument("scene stream needs at least one model"); + if (p.speaker_embed) + speaker_ = std::make_unique(p.speaker_embed, p.registry, p.speaker_opts); if (p.diar) diar_ = std::make_unique(*p.diar, p.diar_latency); if (p.asr) { const Model* m = p.asr; @@ -38,6 +48,18 @@ SceneUpdate SceneStream::feed(const float* pcm, int n, bool is_last) { u.speakers.push_back({c.speaker, c.start, c.end}); } } + if (speaker_) { + part_ = ScenePart::Speaker; + speaker_->push_pcm(pcm, n); + std::vector closed_segs, open_segs; + for (const auto& c : closed) closed_segs.push_back({c.speaker, c.start, c.end}); + for (const auto& o : diar_->open_segments()) open_segs.push_back({o.speaker, o.start, o.end}); + speaker_->update(closed_segs, open_segs, is_last); + u.names = speaker_->names(); + u.named = true; + // A failure past this point is charged to the next part. + part_ = asr_ ? ScenePart::Asr : ScenePart::Diarization; + } if (asr_) { asr_->push(pcm, n); // With diarization, ASR follows how far diarization has got, and only @@ -54,6 +76,12 @@ SceneUpdate SceneStream::feed(const float* pcm, int n, bool is_last) { if (diar_) for (const auto& o : diar_->open_segments()) segs.push_back({o.speaker, o.start, o.end}); u.words = merge_asr_diarization(committed, segs); + if (speaker_) + for (SpeakerWord& w : u.words) { + const SlotName sn = speaker_->name(w.speaker); + w.name = sn.name; + w.name_score = sn.score; + } u.utterances = group_speaker_words(u.words); // Segments that ended before the commit point can no longer match a word. const double commit_sec = asr_->commit_sec(); @@ -95,24 +123,45 @@ std::vector SceneStream::drain_windows() { namespace { -void append_speaker_segment(std::string& out, const SpeakerSegment& s) { +// ,"name":"alice","name_score":0.7100 : only when a speaker model ran (`named`). +void append_name(std::string& out, bool named, const std::map& names, int slot, + const std::string& own_name, float own_score, bool use_own) { + if (!named) return; + std::string n = own_name; + float sc = own_score; + if (!use_own) { + auto it = names.find(slot); + n = it == names.end() ? std::string() : it->second.name; + sc = it == names.end() ? 0.0f : it->second.score; + } + out += ",\"name\":"; append_json_string(out, n); + out += ",\"name_score\":"; append_json_float(out, "%.4f", sc); +} + +void append_speaker_segment(std::string& out, const SpeakerSegment& s, + bool named, const std::map& names) { out += "{\"speaker\":"; append_json_int(out, s.speaker); + append_name(out, named, names, s.speaker, std::string(), 0.0f, false); out += ",\"start\":"; append_json_float(out, "%.3f", s.start); out += ",\"end\":"; append_json_float(out, "%.3f", s.end); out += "}"; } -void append_active_speaker(std::string& out, const StreamingSpeakerSegment& s) { +void append_active_speaker(std::string& out, const StreamingSpeakerSegment& s, + bool named, const std::map& names) { out += "{\"speaker\":"; append_json_int(out, s.speaker); + append_name(out, named, names, s.speaker, std::string(), 0.0f, false); out += ",\"start\":"; append_json_float(out, "%.3f", s.start); out += "}"; } -std::string utterances_to_json(const std::vector& utts) { +std::string utterances_to_json(const std::vector& utts, + bool named, const std::map& names) { std::string out = "["; for (size_t i = 0; i < utts.size(); ++i) { if (i) out += ","; out += "{\"speaker\":"; append_json_int(out, utts[i].speaker); + append_name(out, named, names, utts[i].speaker, utts[i].name, utts[i].name_score, true); out += ",\"text\":"; append_json_string(out, utts[i].text); out += ",\"start\":"; append_json_float(out, "%.3f", utts[i].start); out += ",\"end\":"; append_json_float(out, "%.3f", utts[i].end); @@ -122,7 +171,8 @@ std::string utterances_to_json(const std::vector& utts) { return out + "]"; } -std::string words_to_json(const std::vector& words) { +std::string words_to_json(const std::vector& words, + bool named, const std::map& names) { std::string out = "["; for (size_t i = 0; i < words.size(); ++i) { if (i) out += ","; @@ -131,16 +181,18 @@ std::string words_to_json(const std::vector& words) { out += ",\"end\":"; append_json_float(out, "%.3f", words[i].end); out += ",\"conf\":"; append_json_float(out, "%.4f", words[i].conf); out += ",\"speaker\":"; append_json_int(out, words[i].speaker); + append_name(out, named, names, words[i].speaker, words[i].name, words[i].name_score, true); out += "}"; } return out + "]"; } -std::string speakers_to_json(const std::vector& segs) { +std::string speakers_to_json(const std::vector& segs, + bool named, const std::map& names) { std::string out = "["; for (size_t i = 0; i < segs.size(); ++i) { if (i) out += ","; - append_speaker_segment(out, segs[i]); + append_speaker_segment(out, segs[i], named, names); } return out + "]"; } @@ -150,14 +202,31 @@ std::string speakers_to_json(const std::vector& segs) { std::string scene_update_to_json(const SceneUpdate& u, const std::function& label) { std::string out = "{\"t\":"; append_json_float(out, "%.3f", (float)u.t); - out += ",\"utterances\":" + utterances_to_json(u.utterances); - out += ",\"words\":" + words_to_json(u.words); - out += ",\"speakers\":" + speakers_to_json(u.speakers); + // No speaker part: no names at all (the shape from before speaker + // identification). With one: names on every update, possibly empty. + const bool named = u.named || !u.names.empty(); + if (named) { + out += ",\"names\":{"; + bool first = true; + for (const auto& kv : u.names) { + if (!first) out += ","; + first = false; + out += "\"" + std::to_string(kv.first) + "\":{\"name\":"; + append_json_string(out, kv.second.name); + out += ",\"score\":"; + append_json_float(out, "%.4f", kv.second.score); + out += "}"; + } + out += "}"; + } + out += ",\"utterances\":" + utterances_to_json(u.utterances, named, u.names); + out += ",\"words\":" + words_to_json(u.words, named, u.names); + out += ",\"speakers\":" + speakers_to_json(u.speakers, named, u.names); out += ",\"sounds\":" + sound_segments_to_json(u.sounds, label); out += ",\"active\":{\"speakers\":["; for (size_t i = 0; i < u.active_speakers.size(); ++i) { if (i) out += ","; - append_active_speaker(out, u.active_speakers[i]); + append_active_speaker(out, u.active_speakers[i], named, u.names); } out += "],\"sounds\":" + sound_segments_to_json(u.active_sounds, label); out += "}}"; diff --git a/src/scene_stream.hpp b/src/scene_stream.hpp index b123390..78efff6 100644 --- a/src/scene_stream.hpp +++ b/src/scene_stream.hpp @@ -1,10 +1,12 @@ #pragma once #include "asr_committer.hpp" #include "diar_pcm_stream.hpp" +#include "speaker_identifier.hpp" // pk::SlotName #include "sas_merge.hpp" // pk::SpeakerWord, pk::SpeakerUtterance #include "sound_stream.hpp" // pk::SoundOpts, pk::SoundSegment, pk::SoundWindow #include +#include #include #include #include @@ -21,6 +23,9 @@ struct SceneParts { DiarLatency diar_latency = DiarLatency::Model; CedTagger* tagger = nullptr; // sound events SoundOpts sound; // sound events + SpeakerEmbed speaker_embed; // speaker identification; needs diar and registry + const SpeakerRegistry* registry = nullptr; // borrowed, must outlive the stream + SpeakerIdOpts speaker_opts; }; // What one feed finalized. All times are seconds on the stream clock. @@ -38,18 +43,24 @@ struct SceneUpdate { std::vector sounds; // closed this call std::vector active_speakers; std::vector active_sounds; + // Current identity of every diarization slot the speaker part has seen + // (unknown slots have an empty name). Empty without a speaker model. + std::map names; + // True on every update of a stream that has a speaker part, even before + // any slot is seen, so the JSON keeps one shape for the whole stream. + bool named = false; }; // The part running when feed() threw, so a caller can attribute the error. -enum class ScenePart { None, Diarization, Asr, Sound }; +enum class ScenePart { None, Diarization, Asr, Sound, Speaker }; // Speech, speakers and sound events over one live 16 kHz mono PCM stream. // Each feed gives the PCM to every part, then collects what each one // finalized. Not thread-safe. // -// Error paths: feed() runs diarization, then ASR, then the sound part, in -// that order, and does not catch between them, so a part that throws loses -// the rest of that call. If the sound part throws, anything diarization or +// Error paths: feed() runs diarization, then the speaker part (when there +// is one), then ASR, then the sound part, in that order, and does not catch +// between them, so a part that throws loses the rest of that call. If the sound part throws, anything diarization or // ASR already finalized in this call (including words the commit window // released) is lost with it, not returned before the exception propagates. // If ASR throws, the sound part for that chunk never runs (skipped, not @@ -59,12 +70,16 @@ enum class ScenePart { None, Diarization, Asr, Sound }; // stream ends without flushing whatever the throwing part (or anything // after it) would otherwise have flushed on that final call. // +// If the speaker part throws, ASR and sound for that chunk are skipped like +// any other part that throws. +// // After a feed() that throws, later timestamps may be misaligned: the parts // that did not see the failed chunk lag behind the ones that did. Callers // should end the stream after an error rather than keep feeding it. class SceneStream { public: - explicit SceneStream(const SceneParts& p); // throws std::invalid_argument when no part is given + explicit SceneStream(const SceneParts& p); // throws std::invalid_argument when no part is given, + // or a speaker part lacks diarization, a registry or valid options ~SceneStream(); SceneStream(const SceneStream&) = delete; SceneStream& operator=(const SceneStream&) = delete; @@ -83,6 +98,7 @@ class SceneStream { std::unique_ptr diar_; std::unique_ptr asr_; std::unique_ptr sound_; + std::unique_ptr speaker_; std::vector segs_; // closed diarization segments not yet behind the commit point double t_ = 0.0; // stream time consumed bool finished_ = false; diff --git a/src/speaker_encoder.cpp b/src/speaker_encoder.cpp new file mode 100644 index 0000000..6962b86 --- /dev/null +++ b/src/speaker_encoder.cpp @@ -0,0 +1,78 @@ +#include "speaker_encoder.hpp" + +#include "gguf.h" + +#ifdef PARAKEET_WITH_VOICEDETECT +#include "voicedetect_capi.h" +#endif + +namespace pk { + +bool gguf_is_voicedetect(const std::string& path) { + gguf_init_params p{/*no_alloc=*/true, /*ctx=*/nullptr}; + gguf_context* g = gguf_init_from_file(path.c_str(), p); + if (!g) return false; + const int64_t id = gguf_find_key(g, "general.architecture"); + const bool vd = id >= 0 && gguf_get_kv_type(g, id) == GGUF_TYPE_STRING && + std::string(gguf_get_val_str(g, id)) == "voicedetect"; + gguf_free(g); + return vd; +} + +#ifdef PARAKEET_WITH_VOICEDETECT + +bool SpeakerEncoder::available() { return true; } + +std::unique_ptr SpeakerEncoder::load(const std::string& path) { + voicedetect_ctx* c = voicedetect_capi_load(path.c_str()); + if (!c) return nullptr; + const int dim = voicedetect_capi_embedding_dim(c); + if (dim <= 0) { // an analyze-only model has no speaker embedding + voicedetect_capi_free(c); + return nullptr; + } + std::unique_ptr e(new SpeakerEncoder()); + e->ctx_ = c; + e->dim_ = dim; + return e; +} + +SpeakerEncoder::~SpeakerEncoder() { voicedetect_capi_free(static_cast(ctx_)); } + +bool SpeakerEncoder::embed(const float* pcm, int n, std::vector& emb) { + last_error_.clear(); // describes the latest call only + if (!pcm || n <= 0) { + last_error_ = "no audio to embed"; + return false; + } + auto* c = static_cast(ctx_); + float* v = nullptr; + int d = 0; + if (voicedetect_capi_embed_pcm(c, pcm, n, 16000, &v, &d) != 0 || !v || d != dim_) { + const char* m = voicedetect_capi_last_error(c); + last_error_ = (m && *m) ? m : "speaker embedding failed"; + voicedetect_capi_free_vec(v); + return false; + } + emb.assign(v, v + d); + voicedetect_capi_free_vec(v); + return true; +} + +#else // PARAKEET_WITH_VOICEDETECT + +bool SpeakerEncoder::available() { return false; } +std::unique_ptr SpeakerEncoder::load(const std::string&) { return nullptr; } +SpeakerEncoder::~SpeakerEncoder() = default; +bool SpeakerEncoder::embed(const float*, int, std::vector&) { + last_error_ = "built without speaker identification (PARAKEET_WITH_VOICEDETECT=OFF)"; + return false; +} + +#endif + +SpeakerEmbed SpeakerEncoder::embedder() { + return [this](const float* pcm, int n, std::vector& emb) { return embed(pcm, n, emb); }; +} + +} // namespace pk diff --git a/src/speaker_encoder.hpp b/src/speaker_encoder.hpp new file mode 100644 index 0000000..165127f --- /dev/null +++ b/src/speaker_encoder.hpp @@ -0,0 +1,42 @@ +#pragma once +#include "speaker_identifier.hpp" // pk::SpeakerEmbed + +#include +#include +#include + +namespace pk { + +// A loaded voice-detect.cpp speaker encoder (WeSpeaker, CAM++, ECAPA or ERes2Net). +// The only parakeet code that talks to voice-detect.cpp, and only through +// voicedetect_capi.h. Not thread-safe: one stream at a time per encoder, like +// the other contexts. +class SpeakerEncoder { +public: + // False when parakeet was built with PARAKEET_WITH_VOICEDETECT=OFF. + static bool available(); + // nullptr on failure, when unavailable, or when the GGUF has no speaker + // embedding (for example an age/gender/emotion model). + static std::unique_ptr load(const std::string& gguf_path); + ~SpeakerEncoder(); + SpeakerEncoder(const SpeakerEncoder&) = delete; + SpeakerEncoder& operator=(const SpeakerEncoder&) = delete; + + int dim() const { return dim_; } + // L2-normalized embedding of 16 kHz mono PCM. False on failure (see last_error). + bool embed(const float* pcm, int n, std::vector& emb); + // A SpeakerEmbed bound to this encoder; valid while the encoder lives. + SpeakerEmbed embedder(); + const std::string& last_error() const { return last_error_; } + +private: + SpeakerEncoder() = default; + void* ctx_ = nullptr; // voicedetect_ctx* + int dim_ = 0; + std::string last_error_; +}; + +// True when the GGUF's general.architecture is "voicedetect". Reads only the header. +bool gguf_is_voicedetect(const std::string& gguf_path); + +} // namespace pk diff --git a/src/speaker_identifier.cpp b/src/speaker_identifier.cpp new file mode 100644 index 0000000..359452b --- /dev/null +++ b/src/speaker_identifier.cpp @@ -0,0 +1,175 @@ +#include "speaker_identifier.hpp" + +#include +#include +#include +#include + +namespace pk { + +namespace { +constexpr int kSr = 16000; +constexpr double kMinPieceSec = 0.2; // shorter clean slivers carry no usable voice +} // namespace + +std::string validate_speaker_opts(const SpeakerIdOpts& o) { + if (!(o.min_voice_sec > 0.0f)) return "min_voice_sec must be > 0"; + if (!(o.refresh_sec > 0.0f)) return "refresh_sec must be > 0"; + if (!(o.max_voice_sec >= o.min_voice_sec)) return "max_voice_sec must be >= min_voice_sec"; + if (!(o.ring_sec >= o.max_voice_sec)) return "ring_sec must be >= max_voice_sec"; + if (!(o.accept_threshold >= -1.0f && o.accept_threshold <= 1.0f)) + return "accept_threshold must be in [-1, 1]"; + if (!(o.margin >= 0.0f)) return "margin must be >= 0"; + return ""; +} + +std::vector clean_intervals(const Interval& seg, const std::vector& others, + double min_len) { + std::vector cover; + for (const Interval& o : others) + if (o.end > seg.start && o.start < seg.end) cover.push_back(o); + std::sort(cover.begin(), cover.end(), + [](const Interval& a, const Interval& b) { return a.start < b.start; }); + std::vector out; + double cursor = seg.start; + auto emit = [&](double a, double b) { + if (b - a >= min_len) out.push_back({a, b}); + }; + for (const Interval& c : cover) { + if (c.start > cursor) emit(cursor, std::min(c.start, seg.end)); + cursor = std::max(cursor, c.end); + if (cursor >= seg.end) break; + } + if (cursor < seg.end) emit(cursor, seg.end); + return out; +} + +SpeakerIdentifier::SpeakerIdentifier(SpeakerEmbed embed, const SpeakerRegistry* registry, + SpeakerIdOpts opts) + : embed_(std::move(embed)), registry_(registry), opts_(opts) { + const std::string err = validate_speaker_opts(opts_); + if (!err.empty()) throw std::invalid_argument("invalid speaker options: " + err); + if (!embed_ || !registry_) throw std::invalid_argument("speaker identifier needs an embedder and a registry"); +} + +void SpeakerIdentifier::push_pcm(const float* pcm, int n) { + if (n <= 0 || !pcm) return; + ring_.insert(ring_.end(), pcm, pcm + n); + total_ += n; + const size_t cap = (size_t)((double)opts_.ring_sec * kSr); + if (ring_.size() > cap + (size_t)kSr) { // trim in 1 s steps so the erase is amortized + const size_t drop = ring_.size() - cap; + ring_.erase(ring_.begin(), ring_.begin() + (long)drop); + ring_base_ += (long long)drop; + } +} + +void SpeakerIdentifier::add_audio(int slot, const Interval& iv) { + long long a = std::llround(iv.start * kSr); + long long b = std::llround(iv.end * kSr); + a = std::max(a, ring_base_); + b = std::min(b, total_); + if (b <= a) return; + Slot& s = slots_[slot]; + const float* src = ring_.data() + (a - ring_base_); + s.voice.insert(s.voice.end(), src, src + (b - a)); + s.gained_sec += (double)(b - a) / kSr; + const size_t cap = (size_t)((double)opts_.max_voice_sec * kSr); + if (s.voice.size() > cap) s.voice.erase(s.voice.begin(), s.voice.end() - (long)cap); +} + +void SpeakerIdentifier::apply(Slot& s, const SpeakerMatch& m) { + if (m.name.empty()) { // unknown: keep the current name, break any pending run + s.pending.clear(); + return; + } + if (s.current.name.empty() || m.name == s.current.name) { + s.current = {m.name, m.score}; + s.pending.clear(); + return; + } + if (m.name == s.pending) { // a different name won twice in a row + s.current = {m.name, m.score}; + s.pending.clear(); + } else { + s.pending = m.name; + } +} + +void SpeakerIdentifier::maybe_embed(Slot& s, bool is_last) { + if ((double)s.voice.size() / kSr < (double)opts_.min_voice_sec) return; + if (s.gained_sec <= 0.0) return; + const bool first = !s.embedded; + const bool due = first || s.gained_sec >= (double)opts_.refresh_sec || is_last; + if (!due) return; + std::vector emb; + if (!embed_(s.voice.data(), (int)s.voice.size(), emb)) + throw std::runtime_error("speaker embedding failed"); + s.embedded = true; + s.gained_sec = 0.0; + apply(s, registry_->identify(emb, opts_.accept_threshold, opts_.margin)); +} + +void SpeakerIdentifier::consume(const SpeakerSegment& seg, const std::vector& open, + bool growing) { + Slot& s = slots_[seg.speaker]; // a slot is known as soon as it has any segment + const double from = std::max((double)seg.start, s.consumed_until); + const double to = seg.end; + if (to <= from) return; + std::vector others; + for (const SpeakerSegment& h : history_) + if (h.speaker != seg.speaker) others.push_back({h.start, h.end}); + for (const SpeakerSegment& o : open) + if (o.speaker != seg.speaker) others.push_back({o.start, o.end}); + double cursor = to; + // The tail of a growing segment is taken with min_len 0 so a short clean + // piece at the end can be held back and joined to the audio that follows. + std::vector pieces = clean_intervals({from, to}, others, growing ? 0.0 : kMinPieceSec); + if (growing && !pieces.empty() && pieces.back().end >= to && + pieces.back().end - pieces.back().start < kMinPieceSec) { + cursor = pieces.back().start; + pieces.pop_back(); + } + for (const Interval& iv : pieces) + if (iv.end - iv.start >= kMinPieceSec) add_audio(seg.speaker, iv); + s.consumed_until = std::max(s.consumed_until, cursor); +} + +void SpeakerIdentifier::update(const std::vector& closed, + const std::vector& open, bool is_last) { + for (const SpeakerSegment& c : closed) history_.push_back(c); + for (const SpeakerSegment& c : closed) consume(c, open, false); + for (const SpeakerSegment& o : open) consume(o, open, true); + const double horizon = (double)total_ / kSr - (double)opts_.ring_sec; + history_.erase(std::remove_if(history_.begin(), history_.end(), + [&](const SpeakerSegment& h) { return h.end < horizon; }), + history_.end()); + for (auto& kv : slots_) maybe_embed(kv.second, is_last); +} + +SlotName SpeakerIdentifier::name(int slot) const { + auto it = slots_.find(slot); + return it == slots_.end() ? SlotName{} : it->second.current; +} + +std::map SpeakerIdentifier::names() const { + std::map out; + for (const auto& kv : slots_) out[kv.first] = kv.second.current; + return out; +} + +std::map identify_offline(const std::vector& pcm16k, + const std::vector& segs, + const SpeakerEmbed& embed, const SpeakerRegistry& reg, + const SpeakerIdOpts& opts) { + SpeakerIdOpts o = opts; + o.ring_sec = std::max(o.ring_sec, (float)pcm16k.size() / (float)kSr + 1.0f); // keep the whole recording + o.max_voice_sec = std::max(o.max_voice_sec, 30.0f); // a whole recording can afford more voice + o.ring_sec = std::max(o.ring_sec, o.max_voice_sec); + SpeakerIdentifier id(embed, ®, o); + id.push_pcm(pcm16k.data(), (int)pcm16k.size()); + id.update(segs, {}, true); + return id.names(); +} + +} // namespace pk diff --git a/src/speaker_identifier.hpp b/src/speaker_identifier.hpp new file mode 100644 index 0000000..1ad7dbd --- /dev/null +++ b/src/speaker_identifier.hpp @@ -0,0 +1,112 @@ +#pragma once +#include "diarization.hpp" // pk::SpeakerSegment +#include "speaker_registry.hpp" + +#include +#include +#include +#include + +namespace pk { + +struct SpeakerIdOpts { + float min_voice_sec = 2.0f; // clean audio a slot needs before it is embedded + float refresh_sec = 3.0f; // new clean audio that triggers another embedding + float max_voice_sec = 10.0f; // a slot's newest audio kept for embedding + float accept_threshold = 0.5f; // minimum cosine to take a name + float margin = 0.05f; // best must beat the runner-up by this much + float ring_sec = 60.0f; // PCM history kept to slice segments out of +}; + +// "" when valid, else what is wrong. All of: min_voice_sec > 0, refresh_sec > 0, +// max_voice_sec >= min_voice_sec, ring_sec >= max_voice_sec, accept_threshold in +// [-1, 1], margin >= 0. +std::string validate_speaker_opts(const SpeakerIdOpts& o); + +// A diarization slot's current identity. Empty name = unknown. +struct SlotName { + std::string name; + float score = 0.0f; +}; + +// Embeds one 16 kHz mono window. False on failure. +using SpeakerEmbed = std::function& emb)>; + +struct Interval { + double start; + double end; +}; + +// The parts of `seg` not covered by any of `others`, each at least `min_len` long. +std::vector clean_intervals(const Interval& seg, const std::vector& others, + double min_len); + +// Names diarization slots by embedding each slot's clean (single speaker) audio +// and matching it against a registry. Drive it with the stream's PCM and the +// diarizer's segments; it never talks to a model itself. Not thread-safe. +class SpeakerIdentifier { +public: + // `registry` is borrowed and must outlive the identifier. + SpeakerIdentifier(SpeakerEmbed embed, const SpeakerRegistry* registry, SpeakerIdOpts opts); + + // Appends 16 kHz mono PCM (the same audio diarization sees). + void push_pcm(const float* pcm, int n); + // `closed`: segments that closed since the last call. `open`: segments still + // open now, with `end` at the diarizer's current position (`frames_done`). + // Both are consumed: each slot keeps a cursor (consumed_until), and a call + // adds the clean audio of [max(start, cursor), end] for every closed and + // open segment of that slot, then moves the cursor to `end`. So a slot that + // talks without a pause is embedded while its segment is still open, and + // audio taken while a segment was open is not added again when it closes. + // A short clean piece (under 0.2 s) at the growing end of an open segment is + // held back and taken with the audio that follows it. + // Clean means not overlapped by another slot's closed segment (kept for + // ring_sec) or open segment. Overlap with any segment that started before + // the current position is masked, because such a segment is in `open` or + // `closed` already. Known limit: audio up to a cursor is never re-examined, + // so a segment of another slot reported later with a start earlier than + // that cursor is not masked retroactively. The streaming diarizer never does + // this (it marks an onset at the frame where it happens), but a caller that + // reports onsets late would mix that overlap in. + // Contract: `open` must list every segment that has started and not yet + // closed. Throws std::runtime_error when the embed callback fails. + void update(const std::vector& closed, const std::vector& open, + bool is_last); + + SlotName name(int slot) const; // unknown for a slot never seen + std::map names() const; // every slot that has been seen + +private: + struct Slot { + std::vector voice; // newest clean audio, at most max_voice_sec + double gained_sec = 0.0; // clean audio added since the last embedding + bool embedded = false; + SlotName current; + std::string pending; // a different known name that won last time + double consumed_until = 0.0; // seconds of this slot's segments already taken + }; + + void add_audio(int slot, const Interval& iv); + void consume(const SpeakerSegment& seg, const std::vector& open, bool growing); + void maybe_embed(Slot& s, bool is_last); + void apply(Slot& s, const SpeakerMatch& m); + + SpeakerEmbed embed_; + const SpeakerRegistry* registry_; + SpeakerIdOpts opts_; + std::vector ring_; + long long ring_base_ = 0; // absolute sample index of ring_[0] + long long total_ = 0; // samples pushed so far + std::vector history_; // recent closed segments, for overlap checks + std::map slots_; +}; + +// Names slots for a finished recording: runs the same logic once over all +// segments. Returns every slot that appears in `segs`. It raises max_voice_sec +// to at least 30 s and ring_sec to cover the whole recording. +std::map identify_offline(const std::vector& pcm16k, + const std::vector& segs, + const SpeakerEmbed& embed, const SpeakerRegistry& reg, + const SpeakerIdOpts& opts); + +} // namespace pk diff --git a/src/speaker_registry.cpp b/src/speaker_registry.cpp new file mode 100644 index 0000000..ccc1e3d --- /dev/null +++ b/src/speaker_registry.cpp @@ -0,0 +1,162 @@ +#include "speaker_registry.hpp" + +#include +#include +#include +#include + +namespace pk { + +namespace { + +constexpr char kMagic[4] = {'P', 'K', 'S', 'R'}; +constexpr uint32_t kVersion = 1; +constexpr uint32_t kMaxSpeakers = 1u << 20; +constexpr uint32_t kMaxNameLen = 4096; +constexpr int kMaxDim = 1 << 16; + +// L2-normalized copy; empty when the norm is 0 or a value is NaN or Inf. +std::vector normalized(const std::vector& v) { + double n2 = 0.0; + for (float x : v) { + if (!std::isfinite(x)) return {}; + n2 += (double)x * x; + } + if (!(n2 > 0.0) || !std::isfinite(n2)) return {}; + const float inv = (float)(1.0 / std::sqrt(n2)); + std::vector out(v.size()); + for (size_t i = 0; i < v.size(); ++i) out[i] = v[i] * inv; + return out; +} + +void put(std::string& s, const void* p, size_t n) { s.append(static_cast(p), n); } + +struct Reader { + const std::string& s; + size_t pos = 0; + void get(void* dst, size_t n) { + if (n > s.size() - pos) throw std::runtime_error("speaker registry: truncated"); + std::memcpy(dst, s.data() + pos, n); + pos += n; + } +}; + +} // namespace + +std::vector SpeakerRegistry::names() const { + std::vector out; + out.reserve(entries_.size()); + for (const Entry& e : entries_) out.push_back(e.name); + return out; +} + +void SpeakerRegistry::enroll(const std::string& name, const std::vector& emb) { + if (name.empty()) throw std::invalid_argument("speaker name is empty"); + if (emb.empty()) throw std::invalid_argument("speaker embedding is empty"); + if (dim_ != 0 && (int)emb.size() != dim_) + throw std::invalid_argument("speaker embedding has " + std::to_string(emb.size()) + + " values, registry expects " + std::to_string(dim_)); + const std::vector n = normalized(emb); + if (n.empty()) throw std::invalid_argument("speaker embedding is all zero or not finite"); + if (dim_ == 0) dim_ = (int)emb.size(); + for (Entry& e : entries_) { + if (e.name != name) continue; + for (size_t i = 0; i < n.size(); ++i) e.sum[i] += n[i]; + ++e.count; + return; + } + entries_.push_back({name, n, 1}); +} + +bool SpeakerRegistry::remove(const std::string& name) { + for (size_t i = 0; i < entries_.size(); ++i) { + if (entries_[i].name == name) { + entries_.erase(entries_.begin() + (long)i); + return true; + } + } + return false; +} + +SpeakerMatch SpeakerRegistry::identify(const std::vector& emb, float accept, + float margin) const { + if (dim_ != 0 && (int)emb.size() != dim_) + throw std::invalid_argument("speaker embedding has " + std::to_string(emb.size()) + + " values, registry expects " + std::to_string(dim_)); + SpeakerMatch out; + const std::vector q = normalized(emb); + if (q.empty() || entries_.empty()) return out; + float best = -2.0f, second = -2.0f; + const Entry* best_e = nullptr; + for (const Entry& e : entries_) { + const std::vector c = normalized(e.sum); + if (c.empty()) continue; // enrollments that cancel out exactly + float dot = 0.0f; + for (size_t i = 0; i < q.size(); ++i) dot += q[i] * c[i]; + if (dot > best) { second = best; best = dot; best_e = &e; } + else if (dot > second) { second = dot; } + } + if (!best_e) return out; + out.score = best; + if (best < accept) return out; + if (entries_.size() >= 2 && second > -2.0f && best - second < margin) return out; + out.name = best_e->name; + return out; +} + +std::string SpeakerRegistry::serialize() const { + std::string s; + put(s, kMagic, 4); + const uint32_t ver = kVersion; + put(s, &ver, 4); + const int32_t dim = dim_; + put(s, &dim, 4); + const uint32_t n = (uint32_t)entries_.size(); + put(s, &n, 4); + for (const Entry& e : entries_) { + const uint32_t len = (uint32_t)e.name.size(); + put(s, &len, 4); + put(s, e.name.data(), e.name.size()); + const int32_t count = e.count; + put(s, &count, 4); + put(s, e.sum.data(), e.sum.size() * sizeof(float)); + } + return s; +} + +SpeakerRegistry SpeakerRegistry::deserialize(const std::string& blob) { + Reader r{blob}; + char magic[4]; + r.get(magic, 4); + if (std::memcmp(magic, kMagic, 4) != 0) + throw std::runtime_error("speaker registry: bad magic"); + uint32_t ver = 0; + r.get(&ver, 4); + if (ver != kVersion) throw std::runtime_error("speaker registry: unsupported version"); + int32_t dim = 0; + r.get(&dim, 4); + uint32_t n = 0; + r.get(&n, 4); + if (dim < 0 || dim > kMaxDim || n > kMaxSpeakers || (dim < 1 && n > 0)) + throw std::runtime_error("speaker registry: implausible header"); + SpeakerRegistry out(dim); + for (uint32_t i = 0; i < n; ++i) { + uint32_t len = 0; + r.get(&len, 4); + if (len == 0 || len > kMaxNameLen) throw std::runtime_error("speaker registry: bad name"); + Entry e; + e.name.resize(len); + r.get(&e.name[0], len); + int32_t count = 0; + r.get(&count, 4); + if (count < 1) throw std::runtime_error("speaker registry: bad count"); + e.count = count; + e.sum.resize((size_t)dim); + r.get(e.sum.data(), (size_t)dim * sizeof(float)); + out.entries_.push_back(std::move(e)); + } + if (r.pos != blob.size()) throw std::runtime_error("speaker registry: trailing bytes"); + return out; +} + +} // namespace pk diff --git a/src/speaker_registry.hpp b/src/speaker_registry.hpp new file mode 100644 index 0000000..d96d576 --- /dev/null +++ b/src/speaker_registry.hpp @@ -0,0 +1,54 @@ +#pragma once +#include +#include + +namespace pk { + +// Result of matching one embedding against the registry. An empty `name` means +// unknown; `score` is then the best cosine seen (0 for an empty registry). +struct SpeakerMatch { + std::string name; + float score = 0.0f; +}; + +// Enrolled speakers. Each speaker is the L2-normalized mean of the L2-normalized +// embeddings enrolled under its name (a centroid), so enrolling more clips +// tightens the match. Model-independent: it never sees audio. Not thread-safe. +class SpeakerRegistry { +public: + // dim 0 means "fixed by the first enrollment". + explicit SpeakerRegistry(int dim = 0) : dim_(dim) {} + + int dim() const { return dim_; } + size_t size() const { return entries_.size(); } + std::vector names() const; // enrollment order + + // Throws std::invalid_argument: empty name, empty or all-zero embedding, + // or a size different from dim() when dim() != 0. + void enroll(const std::string& name, const std::vector& emb); + bool remove(const std::string& name); + + // Best speaker by cosine. Known only if the best score >= accept and, when + // two or more speakers are enrolled, it beats the runner-up by >= margin. + // Throws std::invalid_argument on a size different from dim() (when dim() + // is set). An all-zero embedding is unknown, not an error. + SpeakerMatch identify(const std::vector& emb, float accept, float margin) const; + + // Binary blob: "PKSR", u32 version 1, i32 dim, u32 n, then per speaker + // u32 name length, name bytes, i32 count, dim x f32 sum. Little-endian. + std::string serialize() const; + // Throws std::runtime_error on bad magic or version, truncation, an absurd + // count, or trailing bytes. + static SpeakerRegistry deserialize(const std::string& blob); + +private: + struct Entry { + std::string name; + std::vector sum; // sum of L2-normalized enrollments + int count = 0; + }; + int dim_; + std::vector entries_; +}; + +} // namespace pk diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index b143508..245eb84 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -74,6 +74,18 @@ pk_add_test(test_diarization) pk_add_test(test_diarization_accuracy) pk_add_test(test_sas_merge) pk_add_test(test_asr_committer) +pk_add_test(test_speaker_registry) +pk_add_test(test_speaker_identifier) +pk_add_test(test_write_atomic) +pk_add_test(test_speaker_encoder) +target_compile_definitions(test_speaker_encoder PRIVATE PK_SOURCE_DIR="${CMAKE_SOURCE_DIR}") +set_tests_properties(test_speaker_encoder PROPERTIES LABELS "model") +pk_add_test(test_speaker_identify) +target_compile_definitions(test_speaker_identify PRIVATE PK_SOURCE_DIR="${CMAKE_SOURCE_DIR}") +set_tests_properties(test_speaker_identify PROPERTIES LABELS "model") +pk_add_test(test_capi_speaker) +target_compile_definitions(test_capi_speaker PRIVATE PK_SOURCE_DIR="${CMAKE_SOURCE_DIR}") +set_tests_properties(test_capi_speaker PROPERTIES LABELS "model") pk_add_test(test_combined_offline) pk_add_test(test_streaming_diarization) pk_add_test(test_sound_stream) diff --git a/tests/test_capi_speaker.cpp b/tests/test_capi_speaker.cpp new file mode 100644 index 0000000..435a879 --- /dev/null +++ b/tests/test_capi_speaker.cpp @@ -0,0 +1,228 @@ +// Speaker identification through the flat C-API. +// PARAKEET_TEST_DIAR_GGUF, PARAKEET_TEST_VD_GGUF (both required, else skip 77) +// PARAKEET_TEST_GGUF optional ASR GGUF; adds the named speaker-attributed ASR check +#include "parakeet_capi.h" + +#include "audio_io.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include + +static int failures = 0; +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL: %s (line %d)\n", #cond, __LINE__); \ + ++failures; \ + } \ + } while (0) + +static std::vector slice(const std::vector& x, double a, double b) { + return std::vector(x.begin() + (long)(a * 16000), x.begin() + (long)(b * 16000)); +} + +int main() { + const char* diar_path = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + const char* vd_path = std::getenv("PARAKEET_TEST_VD_GGUF"); + if (!diar_path || !vd_path) return 77; + + CHECK(parakeet_capi_abi_version() == 9); + + parakeet_ctx* spk = parakeet_capi_load(vd_path); + if (!spk) { std::fprintf(stderr, "FAIL: load speaker model (built without voice-detect?)\n"); return 77; } + CHECK(parakeet_capi_model_kind(spk) == PARAKEET_MODEL_KIND_SPEAKER); + CHECK(parakeet_capi_speaker_dim(spk) >= 128); + parakeet_ctx* diar = parakeet_capi_load(diar_path); + CHECK(diar && parakeet_capi_speaker_dim(diar) == -1); + CHECK(parakeet_capi_model_kind(diar) == PARAKEET_MODEL_KIND_DIARIZATION); + + pk::Audio wav; + if (!pk::load_audio_16k_mono(std::string(PK_SOURCE_DIR) + "/tests/fixtures/two_speakers.wav", wav)) return 1; + const auto a0 = slice(wav.samples, 0.6, 4.6), b0 = slice(wav.samples, 6.9, 10.9); + + parakeet_speaker_registry* reg = parakeet_capi_speaker_registry_new(); + CHECK(parakeet_capi_speaker_enroll(reg, spk, "speaker_a", a0.data(), (int)a0.size(), 16000) == 0); + CHECK(parakeet_capi_speaker_enroll(reg, spk, "speaker_b", b0.data(), (int)b0.size(), 16000) == 0); + CHECK(parakeet_capi_speaker_registry_size(reg) == 2); + // Errors: empty name, no audio, wrong ctx kind. + CHECK(parakeet_capi_speaker_enroll(reg, spk, "", a0.data(), (int)a0.size(), 16000) != 0); + CHECK(parakeet_capi_speaker_enroll(reg, spk, "x", nullptr, 0, 16000) != 0); + CHECK(parakeet_capi_speaker_enroll(reg, diar, "x", a0.data(), (int)a0.size(), 16000) != 0); + CHECK(parakeet_capi_speaker_registry_size(reg) == 2); + + // One-shot identify of clips not used for enrollment. + { + const auto a1 = slice(wav.samples, 14.9, 18.7); + char* j = parakeet_capi_speaker_identify_pcm_json(reg, spk, a1.data(), (int)a1.size(), 16000); + CHECK(j && std::strstr(j, "\"name\":\"speaker_a\"") != nullptr); + parakeet_capi_free_string(j); + const auto b1 = slice(wav.samples, 20.2, 23.5); + j = parakeet_capi_speaker_identify_pcm_json(reg, spk, b1.data(), (int)b1.size(), 16000); + CHECK(j && std::strstr(j, "\"name\":\"speaker_b\"") != nullptr); + parakeet_capi_free_string(j); + } + + // Save and load round trip, and a corrupt file is refused without crashing. + const std::filesystem::path tmp_dir = std::filesystem::temp_directory_path(); + const std::string path = (tmp_dir / "pk_test_capi_speaker_registry.bin").string(); + std::filesystem::remove(path); + CHECK(parakeet_capi_speaker_registry_save(reg, path.c_str()) == 0); + parakeet_speaker_registry* back = parakeet_capi_speaker_registry_load(path.c_str()); + CHECK(back && parakeet_capi_speaker_registry_size(back) == 2); + parakeet_capi_speaker_registry_free(back); + // Saving over an existing registry replaces it, and leaves no tmp file. + { + parakeet_speaker_registry* one = parakeet_capi_speaker_registry_new(); + CHECK(parakeet_capi_speaker_enroll(one, spk, "only_a", a0.data(), (int)a0.size(), 16000) == 0); + CHECK(parakeet_capi_speaker_registry_save(one, path.c_str()) == 0); + CHECK(!std::filesystem::exists(path + ".tmp")); + parakeet_speaker_registry* again = parakeet_capi_speaker_registry_load(path.c_str()); + CHECK(again && parakeet_capi_speaker_registry_size(again) == 1); + parakeet_capi_speaker_registry_free(again); + // A failed save (directory that does not exist) reports an error and + // leaves the file that was there loadable. + const std::string bad = (tmp_dir / "pk_test_no_such_dir" / "registry.bin").string(); + CHECK(parakeet_capi_speaker_registry_save(reg, bad.c_str()) != 0); + CHECK(std::strlen(parakeet_capi_speaker_registry_last_error(reg)) > 0); + parakeet_speaker_registry* still = parakeet_capi_speaker_registry_load(path.c_str()); + CHECK(still && parakeet_capi_speaker_registry_size(still) == 1); + parakeet_capi_speaker_registry_free(still); + parakeet_capi_speaker_registry_free(one); + } + { FILE* f = std::fopen(path.c_str(), "wb"); std::fputs("garbage", f); std::fclose(f); } + CHECK(parakeet_capi_speaker_registry_load(path.c_str()) == nullptr); + CHECK(parakeet_capi_speaker_registry_load((tmp_dir / "pk_test_no_such_dir" / "registry.bin").string().c_str()) == nullptr); + std::filesystem::remove(path); + + // Scene stream with diarization + speaker. + { + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, reg, &o); + CHECK(s != nullptr); + std::string last; + const int chunk = 3200, n = (int)wav.samples.size(); + for (int lo = 0; lo < n && s; lo += chunk) { + const int len = std::min(chunk, n - lo); + char* j = parakeet_capi_scene_stream_feed_json(s, wav.samples.data() + lo, len, lo + len >= n); + CHECK(j != nullptr); + if (j) { last = j; parakeet_capi_free_string(j); } + } + CHECK(last.find("\"names\":{") != std::string::npos); + CHECK(last.find("\"name\":\"speaker_a\"") != std::string::npos); + CHECK(last.find("\"name\":\"speaker_b\"") != std::string::npos); + parakeet_capi_scene_stream_free(s); + } + // A speaker ctx without diarization, or a registry missing, is refused with an error on the ctx. + { + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + CHECK(parakeet_capi_scene_stream_begin_speaker(nullptr, nullptr, nullptr, spk, reg, &o) == nullptr); + CHECK(std::strlen(parakeet_capi_last_error(spk)) > 0); + CHECK(parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, nullptr, &o) == nullptr); + // Old entry point still works exactly as before. + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin(nullptr, diar, nullptr, &o); + CHECK(s != nullptr); + parakeet_capi_scene_stream_free(s); + } + // An old-sized opts struct (no speaker fields) is accepted and uses defaults. + { + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + o.size = (int)offsetof(parakeet_scene_opts, speaker_accept_threshold); + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, reg, &o); + CHECK(s != nullptr); + parakeet_capi_scene_stream_free(s); + } + // Invalid speaker fields are ignored when `size` does not cover them, and + // rejected when it does. + { + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + o.speaker_refresh_sec = -1.0f; + o.speaker_min_voice_sec = -1.0f; + o.speaker_accept_threshold = 5.0f; + o.size = (int)offsetof(parakeet_scene_opts, speaker_accept_threshold); + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, reg, &o); + CHECK(s != nullptr); + parakeet_capi_scene_stream_free(s); + o.size = (int)sizeof(o); + CHECK(parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, reg, &o) == nullptr); + CHECK(std::strstr(parakeet_capi_last_error(spk), "invalid speaker options") != nullptr); + } + // Memory safety: a caller built against the v8 header owns a buffer that ends + // at `flags`. Nothing past it may be read (run under AddressSanitizer). + { + struct OldOpts { + int size; + int diar_latency; + parakeet_sound_opts sound; + int flags; + }; + parakeet_scene_opts full; + parakeet_capi_scene_opts_default(&full); + OldOpts* old = (OldOpts*)std::malloc(sizeof(OldOpts)); + CHECK(old != nullptr); + if (old) { + old->size = (int)sizeof(OldOpts); + old->diar_latency = full.diar_latency; + old->sound = full.sound; + old->flags = 0; + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin_speaker( + nullptr, diar, nullptr, spk, reg, (const parakeet_scene_opts*)old); + CHECK(s != nullptr); + parakeet_capi_scene_stream_free(s); + std::free(old); + } + } + // A registry built by a model with a different embedding size is refused with a + // clear error. Needs a second speaker GGUF of another size (for example WeSpeaker + // 256 versus CAM++ 192): PARAKEET_TEST_VD_GGUF_ALT. Skipped when unset. + if (const char* alt_path = std::getenv("PARAKEET_TEST_VD_GGUF_ALT")) { + parakeet_ctx* alt = parakeet_capi_load(alt_path); + CHECK(alt != nullptr); + if (alt && parakeet_capi_speaker_dim(alt) != parakeet_capi_speaker_dim(spk)) { + parakeet_speaker_registry* wrong = parakeet_capi_speaker_registry_new(); + CHECK(parakeet_capi_speaker_enroll(wrong, alt, "x", a0.data(), (int)a0.size(), 16000) == 0); + // identify with the main model against the alt-sized registry + CHECK(parakeet_capi_speaker_identify_pcm_json(wrong, spk, a0.data(), (int)a0.size(), 16000) == nullptr); + CHECK(std::strstr(parakeet_capi_last_error(spk), "expects") != nullptr); + // and a scene stream refuses it up front + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + CHECK(parakeet_capi_scene_stream_begin_speaker(nullptr, diar, nullptr, spk, wrong, &o) == nullptr); + CHECK(std::strstr(parakeet_capi_last_error(spk), "embeddings") != nullptr); + parakeet_capi_speaker_registry_free(wrong); + } + parakeet_capi_free(alt); + } + + // Named speaker-attributed ASR (optional: needs an ASR model). + if (const char* asr_path = std::getenv("PARAKEET_TEST_GGUF")) { + parakeet_ctx* asr = parakeet_capi_load(asr_path); + CHECK(asr != nullptr); + char* j = parakeet_capi_transcribe_and_diarize_named_json(asr, diar, spk, reg, wav.samples.data(), + (int)wav.samples.size(), 16000); + CHECK(j != nullptr); + if (j) { + CHECK(std::strstr(j, "\"name\":\"speaker_a\"") != nullptr); + CHECK(std::strstr(j, "\"name\":\"speaker_b\"") != nullptr); + CHECK(std::strstr(j, "\"names\":{") != nullptr); + parakeet_capi_free_string(j); + } + parakeet_capi_free(asr); + } + + parakeet_capi_speaker_registry_free(reg); + parakeet_capi_free(diar); + parakeet_capi_free(spk); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_capi_speaker: PASS\n"); + return 0; +} diff --git a/tests/test_sas_merge.cpp b/tests/test_sas_merge.cpp index e409593..7182d45 100644 --- a/tests/test_sas_merge.cpp +++ b/tests/test_sas_merge.cpp @@ -226,6 +226,24 @@ static void test_multiple_words_same_speaker() { CHECK(utts[0].conf == 0.7f); // min } +// Names travel with words into utterances +static void test_names_grouped() { + std::vector w = { + {0, "hello", 0.0f, 0.3f, 0.9f}, {0, "there", 0.35f, 0.6f, 0.8f}, {1, "hi", 1.0f, 1.2f, 0.9f}, + }; + w[0].name = "alice"; w[0].name_score = 0.71f; + w[1].name = "alice"; w[1].name_score = 0.71f; + w[2].name = ""; // slot 1 unknown + auto u = group_speaker_words(w); + CHECK(u.size() == 2); + CHECK(u[0].name == "alice"); + CHECK(u[0].name_score > 0.7f && u[0].name_score < 0.72f); + CHECK(u[1].name.empty()); + // Aggregate initialization without the new fields still compiles and defaults them. + SpeakerWord plain{0, "x", 0.0f, 0.1f, 0.5f}; + CHECK(plain.name.empty() && plain.name_score == 0.0f); +} + int main() { test_basic_assignment(); test_dominant_speaker(); @@ -238,6 +256,7 @@ int main() { test_empty(); test_boundary(); test_multiple_words_same_speaker(); + test_names_grouped(); if (failures == 0) { std::printf("All SAS merge tests passed.\n"); diff --git a/tests/test_scene_render.cpp b/tests/test_scene_render.cpp index b9bb681..c26a5aa 100644 --- a/tests/test_scene_render.cpp +++ b/tests/test_scene_render.cpp @@ -2,11 +2,104 @@ #include "scene_stream.hpp" #include +#include +#include using namespace pk; static int failures = 0; #define CHECK(c) do { if (!(c)) { std::fprintf(stderr, "FAIL: %s (line %d)\n", #c, __LINE__); ++failures; } } while (0) +// Speaker names replace "Speaker N" when known +static void test_render_names() { + pk::SceneUpdate u; + pk::SpeakerUtterance a{0, "hello there", 0.0f, 1.0f, 0.9f}; + a.name = "alice"; + pk::SpeakerUtterance b{1, "hi", 2.0f, 2.5f, 0.9f}; // slot 1 still unknown + u.utterances = {a, b}; + u.safe_until = 10.0; + pk::SceneRenderer r(/*has_diar=*/true, /*show_speech=*/false, nullptr, /*has_asr=*/true); + r.add(u); + const auto lines = r.flush(10.0); + CHECK(lines.size() == 2); + CHECK(lines[0].find("alice: hello there") != std::string::npos); + CHECK(lines[0].find("Speaker") == std::string::npos); + CHECK(lines[1].find("Speaker 1: hi") != std::string::npos); + + // Speaker-only lines use the slot name too. + pk::SceneUpdate d; + d.speakers = {{0, 1.0f, 2.0f}, {1, 3.0f, 4.0f}}; + d.names = {{0, {"alice", 0.7f}}}; + pk::SceneRenderer dz(true, false, nullptr, /*has_asr=*/false); + dz.add(d); + const auto dl = dz.flush_all(); + CHECK(dl.size() == 2 && dl[0] == "[00:01.0 - 00:02.0] alice" && dl[1] == "[00:03.0 - 00:04.0] Speaker 1"); +} + +static void test_json_names() { + pk::SceneUpdate u; + u.t = 3.0; + pk::SpeakerUtterance a{0, "hello", 0.0f, 1.0f, 0.9f}; + a.name = "alice"; a.name_score = 0.71f; + u.utterances = {a}; + u.speakers = {{0, 0.0f, 1.0f}}; + u.active_speakers = {{1, 2.0f, 3.0f}}; + u.names = {{0, {"alice", 0.71f}}, {1, {"", 0.0f}}}; + const std::string j = pk::scene_update_to_json(u, [](int) -> const char* { return nullptr; }); + CHECK(j.find("\"names\":{\"0\":{\"name\":\"alice\",\"score\":0.7100},\"1\":{\"name\":\"\",\"score\":0.0000}}") != std::string::npos); + CHECK(j.find("\"utterances\":[{\"speaker\":0,\"name\":\"alice\",\"name_score\":0.7100,\"text\":\"hello\"") != std::string::npos); + CHECK(j.find("\"speakers\":[{\"speaker\":0,\"name\":\"alice\",\"name_score\":0.7100,\"start\"") != std::string::npos); + CHECK(j.find("\"active\":{\"speakers\":[{\"speaker\":1,\"name\":\"\",\"name_score\":0.0000,\"start\":2.000}") != std::string::npos); + // With no names the document keeps today's exact shape. + u.names.clear(); + const std::string plain = pk::scene_update_to_json(u, [](int) -> const char* { return nullptr; }); + CHECK(plain.find("\"name") == std::string::npos); +} + + +// A fixed update with every field the JSON writer prints. +static pk::SceneUpdate fixed_update() { + pk::SceneUpdate u; + u.t = 2.5; + u.safe_until = 1.0; + u.utterances = {{0, "hello there", 0.25f, 1.5f, 0.875f}}; + u.words = {{0, "hello", 0.25f, 0.75f, 0.9f}, {0, "there", 0.8f, 1.5f, 0.875f}}; + u.speakers = {{0, 0.2f, 1.6f}}; + u.sounds = {{359, 0.0f, 0.96f, 0.81f}}; + u.active_speakers = {{1, 2.0f, 2.5f}}; + u.active_sounds = {{0, 0.5f, 2.5f, 0.9f}}; + return u; +} + +static const char* fixed_label(int i) { return i == 0 ? "Speech" : i == 359 ? "Knock" : "Other"; } + +// Without a speaker part the document is byte for byte what it was before the +// `named` flag. The expected string was produced by the serializer at the +// commit before that change. +static void test_json_golden_unnamed() { + const pk::SceneUpdate u = fixed_update(); + CHECK(!u.named); + const std::string j = pk::scene_update_to_json(u, fixed_label); + if (std::getenv("PK_PRINT_GOLDEN")) std::printf("%s\n", j.c_str()); + const std::string want = + R"({"t":2.500,"utterances":[{"speaker":0,"text":"hello there","start":0.250,"end":1.500,"conf":0.8750}],"words":[{"text":"hello","start":0.250,"end":0.750,"conf":0.9000,"speaker":0},{"text":"there","start":0.800,"end":1.500,"conf":0.8750,"speaker":0}],"speakers":[{"speaker":0,"start":0.200,"end":1.600}],"sounds":[{"index":359,"label":"Knock","start":0.000,"end":0.960,"peak":0.8100}],"active":{"speakers":[{"speaker":1,"start":2.000}],"sounds":[{"index":0,"label":"Speech","start":0.500,"end":2.500,"peak":0.9000}]}})"; + CHECK(j == want); +} + +// With a speaker part the shape is fixed from the first document: "names" is +// present (possibly empty) and every utterance, word and speaker carries +// name and name_score. +static void test_json_named_empty() { + pk::SceneUpdate u = fixed_update(); + u.named = true; + const std::string j = pk::scene_update_to_json(u, fixed_label); + CHECK(j.find("{\"t\":2.500,\"names\":{},\"utterances\":") == 0); + CHECK(j.find("\"utterances\":[{\"speaker\":0,\"name\":\"\",\"name_score\":0.0000,\"text\":\"hello there\"") != std::string::npos); + CHECK(j.find("{\"text\":\"hello\",\"start\":0.250,\"end\":0.750,\"conf\":0.9000,\"speaker\":0,\"name\":\"\",\"name_score\":0.0000}") != std::string::npos); + CHECK(j.find("{\"text\":\"there\",\"start\":0.800,\"end\":1.500,\"conf\":0.8750,\"speaker\":0,\"name\":\"\",\"name_score\":0.0000}") != std::string::npos); + CHECK(j.find("\"speakers\":[{\"speaker\":0,\"name\":\"\",\"name_score\":0.0000,\"start\":0.200") != std::string::npos); + CHECK(j.find("\"active\":{\"speakers\":[{\"speaker\":1,\"name\":\"\",\"name_score\":0.0000,\"start\":2.000}") != std::string::npos); +} + int main() { auto label = [](int i) -> const char* { return i == 0 ? "Speech" : i == 359 ? "Knock" : i == 42 ? "Speech synthesizer" : "Other"; @@ -76,6 +169,11 @@ int main() { auto wl = withasr.flush_all(); CHECK(wl.size() == 1 && wl[0] == "[00:03.0 - 00:03.5] (Knock 0.70)"); + test_render_names(); + test_json_names(); + test_json_golden_unnamed(); + test_json_named_empty(); + if (failures) return 1; std::fprintf(stderr, "PASS\n"); return 0; diff --git a/tests/test_speaker_encoder.cpp b/tests/test_speaker_encoder.cpp new file mode 100644 index 0000000..d073ecb --- /dev/null +++ b/tests/test_speaker_encoder.cpp @@ -0,0 +1,120 @@ +// pk::SpeakerEncoder against a real voice-detect GGUF. +// +// PARAKEET_TEST_VD_GGUF speaker encoder GGUF (required, else skip 77) +// PARAKEET_TEST_VD_REF_WAV optional: a WAV whose reference embedding is in ... +// PARAKEET_TEST_VD_REF_JSON ... this file, the output of +// `voicedetect-cli embed --model --input --json` +// from a standalone voice-detect.cpp build. When both are +// set the folded encoder must match it (cosine >= 0.9999). +// +// The functional check uses tests/fixtures/two_speakers.wav (LibriSpeech 1272 and +// 2086, A-B-A-B). NeMo's segments for it: A 0.50-5.52 and 14.78-18.75, B 6.85-10.82 +// and 20.10-23.60. Two clips of the same voice must score higher than two clips of +// different voices. +#include "audio_io.hpp" +#include "speaker_encoder.hpp" + +#include +#include +#include +#include +#include +#include +#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) + +static double cosine(const std::vector& a, const std::vector& b) { + double d = 0, na = 0, nb = 0; + for (size_t i = 0; i < a.size(); ++i) { d += (double)a[i] * b[i]; na += (double)a[i] * a[i]; nb += (double)b[i] * b[i]; } + return d / std::sqrt(na * nb); +} + +static std::vector slice(const std::vector& x, double a, double b) { + return std::vector(x.begin() + (long)(a * 16000), x.begin() + (long)(b * 16000)); +} + +// Reads "embedding":[...] from voicedetect-cli --json output. +static std::vector read_ref_json(const std::string& path) { + std::ifstream f(path); + std::stringstream ss; + ss << f.rdbuf(); + const std::string s = ss.str(); + const size_t p = s.find("\"embedding\":["); + std::vector v; + if (p == std::string::npos) return v; + const char* c = s.c_str() + p + std::strlen("\"embedding\":["); + char* end = nullptr; + while (*c && *c != ']') { + v.push_back(std::strtof(c, &end)); + c = end; + if (*c == ',') ++c; + } + return v; +} + +int main() { + const char* gguf = std::getenv("PARAKEET_TEST_VD_GGUF"); + if (!gguf) return 77; + if (!SpeakerEncoder::available()) { std::printf("built without PARAKEET_WITH_VOICEDETECT\n"); return 77; } + CHECK(gguf_is_voicedetect(gguf)); + CHECK(!gguf_is_voicedetect("/nonexistent.gguf")); + + auto enc = SpeakerEncoder::load(gguf); + if (!enc) { std::fprintf(stderr, "FAIL: load %s\n", gguf); return 1; } + CHECK(enc->dim() >= 128 && enc->dim() <= 1024); + + Audio wav; + if (!load_audio_16k_mono(std::string(PK_SOURCE_DIR) + "/tests/fixtures/two_speakers.wav", wav)) { + std::fprintf(stderr, "FAIL: load two_speakers.wav\n"); + return 1; + } + auto emb = [&](double a, double b) { + std::vector e; + const auto pcm = slice(wav.samples, a, b); + if (!enc->embed(pcm.data(), (int)pcm.size(), e)) { std::fprintf(stderr, "embed failed: %s\n", enc->last_error().c_str()); ++failures; } + return e; + }; + const auto a1 = emb(0.6, 5.4), a2 = emb(14.9, 18.7), b1 = emb(6.9, 10.7), b2 = emb(20.2, 23.5); + CHECK((int)a1.size() == enc->dim()); + double n2 = 0; + for (float x : a1) n2 += (double)x * x; + CHECK(std::fabs(n2 - 1.0) < 1e-3); // L2-normalized + const double same = 0.5 * (cosine(a1, a2) + cosine(b1, b2)); + const double diff = 0.25 * (cosine(a1, b1) + cosine(a1, b2) + cosine(a2, b1) + cosine(a2, b2)); + std::printf("same-speaker cosine %.3f, different-speaker cosine %.3f\n", same, diff); + CHECK(same > diff + 0.1); + + // Empty and tiny inputs fail cleanly instead of crashing. + std::vector e; + CHECK(!enc->embed(nullptr, 0, e)); + CHECK(!enc->last_error().empty()); + + const char* ref_wav = std::getenv("PARAKEET_TEST_VD_REF_WAV"); + const char* ref_json = std::getenv("PARAKEET_TEST_VD_REF_JSON"); + if (ref_wav && ref_json) { + Audio r; + CHECK(load_audio_16k_mono(ref_wav, r)); + std::vector got; + CHECK(enc->embed(r.samples.data(), (int)r.samples.size(), got)); + const auto want = read_ref_json(ref_json); + CHECK(want.size() == got.size()); + const double c = want.size() == got.size() ? cosine(got, want) : 0.0; + std::printf("folded vs standalone cosine %.6f\n", c); + CHECK(c >= 0.9999); + } + + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_speaker_encoder: PASS\n"); + return 0; +} diff --git a/tests/test_speaker_identifier.cpp b/tests/test_speaker_identifier.cpp new file mode 100644 index 0000000..af6ba67 --- /dev/null +++ b/tests/test_speaker_identifier.cpp @@ -0,0 +1,375 @@ +// Unit test for pk::SpeakerIdentifier with a fake embedder. No model or audio. +// +// Fake audio: speaker k is a constant sample value 0.1*(k+1). The fake embedder +// maps the mean sample value back to a one-hot vector, so a clip that mixes two +// speakers (mean 0.15) lands between two voices and would be caught by the checks. +#include "speaker_identifier.hpp" + +#include +#include +#include +#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) + +static const int kSr = 16000; + +struct Fake { + int calls = 0; + int last_n = 0; + std::vector ns; // every n the embedder was called with, in order + SpeakerEmbed fn() { + return [this](const float* pcm, int n, std::vector& emb) { + ++calls; + last_n = n; + ns.push_back(n); + double sum = 0.0; + for (int i = 0; i < n; ++i) sum += pcm[i]; + const double mean = n ? sum / n : 0.0; + const int hot = (int)std::lround(mean * 10.0) - 1; // 0.1 -> 0, 0.2 -> 1, 0.3 -> 2 + emb.assign(4, 0.0f); + if (hot >= 0 && hot < 4 && std::fabs(mean * 10.0 - std::round(mean * 10.0)) < 0.2) + emb[(size_t)hot] = 1.0f; + else + emb[3] = 1.0f; // ambiguous audio (a mixture) goes to a voice nobody enrolled + return true; + }; + } +}; + +static SpeakerRegistry make_registry() { + SpeakerRegistry r; + r.enroll("alice", {1, 0, 0, 0}); + r.enroll("bob", {0, 1, 0, 0}); + return r; +} + +// Builds a PCM stream of `total_sec` where each listed segment adds its speaker's value. +struct Seg { int spk; float start, end; }; +static std::vector make_pcm(double total_sec, const std::vector& segs) { + std::vector pcm((size_t)(total_sec * kSr), 0.0f); + for (const Seg& s : segs) + for (int i = (int)(s.start * kSr); i < (int)(s.end * kSr) && i < (int)pcm.size(); ++i) + pcm[(size_t)i] += 0.1f * (float)(s.spk + 1); + return pcm; +} + +static SpeakerIdOpts opts() { + SpeakerIdOpts o; + o.min_voice_sec = 2.0f; + o.refresh_sec = 3.0f; + o.max_voice_sec = 10.0f; + return o; +} + +static void test_clean_intervals() { + auto len = [](const std::vector& v) { double t = 0; for (auto& i : v) t += i.end - i.start; return t; }; + auto a = clean_intervals({0, 4}, {{3, 6}}, 0.2); + CHECK(a.size() == 1 && std::fabs(a[0].start) < 1e-9 && std::fabs(a[0].end - 3) < 1e-9); + auto b = clean_intervals({0, 4}, {{1, 2}}, 0.2); + CHECK(b.size() == 2 && std::fabs(len(b) - 3) < 1e-9); + CHECK(clean_intervals({0, 4}, {{0, 4}}, 0.2).empty()); + CHECK(clean_intervals({0, 4}, {{0.1, 4}}, 0.2).empty()); // 0.1 s sliver dropped + auto c = clean_intervals({0, 4}, {{2, 5}, {1, 3}}, 0.2); // overlapping others merge + CHECK(c.size() == 1 && std::fabs(c[0].end - 1) < 1e-9); + auto d = clean_intervals({0, 4}, {}, 0.2); + CHECK(d.size() == 1 && std::fabs(d[0].end - 4) < 1e-9); +} + +static void test_names_two_speakers() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(12, {{0, 0, 5}, {1, 6, 11}}); + id.push_pcm(pcm.data(), 5 * kSr); + id.update({{0, 0.0f, 5.0f}}, {}, false); + CHECK(id.name(0).name == "alice"); + CHECK(id.name(1).name.empty()); + id.push_pcm(pcm.data() + 5 * kSr, 6 * kSr); + id.update({{1, 6.0f, 11.0f}}, {}, false); + CHECK(id.name(1).name == "bob"); + CHECK(id.names().size() == 2); + CHECK(id.name(7).name.empty()); // a slot never seen is unknown, not an error +} + +static void test_min_voice() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(3, {{0, 0, 1.5}}); + id.push_pcm(pcm.data(), 3 * kSr); + id.update({{0, 0.0f, 1.5f}}, {}, true); // 1.5 s < min_voice_sec, even at end of stream + CHECK(f.calls == 0); + CHECK(id.name(0).name.empty()); +} + +static void test_refresh_and_last() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(12, {{0, 0, 2.5}, {0, 3, 4}, {0, 5, 6.5}, {0, 7, 8}, {0, 9, 10}}); + auto feed_to = [&](int from, int to) { id.push_pcm(pcm.data() + from * kSr, (to - from) * kSr); }; + feed_to(0, 3); id.update({{0, 0.0f, 2.5f}}, {}, false); + CHECK(f.calls == 1 && id.name(0).name == "alice"); // first time past min_voice + feed_to(3, 5); id.update({{0, 3.0f, 4.0f}}, {}, false); + CHECK(f.calls == 1); // gained 1.0 s < refresh 3 s + feed_to(5, 7); id.update({{0, 5.0f, 6.5f}}, {}, false); + CHECK(f.calls == 1); // gained 2.5 s + feed_to(7, 9); id.update({{0, 7.0f, 8.0f}}, {}, false); + CHECK(f.calls == 2); // gained 3.5 s >= 3 s + feed_to(9, 11); id.update({{0, 9.0f, 10.0f}}, {}, true); + CHECK(f.calls == 3); // end of stream flushes the 1.0 s gained +} + +static void test_overlap_skipped() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + // Speaker 0 talks 0-4 s, speaker 1 talks 3-6 s (still open when 0 closes). + const auto pcm = make_pcm(6, {{0, 0, 4}, {1, 3, 6}}); + id.push_pcm(pcm.data(), 4 * kSr); + id.update({{0, 0.0f, 4.0f}}, {{1, 3.0f, 4.0f}}, false); + CHECK(f.calls == 1); + CHECK(std::abs(f.last_n - 3 * kSr) <= 2); // only the 3 s that speaker 0 had alone + CHECK(id.name(0).name == "alice"); // mixing 3-4 s would have given mean 0.3 -> unknown +} + +static void test_overlap_same_call_close() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(6, {{0, 0, 4}, {1, 3, 6}}); + id.push_pcm(pcm.data(), 6 * kSr); + id.update({{1, 3.0f, 6.0f}, {0, 0.0f, 4.0f}}, {}, false); // both close in one call + CHECK(f.ns.size() == 2); + if (f.ns.size() == 2) { + CHECK(std::abs(f.ns[0] - 3 * kSr) <= 2); // slot 0 (map order): 0-3 s alone + CHECK(std::abs(f.ns[1] - 2 * kSr) <= 2); // slot 1: 4-6 s alone + } + CHECK(id.name(0).name == "alice"); + CHECK(id.name(1).name == "bob"); +} + +static void test_overlap_earlier_call_close() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(6, {{0, 0, 4}, {1, 3, 6}}); + id.push_pcm(pcm.data(), 6 * kSr); + // Slot 0 still open: slot 1 keeps 4-6 s, and slot 0's open 0-4 s is taken + // now (open segments are consumed), minus the 3-4 s that slot 1 shares. + id.update({{1, 3.0f, 6.0f}}, {{0, 0.0f, 4.0f}}, false); + CHECK(id.name(1).name == "bob"); + CHECK(f.ns.size() == 2); + if (f.ns.size() == 2) { + CHECK(std::abs(f.ns[0] - 3 * kSr) <= 2); // slot 0 (map order): 3-4 s overlap excluded + CHECK(std::abs(f.ns[1] - 2 * kSr) <= 2); // slot 1: 4-6 s alone + } + CHECK(id.name(0).name == "alice"); + id.update({{0, 0.0f, 4.0f}}, {}, false); // closing adds nothing new: 0-4 s was consumed + CHECK(f.ns.size() == 2); + CHECK(id.name(0).name == "alice"); +} + +static void test_unknown_voice() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(5, {{2, 0, 4}}); // a third voice nobody enrolled + id.push_pcm(pcm.data(), 5 * kSr); + id.update({{2, 0.0f, 4.0f}}, {}, true); + CHECK(f.calls >= 1); + CHECK(id.name(2).name.empty()); +} + +static void test_hysteresis() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdOpts o = opts(); + o.max_voice_sec = 3.0f; // the buffer holds only the newest 3 s, so each refresh sees one voice + SpeakerIdentifier id(f.fn(), ®, o); + // Slot 0 is alice's audio first, then bob's audio arrives on the same slot. + const auto pcm = make_pcm(30, {{0, 0, 3}, {1, 3, 6}, {1, 6, 9}, {1, 9, 12}}); + auto feed_to = [&](int from, int to) { id.push_pcm(pcm.data() + from * kSr, (to - from) * kSr); }; + feed_to(0, 3); id.update({{0, 0.0f, 3.0f}}, {}, false); + CHECK(id.name(0).name == "alice"); + feed_to(3, 6); id.update({{0, 3.0f, 6.0f}}, {}, false); + CHECK(id.name(0).name == "alice"); // bob won once: only pending + feed_to(6, 9); id.update({{0, 6.0f, 9.0f}}, {}, false); + CHECK(id.name(0).name == "bob"); // bob won twice in a row +} + +static void test_hysteresis_reset_by_unknown() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdOpts o = opts(); + o.max_voice_sec = 3.0f; + SpeakerIdentifier id(f.fn(), ®, o); + const auto pcm = make_pcm(30, {{0, 0, 3}, {1, 3, 6}, {2, 6, 9}, {1, 9, 12}}); + auto feed_to = [&](int from, int to) { id.push_pcm(pcm.data() + from * kSr, (to - from) * kSr); }; + feed_to(0, 3); id.update({{0, 0.0f, 3.0f}}, {}, false); + feed_to(3, 6); id.update({{0, 3.0f, 6.0f}}, {}, false); // bob pending + feed_to(6, 9); id.update({{0, 6.0f, 9.0f}}, {}, false); // unknown voice: breaks the run + feed_to(9, 12); id.update({{0, 9.0f, 12.0f}}, {}, false); // bob again, but only once in a row + CHECK(id.name(0).name == "alice"); +} + +static void test_ring_drops_old_audio() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdOpts o = opts(); + o.ring_sec = 5.0f; + o.max_voice_sec = 4.0f; // ring_sec must be >= max_voice_sec or the constructor rejects the options + SpeakerIdentifier id(f.fn(), ®, o); + const auto pcm = make_pcm(20, {{0, 0, 3}}); + id.push_pcm(pcm.data(), 20 * kSr); + id.update({{0, 0.0f, 3.0f}}, {}, false); // its audio is 17 s old and gone from the ring + CHECK(f.calls == 0); // no audio, no embedding, no crash + CHECK(id.name(0).name.empty()); +} + +static void test_embed_failure_throws() { + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id([](const float*, int, std::vector&) { return false; }, ®, opts()); + const auto pcm = make_pcm(5, {{0, 0, 4}}); + id.push_pcm(pcm.data(), 5 * kSr); + bool threw = false; + try { id.update({{0, 0.0f, 4.0f}}, {}, false); } catch (const std::runtime_error&) { threw = true; } + CHECK(threw); +} + +static void test_offline() { + Fake f; + const SpeakerRegistry reg = make_registry(); + const auto pcm = make_pcm(24, {{0, 0, 5}, {1, 6, 11}, {0, 12, 17}, {1, 18, 23}}); + const std::vector segs = {{0, 0, 5}, {1, 6, 11}, {0, 12, 17}, {1, 18, 23}}; + const auto names = identify_offline(pcm, segs, f.fn(), reg, opts()); + CHECK(names.size() == 2); + CHECK(names.at(0).name == "alice"); + CHECK(names.at(1).name == "bob"); +} + + +// F1: a slot talking without a pause is named while its segment is still open. +static void test_open_segment_named_before_close() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(8, {{0, 0, 8}}); + for (int t = 1; t <= 5; ++t) { + id.push_pcm(pcm.data() + (t - 1) * kSr, kSr); + id.update({}, {{0, 0.0f, (float)t}}, false); + if (t == 1) CHECK(f.calls == 0 && id.name(0).name.empty()); // 1 s < min_voice + if (t == 2) CHECK(f.calls == 1 && id.name(0).name == "alice"); // named while still open + if (t == 3 || t == 4) CHECK(f.calls == 1); // gained < refresh + } + CHECK(id.names().size() == 1); + CHECK(f.ns.size() == 2); + if (f.ns.size() == 2) { + CHECK(std::abs(f.ns[0] - 2 * kSr) <= 2); // first embedding at 2.0 s consumed + CHECK(std::abs(f.ns[1] - 5 * kSr) <= 2); // refresh after 3 s more + } + // The segment closes at 6 s: only 5-6 s is new, 0-5 s was consumed while open. + id.push_pcm(pcm.data() + 5 * kSr, kSr); + id.update({{0, 0.0f, 6.0f}}, {}, true); + CHECK(f.ns.size() == 3); + if (f.ns.size() == 3) CHECK(std::abs(f.ns[2] - 6 * kSr) <= 2); // 6 s, not 11 s + CHECK(id.name(0).name == "alice"); +} + +// A slot is known (listed by names()) as soon as it has an open segment. +static void test_open_slot_is_known() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(2, {{1, 0, 2}}); + id.push_pcm(pcm.data(), kSr); + id.update({}, {{1, 0.0f, 1.0f}}, false); + CHECK(id.names().size() == 1 && id.names().count(1) == 1); + CHECK(id.name(1).name.empty()); +} + +// Overlap with another slot that opened while this one was open is masked. +static void test_open_segment_overlap_masked() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + // Slot 0 talks 0-6 s, slot 1 talks over it 3-4 s. + const auto pcm = make_pcm(6, {{0, 0, 6}, {1, 3, 4}}); + id.push_pcm(pcm.data(), 3 * kSr); + id.update({}, {{0, 0.0f, 3.0f}}, false); + CHECK(f.ns.size() == 1); + if (!f.ns.empty()) CHECK(std::abs(f.ns[0] - 3 * kSr) <= 2); + id.push_pcm(pcm.data() + 3 * kSr, kSr); + id.update({}, {{0, 0.0f, 4.0f}, {1, 3.0f, 4.0f}}, false); // 3-4 s is shared: nothing added + id.push_pcm(pcm.data() + 4 * kSr, kSr); + id.update({{1, 3.0f, 4.0f}}, {{0, 0.0f, 5.0f}}, false); // 4-5 s alone + id.push_pcm(pcm.data() + 5 * kSr, kSr); + id.update({{0, 0.0f, 6.0f}}, {}, true); // 5-6 s alone, end flush + CHECK(f.ns.size() == 2); + if (f.ns.size() == 2) CHECK(std::abs(f.ns[1] - 5 * kSr) <= 2); // 0-3 + 4-6, 3-4 excluded + CHECK(id.name(0).name == "alice"); // mixing 3-4 s in would give mean 0.133, an unknown voice + CHECK(id.name(1).name.empty()); // slot 1 never had clean audio +} + +// A short clean tail at the growing end of an open segment is not lost: it is +// taken again with the audio that follows it. +static void test_open_short_tail_not_lost() { + Fake f; + const SpeakerRegistry reg = make_registry(); + SpeakerIdentifier id(f.fn(), ®, opts()); + const auto pcm = make_pcm(3, {{0, 0, 3}}); + // 0.1 s per update: every step alone is shorter than the 0.2 s minimum piece. + for (int k = 1; k <= 25; ++k) { + id.push_pcm(pcm.data() + (k - 1) * (kSr / 10), kSr / 10); + id.update({}, {{0, 0.0f, 0.1f * (float)k}}, false); + } + CHECK(f.calls == 1); + if (!f.ns.empty()) CHECK(f.ns[0] >= 2 * kSr); + CHECK(id.name(0).name == "alice"); +} + +static void test_validate_opts() { + CHECK(validate_speaker_opts(SpeakerIdOpts{}).empty()); + SpeakerIdOpts o; + o.min_voice_sec = 0.0f; CHECK(!validate_speaker_opts(o).empty()); + o = SpeakerIdOpts{}; o.refresh_sec = -1.0f; CHECK(!validate_speaker_opts(o).empty()); + o = SpeakerIdOpts{}; o.max_voice_sec = 1.0f; CHECK(!validate_speaker_opts(o).empty()); // < min_voice + o = SpeakerIdOpts{}; o.accept_threshold = 1.5f; CHECK(!validate_speaker_opts(o).empty()); + o = SpeakerIdOpts{}; o.margin = -0.1f; CHECK(!validate_speaker_opts(o).empty()); + o = SpeakerIdOpts{}; o.ring_sec = 5.0f; CHECK(!validate_speaker_opts(o).empty()); // < max_voice +} + +int main() { + test_clean_intervals(); + test_names_two_speakers(); + test_min_voice(); + test_refresh_and_last(); + test_overlap_skipped(); + test_overlap_same_call_close(); + test_overlap_earlier_call_close(); + test_unknown_voice(); + test_hysteresis(); + test_hysteresis_reset_by_unknown(); + test_ring_drops_old_audio(); + test_embed_failure_throws(); + test_offline(); + test_validate_opts(); + test_open_segment_named_before_close(); + test_open_slot_is_known(); + test_open_segment_overlap_masked(); + test_open_short_tail_not_lost(); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_speaker_identifier: PASS\n"); + return 0; +} diff --git a/tests/test_speaker_identify.cpp b/tests/test_speaker_identify.cpp new file mode 100644 index 0000000..1091137 --- /dev/null +++ b/tests/test_speaker_identify.cpp @@ -0,0 +1,199 @@ +// Speaker identification end to end: real diarization + real speaker encoder. +// +// PARAKEET_TEST_DIAR_GGUF diarization GGUF (required, else skip 77) +// PARAKEET_TEST_VD_GGUF speaker encoder GGUF (required, else skip 77) +// PARAKEET_TEST_GGUF ASR GGUF (optional: enables the named-utterance block) +// +// Enrolls the two voices of tests/fixtures/two_speakers.wav (LibriSpeech 1272 = A, +// 2086 = B) from their first turns, then streams the whole file and checks the +// diarization slots get the right names. The enrollment clips come from inside the +// streamed segments (same recording), so on its own this is parity-style evidence. +// What makes it discriminating: the voices are enrolled in the reverse of their +// arrival order (so slot i cannot map to registry entry i), and a registry that +// lacks voice A must leave slot 0 unnamed. NeMo's segments for the fixture: A +// 0.50-5.52 and 14.78-18.75, B 6.85-13.49 and 20.10-23.60. +#include "audio_io.hpp" +#include "diarization.hpp" +#include "model.hpp" +#include "scene_stream.hpp" +#include "speaker_encoder.hpp" + +#include +#include +#include +#include +#include +#include +#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) + +static std::vector slice(const std::vector& x, double a, double b) { + return std::vector(x.begin() + (long)(a * 16000), x.begin() + (long)(b * 16000)); +} + +int main() { + const char* diar_path = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + const char* vd_path = std::getenv("PARAKEET_TEST_VD_GGUF"); + if (!diar_path || !vd_path) return 77; + if (!SpeakerEncoder::available()) return 77; + + auto diar = DiarizationModel::load(diar_path); + auto enc = SpeakerEncoder::load(vd_path); + if (!diar || !enc) { std::fprintf(stderr, "FAIL: load models\n"); return 1; } + + Audio wav; + if (!load_audio_16k_mono(std::string(PK_SOURCE_DIR) + "/tests/fixtures/two_speakers.wav", wav)) { + std::fprintf(stderr, "FAIL: load wav\n"); + return 1; + } + + auto a0 = slice(wav.samples, 0.6, 4.6); + auto b0 = slice(wav.samples, 6.9, 10.9); + std::vector e; + + // Streams the file and returns each slot's final name. + auto run = [&](const SpeakerRegistry& reg, const SpeakerIdOpts& opts = SpeakerIdOpts()) { + SceneParts parts; + parts.diar = diar.get(); + parts.speaker_embed = enc->embedder(); + parts.registry = ® + parts.speaker_opts = opts; + SceneStream stream(parts); + std::map names; + const int chunk = 3200; // 200 ms + const int n = (int)wav.samples.size(); + for (int lo = 0; lo < n; lo += chunk) { + const int len = std::min(chunk, n - lo); + const SceneUpdate u = stream.feed(wav.samples.data() + lo, len, lo + len >= n); + for (const auto& kv : u.names) { + names[kv.first] = kv.second.name; + if (std::getenv("PK_TEST_DEBUG")) + std::fprintf(stderr, "DBG t=%.1f slot%d '%s' %.3f\n", u.t, kv.first, + kv.second.name.c_str(), kv.second.score); + } + } + return names; + }; + + // Enroll in reverse arrival order: B's clip first, A's second. + { + SpeakerRegistry reg; + CHECK(enc->embed(b0.data(), (int)b0.size(), e)); reg.enroll("second_voice", e); + CHECK(enc->embed(a0.data(), (int)a0.size(), e)); reg.enroll("first_voice", e); + auto names = run(reg); + // Slot numbers are arrival order: slot 0 is voice A (speaks first), slot 1 is B. + CHECK(names.size() == 2); + CHECK(names[0] == "first_voice"); + CHECK(names[1] == "second_voice"); + } + // Only B enrolled: slot 0 (voice A) must stay unnamed, never get B's name. + { + SpeakerRegistry reg; + CHECK(enc->embed(b0.data(), (int)b0.size(), e)); reg.enroll("second_voice", e); + // ECAPA scored an impostor voice at cosine 0.566 on this fixture, so the + // acceptance threshold is encoder specific and the default is a starting + // point (docs/speaker.md carries the per-encoder numbers). + SpeakerIdOpts strict; + strict.accept_threshold = 0.7f; + auto names = run(reg, strict); + CHECK(names.count(0) == 1 && names[0].empty()); + CHECK(names[1] == "second_voice"); + } + + SpeakerRegistry reg; + CHECK(enc->embed(a0.data(), (int)a0.size(), e)); reg.enroll("first_voice", e); + + // Bad speaker configurations are rejected with the specific message. + auto expect_throw = [&](const SceneParts& p, const char* what) { + try { + SceneStream s(p); + } catch (const std::invalid_argument& ex) { + if (std::string(ex.what()).find(what) != std::string::npos) return; + std::fprintf(stderr, "FAIL: wrong message '%s', wanted '%s'\n", ex.what(), what); + ++failures; + return; + } + std::fprintf(stderr, "FAIL: no throw, wanted '%s'\n", what); + ++failures; + }; + { + SceneParts no_diar; + no_diar.speaker_embed = enc->embedder(); + no_diar.registry = ® + expect_throw(no_diar, "needs a diarization model"); + SceneParts no_reg; + no_reg.diar = diar.get(); + no_reg.speaker_embed = enc->embedder(); + expect_throw(no_reg, "needs a registry"); + SceneParts bad_opts; + bad_opts.diar = diar.get(); + bad_opts.speaker_embed = enc->embedder(); + bad_opts.registry = ® + bad_opts.speaker_opts.min_voice_sec = 0; + expect_throw(bad_opts, "invalid speaker options"); + } + + // With no speaker part the update carries no names (existing behavior). + { + SceneParts plain; + plain.diar = diar.get(); + SceneStream s(plain); + const SceneUpdate u = s.feed(wav.samples.data(), (int)wav.samples.size(), true); + CHECK(u.names.empty()); + } + + // ASR + diarization + speaker: the utterances themselves carry names. Words + // committed before a slot is identified keep their earlier (empty) name, so an + // empty name is allowed; the other voice's name never is. + if (const char* asr_path = std::getenv("PARAKEET_TEST_GGUF")) { + auto asr = Model::load(asr_path); + if (!asr) { std::fprintf(stderr, "FAIL: load asr\n"); return 1; } + SpeakerRegistry areg; + CHECK(enc->embed(b0.data(), (int)b0.size(), e)); areg.enroll("second_voice", e); + CHECK(enc->embed(a0.data(), (int)a0.size(), e)); areg.enroll("first_voice", e); + SceneParts parts; + parts.asr = asr.get(); + parts.diar = diar.get(); + parts.speaker_embed = enc->embedder(); + parts.registry = &areg; + SceneStream stream(parts); + std::vector utts; + const int chunk = 3200; + const int n = (int)wav.samples.size(); + for (int lo = 0; lo < n; lo += chunk) { + const int len = std::min(chunk, n - lo); + const SceneUpdate u = stream.feed(wav.samples.data() + lo, len, lo + len >= n); + utts.insert(utts.end(), u.utterances.begin(), u.utterances.end()); + } + CHECK(!utts.empty()); + int named0 = 0, named1 = 0, wrong = 0; + for (const auto& u : utts) { + if (std::getenv("PK_TEST_DEBUG")) + std::fprintf(stderr, "UTT slot%d '%s' start=%.2f '%s'\n", u.speaker, u.name.c_str(), + u.start, u.text.c_str()); + if (u.speaker == 0 && u.name == "first_voice") ++named0; + if (u.speaker == 1 && u.name == "second_voice") ++named1; + if ((u.speaker == 0 && u.name == "second_voice") || + (u.speaker == 1 && u.name == "first_voice")) + ++wrong; + } + CHECK(named0 >= 1); + CHECK(named1 >= 1); + CHECK(wrong == 0); + } + + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_speaker_identify: PASS\n"); + return 0; +} diff --git a/tests/test_speaker_registry.cpp b/tests/test_speaker_registry.cpp new file mode 100644 index 0000000..7919a0b --- /dev/null +++ b/tests/test_speaker_registry.cpp @@ -0,0 +1,223 @@ +// Unit test for pk::SpeakerRegistry. No model or audio needed. +#include "speaker_registry.hpp" + +#include +#include +#include +#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) + +static std::vector unit(int dim, int hot) { + std::vector v((size_t)dim, 0.0f); + v[(size_t)hot] = 1.0f; + return v; +} + +static void test_enroll_and_identify() { + SpeakerRegistry r; + r.enroll("alice", unit(4, 0)); + r.enroll("bob", unit(4, 1)); + CHECK(r.dim() == 4); + CHECK(r.size() == 2); + SpeakerMatch m = r.identify(unit(4, 0), 0.5f, 0.05f); + CHECK(m.name == "alice"); + CHECK(std::fabs(m.score - 1.0f) < 1e-5f); + m = r.identify(unit(4, 1), 0.5f, 0.05f); + CHECK(m.name == "bob"); +} + +static void test_centroid_averages_enrollments() { + SpeakerRegistry r; + r.enroll("alice", {1.0f, 0.0f}); + r.enroll("alice", {0.0f, 1.0f}); // centroid points at (1,1)/sqrt2 + CHECK(r.size() == 1); + const SpeakerMatch m = r.identify({1.0f, 1.0f}, 0.9f, 0.0f); + CHECK(m.name == "alice"); + CHECK(m.score > 0.999f); +} + +static void test_unknown_below_threshold() { + SpeakerRegistry r; + r.enroll("alice", unit(4, 0)); + const SpeakerMatch m = r.identify(unit(4, 2), 0.5f, 0.05f); // orthogonal: cosine 0 + CHECK(m.name.empty()); + CHECK(std::fabs(m.score) < 1e-5f); +} + +static void test_margin() { + SpeakerRegistry r; + r.enroll("alice", {1.0f, 0.0f}); + r.enroll("bob", {0.0f, 1.0f}); + // Equidistant probe: both score 0.707, so the margin is 0 and it must be unknown. + const SpeakerMatch m = r.identify({1.0f, 1.0f}, 0.5f, 0.05f); + CHECK(m.name.empty()); + CHECK(m.score > 0.7f); + // A single enrolled speaker has no runner-up, so the margin does not apply. + SpeakerRegistry one; + one.enroll("alice", {1.0f, 0.0f}); + CHECK(one.identify({1.0f, 1.0f}, 0.5f, 0.5f).name == "alice"); +} + +static void test_dim_mismatch() { + SpeakerRegistry r; + r.enroll("alice", unit(4, 0)); + bool threw = false; + try { r.enroll("bob", unit(3, 0)); } catch (const std::invalid_argument&) { threw = true; } + CHECK(threw); + threw = false; + try { r.identify(unit(3, 0), 0.5f, 0.05f); } catch (const std::invalid_argument&) { threw = true; } + CHECK(threw); +} + +static void test_bad_enroll() { + SpeakerRegistry r; + auto throws = [&](const std::string& n, const std::vector& e) { + try { r.enroll(n, e); } catch (const std::invalid_argument&) { return true; } + return false; + }; + CHECK(throws("", unit(4, 0))); + CHECK(throws("a", {})); + CHECK(throws("a", {0.0f, 0.0f})); + CHECK(r.size() == 0); + // An all-zero probe is unknown, not an error. + r.enroll("alice", unit(2, 0)); + CHECK(r.identify({0.0f, 0.0f}, 0.5f, 0.05f).name.empty()); + // Failed enroll on non-empty registry leaves state unchanged. + const int orig_dim = r.dim(); + const size_t orig_size = r.size(); + const SpeakerMatch orig_match = r.identify(unit(2, 0), 0.5f, 0.05f); + CHECK(throws("alice", unit(3, 0))); // dim mismatch + CHECK(throws("bob", {0.0f, 0.0f})); // zero vector + CHECK(r.dim() == orig_dim && r.size() == orig_size); + const SpeakerMatch new_match = r.identify(unit(2, 0), 0.5f, 0.05f); + CHECK(new_match.name == orig_match.name && std::fabs(new_match.score - orig_match.score) < 1e-5f); +} + +static void test_remove_and_names() { + SpeakerRegistry r; + r.enroll("bob", unit(2, 1)); + r.enroll("alice", unit(2, 0)); + const auto n = r.names(); + CHECK(n.size() == 2 && n[0] == "bob" && n[1] == "alice"); // enrollment order + CHECK(r.remove("bob")); + CHECK(!r.remove("bob")); + CHECK(r.size() == 1); + // names() still lists remaining speaker in order after remove. + const auto remaining = r.names(); + CHECK(remaining.size() == 1 && remaining[0] == "alice"); +} + +static void test_serialize_roundtrip() { + SpeakerRegistry r; + r.enroll("alice", {1.0f, 0.0f, 0.0f}); + r.enroll("alice", {0.9f, 0.1f, 0.0f}); + r.enroll("bob", {0.0f, 1.0f, 0.0f}); + const SpeakerRegistry back = SpeakerRegistry::deserialize(r.serialize()); + CHECK(back.dim() == 3); + CHECK(back.size() == 2); + const auto a = r.identify({0.95f, 0.05f, 0.0f}, 0.5f, 0.05f); + const auto b = back.identify({0.95f, 0.05f, 0.0f}, 0.5f, 0.05f); + CHECK(a.name == b.name && std::fabs(a.score - b.score) < 1e-6f); + // Enrolling more into the loaded registry keeps averaging (count survived). + SpeakerRegistry loaded = SpeakerRegistry::deserialize(r.serialize()); + loaded.enroll("alice", {0.0f, 0.0f, 1.0f}); + CHECK(loaded.size() == 2); +} + +static void test_deserialize_corrupt() { + SpeakerRegistry r; + r.enroll("alice", {1.0f, 0.0f}); + const std::string good = r.serialize(); + auto throws = [](const std::string& s) { + try { SpeakerRegistry::deserialize(s); } catch (const std::runtime_error&) { return true; } + return false; + }; + CHECK(throws("")); + CHECK(throws("not a registry")); + CHECK(throws(good.substr(0, good.size() - 1))); // truncated + CHECK(throws(good + "x")); // trailing bytes + std::string bad_magic = good; + bad_magic[0] = 'X'; + CHECK(throws(bad_magic)); + std::string huge = good; // absurd speaker count + huge[12] = (char)0xff; huge[13] = (char)0xff; huge[14] = (char)0xff; huge[15] = (char)0x7f; + CHECK(throws(huge)); + // Corrupt: dim 0 with n > 0 (would cause out-of-bounds write on later enroll). + // Build manually: "PKSR" (4 bytes) + version 1 (4 bytes, little-endian) + + // dim 0 (4 bytes) + n 1 (4 bytes) + name length 5 (4 bytes) + "alice" (5 bytes) + + // count 1 (4 bytes) + 0 floats for sum (since dim is 0). + std::string corrupt_dim_zero; + corrupt_dim_zero += 'P'; corrupt_dim_zero += 'K'; corrupt_dim_zero += 'S'; corrupt_dim_zero += 'R'; + // version 1 in little-endian + corrupt_dim_zero += (char)0x01; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; + // dim 0 in little-endian + corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; + // n 1 in little-endian + corrupt_dim_zero += (char)0x01; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; + // name length 5 in little-endian + corrupt_dim_zero += (char)0x05; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; + // name "alice" + corrupt_dim_zero += "alice"; + // count 1 in little-endian + corrupt_dim_zero += (char)0x01; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; corrupt_dim_zero += (char)0x00; + CHECK(throws(corrupt_dim_zero)); + // Empty registry (dim 0, n 0) round-trips and accepts normal enroll. + SpeakerRegistry empty; + const std::string empty_blob = empty.serialize(); + SpeakerRegistry loaded_empty = SpeakerRegistry::deserialize(empty_blob); + CHECK(loaded_empty.dim() == 0 && loaded_empty.size() == 0); + loaded_empty.enroll("charlie", {1.0f, 0.0f}); + CHECK(loaded_empty.size() == 1 && loaded_empty.dim() == 2); +} + +// A NaN or Inf embedding is refused like an all-zero one, and a NaN probe is unknown. +static void test_non_finite() { + SpeakerRegistry r; + r.enroll("alice", unit(2, 0)); + const SpeakerMatch before = r.identify(unit(2, 0), 0.5f, 0.05f); + auto throws = [&](const std::string& n, const std::vector& e, std::string* msg) { + try { r.enroll(n, e); } catch (const std::invalid_argument& ex) { *msg = ex.what(); return true; } + return false; + }; + std::string msg; + CHECK(throws("bob", {std::nanf(""), 1.0f}, &msg)); + CHECK(msg.find("not finite") != std::string::npos); + CHECK(throws("bob", {INFINITY, 0.0f}, &msg)); + CHECK(throws("alice", {-INFINITY, 1.0f}, &msg)); // also for a name already enrolled + CHECK(r.size() == 1 && r.dim() == 2); + const SpeakerMatch after = r.identify(unit(2, 0), 0.5f, 0.05f); + CHECK(after.name == before.name && std::fabs(after.score - before.score) < 1e-6f); + const SpeakerMatch nan_probe = r.identify({std::nanf(""), 1.0f}, 0.5f, 0.05f); + CHECK(nan_probe.name.empty()); + CHECK(r.identify({INFINITY, 0.0f}, 0.5f, 0.05f).name.empty()); + SpeakerRegistry empty; // a failed first enroll does not fix the dimension + try { empty.enroll("x", {std::nanf(""), 0.0f, 0.0f}); } catch (const std::invalid_argument&) {} + CHECK(empty.size() == 0 && empty.dim() == 0); +} + +int main() { + test_enroll_and_identify(); + test_centroid_averages_enrollments(); + test_unknown_below_threshold(); + test_margin(); + test_dim_mismatch(); + test_bad_enroll(); + test_remove_and_names(); + test_serialize_roundtrip(); + test_deserialize_corrupt(); + test_non_finite(); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_speaker_registry: PASS\n"); + return 0; +} diff --git a/tests/test_write_atomic.cpp b/tests/test_write_atomic.cpp new file mode 100644 index 0000000..fcad385 --- /dev/null +++ b/tests/test_write_atomic.cpp @@ -0,0 +1,78 @@ +// Unit test for pk::write_file_atomic. No model needed. +#include "common.hpp" + +#include +#include +#include +#include +#include + +namespace fs = std::filesystem; + +static int failures = 0; +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL: %s (line %d)\n", #cond, __LINE__); \ + ++failures; \ + } \ + } while (0) + +static std::string read_all(const fs::path& p) { + std::ifstream in(p, std::ios::binary); + return std::string(std::istreambuf_iterator(in), std::istreambuf_iterator()); +} + +int main() { + const fs::path dir = fs::temp_directory_path() / "pk_test_write_atomic"; + std::error_code ec; + fs::remove_all(dir, ec); + fs::create_directories(dir); + const fs::path file = dir / "reg.bin"; + const std::string tmp = file.string() + ".tmp"; + + // A new file gets exactly the bytes, including a NUL byte. + const std::string first("first\0bytes", 11); + std::string err; + CHECK(pk::write_file_atomic(file.string(), first, &err)); + CHECK(err.empty()); + CHECK(read_all(file) == first); + CHECK(!fs::exists(tmp)); + + // Overwriting replaces the whole file and leaves no tmp file. + const std::string second = "second, longer than the first one"; + CHECK(pk::write_file_atomic(file.string(), second, &err)); + CHECK(read_all(file) == second); + CHECK(!fs::exists(tmp)); + CHECK(pk::write_file_atomic(file.string(), "x", nullptr)); // err may be null + CHECK(read_all(file) == "x"); + CHECK(pk::write_file_atomic(file.string(), second, nullptr)); + + // A path in a directory that does not exist fails with a reason, and the + // existing file is untouched. + err.clear(); + CHECK(!pk::write_file_atomic((dir / "missing" / "reg.bin").string(), "new", &err)); + CHECK(!err.empty()); + CHECK(read_all(file) == second); + CHECK(!fs::exists(dir / "missing")); + + // A target that is a directory fails, the directory survives, no tmp file. + const fs::path sub = dir / "sub"; + fs::create_directories(sub); + err.clear(); + CHECK(!pk::write_file_atomic(sub.string(), "new", &err)); + CHECK(!err.empty()); + CHECK(fs::is_directory(sub)); + CHECK(!fs::exists(sub.string() + ".tmp")); + CHECK(read_all(file) == second); + + // An empty path is refused. + err.clear(); + CHECK(!pk::write_file_atomic("", "x", &err)); + CHECK(!err.empty()); + + fs::remove_all(dir, ec); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::printf("test_write_atomic: PASS\n"); + return 0; +} diff --git a/third_party/voice-detect.cpp b/third_party/voice-detect.cpp new file mode 160000 index 0000000..b74a896 --- /dev/null +++ b/third_party/voice-detect.cpp @@ -0,0 +1 @@ +Subproject commit b74a896f47c6d04fcca0a962ff317528fd0b0019