diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c844881..6a5c8c9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -32,6 +32,21 @@ 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. + run: | + cmake -B build-noced -DPARAKEET_BUILD_TESTS=ON -DGGML_NATIVE=OFF -DPARAKEET_WITH_CED=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. + # Capture first: steps run with -e -o pipefail. + rc=0 + out=$(build-noced/examples/cli/parakeet-cli scene --sound x --input y 2>&1) || rc=$? + echo "$out" + test "$rc" -eq 2 + grep -q "built without sound tagging" <<< "$out" + # ------------------------------------------------------------------------- # server-e2e: drive the real parakeet-server over HTTP. # diff --git a/.gitmodules b/.gitmodules index 7d6227f..07e9117 100644 --- a/.gitmodules +++ b/.gitmodules @@ -1,3 +1,6 @@ [submodule "third_party/ggml"] path = third_party/ggml url = https://github.com/ggml-org/ggml +[submodule "third_party/ced.cpp"] + path = third_party/ced.cpp + url = https://github.com/localai-org/ced.cpp diff --git a/AGENTS.md b/AGENTS.md index 8f89fd1..c650ce3 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -80,8 +80,14 @@ src/ libparakeet implementation diarization_encoder/head.*, RoPE Transformer encoder + speaker head diarization_streaming.*, NeMo cache-aware streaming diarization (speaker cache + FIFO) sas_merge.hpp/cpp , ASR words x speaker segments -> speaker-attributed utterances + asr_committer.hpp/cpp, shared "finalize a word/utterance once" logic used by SAS and SceneStream + ced_tagger.hpp/cpp , pk::CedTagger: loads a CED GGUF (ced.cpp) into a tagger context, pk::SoundScorer + 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 examples/cli/ parakeet-cli binary - subcommands: info, transcribe (+ --stream), quantize + 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 diarize binary: diarize [--stream] scripts/ Python tooling convert_parakeet_to_gguf.py, .nemo/.hf -> GGUF (--dtype f32|f16|q8_0) @@ -111,6 +117,12 @@ tests/ ctest targets test_streaming_diarization.cpp, streaming diarization == NeMo streaming, every latency mode (same baseline) test_combined_offline.cpp, SAS + streaming diarization/SAS through the C-API test_sas_merge.cpp , SAS merge/grouping (model-independent) + test_asr_committer.cpp , shared word/utterance finalize logic (model-independent) + test_ced_parity.cpp , CedTagger scores == ced.cpp PyTorch baseline (PARAKEET_TEST_CED_GGUF f32 + PARAKEET_TEST_CED_BASELINE) + test_sound_stream.cpp , pk::SoundStream windowing/on-off-min_duration logic (model-independent) + 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) 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 @@ -118,6 +130,10 @@ tests/ ctest targets fixtures/two_speakers.wav, LibriSpeech 1272 + 2086 alternating A-B-A-B, 23.6 s third_party/ vendored deps ggml/ , submodule pinned at v0.13.0 + ced.cpp/ , submodule, CED sound-event tagger (PARAKEET_WITH_CED, on by default); + 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 dr_wav.h , vendored single header models/ output dir for converted GGUFs (gitignored; MANIFEST.md tracks the expected published set) @@ -147,6 +163,7 @@ cmake -B build -DPARAKEET_BUILD_TESTS=ON -DGGML_NATIVE=ON && cmake --build build | `PARAKEET_GGML_METAL` | OFF | Forward GGML_METAL to the submodule | | `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 | Use `-DGGML_NATIVE=OFF` when building for CI or portable binaries. @@ -238,6 +255,7 @@ The binary is at `build/examples/cli/parakeet-cli`. 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] ``` `--timestamps` prints one `- ()` line per word (also @@ -282,6 +300,33 @@ parakeet_capi_free_diar_segments parakeet_capi_sas_stream_begin / _begin_latency / _feed / _free ``` +Sound-event detection (ABI v8, additive; not used by LocalAI yet). A CED GGUF +(ced.cpp) loads into its own `parakeet_ctx` kind (a "tagger") through the same +`parakeet_capi_load`; see `docs/sound.md`: + +``` +parakeet_capi_sound_opts_default +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) +``` + +Combined scene stream (ABI v8, additive; not used by LocalAI yet). One stream +that carries any mix of an ASR context, a diarization context, and a tagger +context, and emits speaker-attributed words/utterances plus sound-event +segments in one time-ordered JSON document per feed; see `docs/sound.md`: + +``` +parakeet_capi_scene_opts_default +parakeet_capi_scene_stream_begin +parakeet_capi_scene_stream_feed_json +parakeet_capi_scene_stream_drain_scores_json +parakeet_capi_scene_stream_last_error +parakeet_capi_scene_stream_free +``` + `parakeet_capi_transcribe_path_json(ctx, wav, decoder)` returns malloc'd UTF-8 JSON `{"text":..,"words":[{"w","start","end","conf"}],"tokens":[{"id","t","conf"}]}` (times in seconds, conf in `(0,1]`), built from diff --git a/CMakeLists.txt b/CMakeLists.txt index 82e6374..cbaa2db 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -17,6 +17,7 @@ option(PARAKEET_GGML_CUDA "Forward GGML_CUDA" OFF) 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) set(GGML_CUDA ${PARAKEET_GGML_CUDA} CACHE BOOL "" FORCE) set(GGML_METAL ${PARAKEET_GGML_METAL} CACHE BOOL "" FORCE) @@ -63,6 +64,33 @@ endif() add_subdirectory(third_party/ggml) +# Single dr_wav implementation for the whole build. src/audio_io.cpp only +# declares dr_wav's functions (#include "dr_wav.h", no DR_WAV_IMPLEMENTATION); +# this object library is the one translation unit that defines them, so both +# libparakeet and, when PARAKEET_WITH_CED is on, ced.cpp (built with +# CED_EXTERNAL_DR_WAV, which likewise only declares them) link against the +# same symbols instead of each pulling in its own copy (a second +# DR_WAV_IMPLEMENTATION define would be a multiple-definition link error). +# An OBJECT library's objects are linked directly into each consumer rather +# than referenced as a separate archive dependency, so this does not create a +# link-time cycle between the parakeet and ced targets. +add_library(dr_wav_impl OBJECT src/dr_wav_impl.cpp) +target_include_directories(dr_wav_impl PRIVATE ${CMAKE_SOURCE_DIR}/third_party) +set_target_properties(dr_wav_impl PROPERTIES POSITION_INDEPENDENT_CODE ON) + +if(PARAKEET_WITH_CED) + if(NOT EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/third_party/ced.cpp/CMakeLists.txt") + message(FATAL_ERROR "third_party/ced.cpp is missing: run `git submodule update --init --recursive`, or configure with -DPARAKEET_WITH_CED=OFF") + endif() + set(CED_BUILD_CLI OFF CACHE BOOL "" FORCE) + set(CED_BUILD_TESTS OFF CACHE BOOL "" FORCE) + set(CED_SHARED OFF CACHE BOOL "" FORCE) + set(CED_EXTERNAL_DR_WAV ON CACHE BOOL "" FORCE) # dr_wav_impl provides it + add_subdirectory(third_party/ced.cpp EXCLUDE_FROM_ALL) + set_target_properties(ced PROPERTIES POSITION_INDEPENDENT_CODE ON) + target_link_libraries(ced PRIVATE dr_wav_impl) +endif() + set(PARAKEET_SRC src/parakeet.cpp src/model.cpp @@ -97,7 +125,13 @@ set(PARAKEET_SRC src/diarization_encoder.cpp src/diarization_head.cpp src/sas_merge.cpp - src/diarization_streaming.cpp) + src/asr_committer.cpp + src/diar_pcm_stream.cpp + src/scene_stream.cpp + src/diarization_streaming.cpp + src/ced_tagger.cpp + src/sound_stream.cpp + src/scene_render.cpp) if(PARAKEET_SHARED) add_library(parakeet SHARED ${PARAKEET_SRC}) @@ -113,6 +147,12 @@ target_include_directories(parakeet PUBLIC include PRIVATE src ${CMAKE_SOURCE_DI target_compile_definitions(parakeet PUBLIC $<$:_USE_MATH_DEFINES>) target_compile_definitions(parakeet PRIVATE PARAKEET_VERSION="${PARAKEET_VERSION}") target_link_libraries(parakeet PUBLIC ggml) +target_link_libraries(parakeet PRIVATE dr_wav_impl) + +if(PARAKEET_WITH_CED) + target_link_libraries(parakeet PRIVATE ced) + target_compile_definitions(parakeet PRIVATE PARAKEET_WITH_CED=1) +endif() if(PARAKEET_BUILD_CLI) add_subdirectory(examples/cli) diff --git a/README.md b/README.md index a71c2ae..0a3503e 100644 --- a/README.md +++ b/README.md @@ -116,6 +116,7 @@ cmake --build build-shared -j | `PARAKEET_GGML_METAL` | OFF | Forward GGML_METAL to the submodule | | `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 | To build for a GPU backend, forward its flag, e.g. Apple Metal: @@ -317,6 +318,27 @@ To batch from code, use the batched entry points (single-clip B=1 is just N=1): --- +## Sound events + +parakeet.cpp can also tag everyday sounds (dog bark, glass breaking, applause, +alarms, music, and the rest of the 527-class AudioSet ontology) with +[CED](https://github.com/RicherMans/CED), through the +[ced.cpp](https://github.com/localai-org/ced.cpp) submodule (`PARAKEET_WITH_CED`, +on by default). `parakeet-cli scene` combines it with ASR and diarization into +one time-ordered feed: + +```sh +parakeet-cli scene --model asr.gguf --diar diar.gguf --sound ced-base-q8_0.gguf \ + --latency low --input audio.wav +[00:00.4 - 00:03.2] Speaker 0: mister Quilter is the apostle of the middle classes, and +[00:24.0 - 00:30.0] (Chicken, rooster 0.86) +``` + +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. + +--- + ## C-API (`libparakeet.so`) `include/parakeet_capi.h` defines a flat, exception-free C-API meant for `dlopen` / FFI / LocalAI integration. Build the shared library with `-DPARAKEET_SHARED=ON`: diff --git a/docs/sound.md b/docs/sound.md new file mode 100644 index 0000000..3e05df4 --- /dev/null +++ b/docs/sound.md @@ -0,0 +1,417 @@ +# Sound-event detection + +parakeet.cpp can tag everyday sounds (dog bark, glass breaking, applause, +alarms, music, speech, and the rest of the 527-class AudioSet ontology) using +[CED](https://github.com/RicherMans/CED) (Consistent Ensemble Distillation), +run through [ced.cpp](https://github.com/localai-org/ced.cpp), the same LocalAI-team +ggml port used standalone. This is separate from ASR: a CED GGUF loads into +its own context (a "tagger"), and `pk::CedTagger` / the sound-stream C-API are +the only code in this repository that talks to it. + +CED GGUFs (tiny, mini, small, base; f16 and q8_0, base also f32) are published +in one collection at +[huggingface.co/mudler/ced-gguf](https://huggingface.co/mudler/ced-gguf). +`ced-tiny` at q8_0 is about 6 MB; `ced-base` at q8_0 is about 88 MB. + +## Loading a tagger + +`parakeet_capi_load` detects the GGUF's architecture and returns a context +that holds either an ASR model, a diarization model, or (for a CED GGUF) a +tagger. Passing a tagger context to the ASR or diarization entry points fails +with a clear "context holds a CED sound model" error, and vice versa. + +```c +parakeet_ctx* tagger = parakeet_capi_load("ced-tiny-q8_0.gguf"); +``` + +On the C++ side this is `pk::CedTagger::load(path)`, which exposes +`n_classes()`, `label(index)`, and a `scorer()` (a `pk::SoundScorer`: scores +one PCM window and fills one probability per class, multi-label, not a +softmax). + +### Device + +ced.cpp picks its own device independently of the rest of parakeet.cpp: on a +GPU build it picks the first GPU it finds and falls back to CPU, and +`CED_DEVICE` overrides that choice the same way `PARAKEET_DEVICE` does for +Parakeet (`CED_DEVICE=cpu` forces CPU; a device name such as `CUDA0`, +`Vulkan0`, or `MTL0` selects that device). `CED_DEVICE` and `PARAKEET_DEVICE` +are read separately, so a tagger and an ASR/diarization context in the same +process can land on different devices. + +## Windowing and timing rules + +`pk::SoundStream` (`src/sound_stream.hpp`) runs a tagger over live PCM as a +sliding window: every `hop_sec` seconds, the most recent `window_sec` seconds +are scored, and a class's score is compared against two thresholds to decide +whether an event is open: + +- A class **opens** when a scored window's score for it is `>= on_threshold`. +- An open class **closes** when a scored window's score drops `< off_threshold`. +- A closed segment shorter than `min_duration_sec` is dropped. + +Because a score only tells you "this class was present somewhere in the last +`window_sec`," segment boundaries follow a **one-hop** grid: the code places +the open boundary at the start of the newest hop inside the triggering window +(the latest boundary consistent with the score), and the close boundary at +the end of the oldest hop inside the triggering window (the earliest boundary +consistent with the score). For a sharp event edge that puts the boundary +within about one `hop_sec` of it. A score that rises or falls slowly crosses +the thresholds later or earlier than the true edge, so it can move a boundary +by more than one hop. + +At end of stream (`is_last`), a final tail window is scored if at least one +CED patch's worth of audio (0.16 s) remains unscored, and every class still +open is closed at the current stream time. + +## Defaults + +```cpp +struct SoundOpts { + float window_sec = 3.0f; + float hop_sec = 1.0f; + float on_threshold = 0.4f; + float off_threshold = 0.3f; + float min_duration_sec = 0.3f; + int top_k = 5; +}; +``` + +CED was trained and evaluated on clips up to about 10 s. A live stream cannot +wait 10 s to say what it is hearing, so the question is how much accuracy a +short window gives up. `examples/cli/sound_window_eval.cpp` (built as +`sound-window-eval`) measures this directly: for each WAV, it takes the +whole-clip (up to 10 s) top-1 class as the reference, then slides windows of +1, 2, 3, and 5 s (hop = window / 2) over the same clip and reports the share +of windows whose top-1 matches the reference, and the share whose top-5 +contains it. + +``` +sound-window-eval # one WAV path per line +``` + +### Measurement + +Data: all 2000 clips of [ESC-50](https://github.com/karoldvl/ESC-50) (5 s, +44.1 kHz, resampled to 16 kHz mono) plus the three +`third_party/ced.cpp/benchmarks/demo/clips` demo clips (rooster, thunder, +guitar; 6 s), for `ced-tiny-q8_0`; ESC-50 folds 1 and 2 (800 clips) plus the +same three demo clips for `ced-base-q8_0`. ESC-50 clips are 5 s, so a 5 s +window is the reference clip itself; that row is trivially 100% and is listed +only for shape, not as a real data point. + +`ced-tiny-q8_0`, 2003 clips: + +| window | windows | top-1 = clip top-1 | clip top-1 in top-5 | +|---|--:|--:|--:| +| 1 s | 18033 | 33.1% | 55.0% | +| 2 s | 8015 | 57.8% | 80.4% | +| 3 s | 4009 | 74.1% | 91.5% | +| 5 s | 2003 | 100.0% | 100.0% | + +`ced-base-q8_0`, 803 clips: + +| window | windows | top-1 = clip top-1 | clip top-1 in top-5 | +|---|--:|--:|--:| +| 1 s | 7233 | 39.2% | 62.3% | +| 2 s | 3215 | 60.7% | 83.5% | +| 3 s | 1609 | 75.1% | 93.0% | +| 5 s | 803 | 100.0% | 100.0% | + +### Decision + +The rule was: keep window 3 s / hop 1 s unless ced-base's 3 s windows agree +with the clip top-1 on fewer than 70% of windows, in which case switch to 5 s +/ hop 1 s. ced-base's 3 s row is 75.1%, above the 70% line, so the defaults +stay at window 3 s / hop 1 s. A 1 s window is noticeably worse (39.2% on +ced-base) and a 2 s window is a real step down too (60.7%); 3 s is the +shortest window that stays reasonably close to what the whole clip would say, +which is the point of picking a default for a live stream instead of always +waiting for the whole clip. + +## Sound-stream C-API + +`include/parakeet_capi.h`, additive since ABI v8: + +```c +typedef struct { + int size; // sizeof(parakeet_sound_opts) + float window_sec, hop_sec; + float on_threshold, off_threshold, min_duration_sec; + int top_k; +} parakeet_sound_opts; +void parakeet_capi_sound_opts_default(parakeet_sound_opts* o); + +typedef struct { + int class_index; + const char* label; // borrowed from the tagger ctx + float start, end, peak; // seconds from stream start +} parakeet_sound_segment; + +typedef struct parakeet_sound_stream parakeet_sound_stream; + +parakeet_sound_stream* parakeet_capi_sound_stream_begin(parakeet_ctx* tagger, + const parakeet_sound_opts* o); +int parakeet_capi_sound_stream_feed(parakeet_sound_stream* s, const float* pcm, int n, + int is_last, parakeet_sound_segment** out, int* n_out); +int parakeet_capi_sound_stream_active(parakeet_sound_stream* s, + parakeet_sound_segment** out, int* n_out); +char* parakeet_capi_sound_stream_drain_scores_json(parakeet_sound_stream* s); +void parakeet_capi_free_sound_segments(parakeet_sound_segment* segs); +void parakeet_capi_sound_stream_free(parakeet_sound_stream* s); + +int parakeet_capi_num_classes(const parakeet_ctx* ctx); +const char* parakeet_capi_class_label(const parakeet_ctx* ctx, int index); + +#define PARAKEET_MODEL_KIND_NONE 0 +#define PARAKEET_MODEL_KIND_ASR 1 +#define PARAKEET_MODEL_KIND_DIARIZATION 2 +#define PARAKEET_MODEL_KIND_SOUND 3 +int parakeet_capi_model_kind(const parakeet_ctx* ctx); +``` + +`parakeet_capi_sound_stream_begin` with `o = NULL` uses the defaults above. +`_feed` accepts any chunk size and returns the segments that closed during +that call (free with `parakeet_capi_free_sound_segments`); pass `is_last = 1` +on the final chunk to flush and close everything still open. `_active` reads +the segments still open right now, with `end` set to the current stream time. +`_drain_scores_json` returns the raw per-window top-k scores scored since the +previous drain, as +`[{"start":..,"end":..,"tags":[{"index":..,"label":..,"score":..}]}]`. + +`parakeet_capi_model_kind` reports which kind of model a `parakeet_capi_load` +context holds (`PARAKEET_MODEL_KIND_NONE` on `NULL`), so a caller loading ASR, +diarization and tagger GGUFs through the same entry point can dispatch without +probing individual functions. +The stream keeps one entry per hop until it is drained, so the queue grows +with the stream: drain regularly, or set `top_k = 0` to keep no scores. + +### C example + +```c +#include "parakeet_capi.h" +#include + +void tag_stream(const char* ced_gguf, const float* pcm, int n_samples) { + parakeet_ctx* tagger = parakeet_capi_load(ced_gguf); + if (!tagger) { fprintf(stderr, "load failed\n"); return; } + + parakeet_sound_stream* s = parakeet_capi_sound_stream_begin(tagger, NULL); + if (!s) { fprintf(stderr, "%s\n", parakeet_capi_last_error(tagger)); parakeet_capi_free(tagger); return; } + + const int chunk = 8000; // 0.5 s at 16 kHz + for (int i = 0; i < n_samples; i += chunk) { + const int n = (i + chunk <= n_samples) ? chunk : n_samples - i; + const int is_last = (i + n >= n_samples); + + parakeet_sound_segment* segs = NULL; + int n_segs = 0; + if (parakeet_capi_sound_stream_feed(s, pcm + i, n, is_last, &segs, &n_segs) != 0) { + fprintf(stderr, "%s\n", parakeet_capi_last_error(tagger)); + break; + } + for (int k = 0; k < n_segs; ++k) + printf("%-24s %6.2f - %6.2f peak %.2f\n", + segs[k].label, segs[k].start, segs[k].end, segs[k].peak); + parakeet_capi_free_sound_segments(segs); + } + + parakeet_capi_sound_stream_free(s); + parakeet_capi_free(tagger); +} +``` + +## `parakeet-cli scene` + +`scene` combines any mix of an ASR model, a diarization model and a CED +tagger into one time-ordered feed. At least one of `--model`, `--diar`, +`--sound` is required, plus `--input`: + +``` +parakeet-cli scene --model --diar --sound \ + --input [--latency model|low|very_low|ultra_low] [--chunk-ms N] \ + [--show-speech] [--json] +``` + +`--latency` picks the diarization streaming mode (see `docs/diarization.md`); +it has no effect without `--diar`. `--chunk-ms` (default 200, capped at 60000) +sets how much PCM is fed to the stream per step. `--show-speech` keeps plain +speech labels (`Speech`, `Speech synthesizer`, `Conversation`, and similar) +in the sound output; by default they are filtered out since they are +redundant with the ASR transcript. `--json` prints one +`parakeet_capi_scene_stream_feed_json` document per step instead of the +rendered text. + +Real output, all three models, `--latency low`, on a demo clip (two +LibriSpeech speakers, a rooster clip, then a second LibriSpeech excerpt): + +``` +$ parakeet-cli scene --model asr.gguf --diar diar.gguf --sound ced-base-q8_0.gguf \ + --latency low --input scene_demo.wav +[00:00.4 - 00:03.2] Speaker 0: mister Quilter is the apostle of the middle classes, and +[00:03.6 - 00:05.4] Speaker 0: we're glad to welcome his gospel. +[00:06.6 - 00:06.7] Speaker 1: Well, +[00:07.2 - 00:09.4] Speaker 1: I don't wish to see it any more, observed Phoebe, +[00:09.7 - 00:10.8] Speaker 1: turning away her eyes +[00:11.4 - 00:12.6] Speaker 1: it is certainly very like +[00:12.9 - 00:13.6] Speaker 1: old portrait. +[00:14.7 - 00:16.2] Speaker 0: Nor is Mr Quilter's +[00:16.4 - 00:18.5] Speaker 0: manner less interesting than his matter. +[00:20.0 - 00:21.5] Speaker 1: Well, I don't wish to see it anymore, +[00:22.0 - 00:23.6] Speaker 1: observed Phoebe, turning away her. +[00:24.0 - 00:30.0] (Fowl 0.58) +[00:24.0 - 00:30.0] (Chicken, rooster 0.86) +[00:25.0 - 00:27.0] (Cluck 0.46) +[00:26.0 - 00:30.0] (Crowing, cock-a-doodle-doo 0.65) +[00:30.0 - 00:32.6] Speaker 1: Well, I don't wish to see it any more, observed Phoebe, +[00:33.0 - 00:36.7] Speaker 1: turning away her eyes it is certainly very like the old portrait +``` + +Each line is `[start - end] ` in stream +order. `Speech synthesizer` is one of the labels filtered out by default +(CED tags clean narration with it at moderate confidence; it is treated as a +plain-speech label alongside `Speech`, `Conversation`, and the rest, so it +does not show up twice next to the transcript). Pass `--show-speech` to see +it. With only `--sound`, the output is the sound lines alone; with only +`--model` (no `--diar`), the utterance lines drop the `Speaker N:` prefix. +With `--diar` but no `--model`, there is no transcript, so each closed +speaker segment prints as `[start - end] Speaker N`, in time order with +the sound lines. A segment line waits while an earlier-starting segment is +still open. + +## Scene stream C-API + +`include/parakeet_capi.h`, additive since ABI v8. One stream carries any mix +of an ASR context, a diarization context and a tagger context (at least one +is required); each `_feed_json` call returns everything the stream finalized +in that call, as one JSON document: + +```c +typedef struct { + int size; // sizeof(parakeet_scene_opts) + 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 +} parakeet_scene_opts; +void parakeet_capi_scene_opts_default(parakeet_scene_opts* o); + +typedef struct parakeet_scene_stream parakeet_scene_stream; + +parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* asr, parakeet_ctx* diar, + parakeet_ctx* tagger, + const parakeet_scene_opts* o); +char* parakeet_capi_scene_stream_feed_json(parakeet_scene_stream* s, const float* pcm, int n, + int is_last); +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); +``` + +The context arguments are borrowed (same lifetime rule as `sas_stream` and +`sound_stream`): free the scene stream before freeing any of the contexts it +was given. Passing `NULL` for a context leaves that part out of the stream; +`diar_latency` and `sound` in `parakeet_scene_opts` are ignored when the +matching context is `NULL`. + +Each `_feed_json` document has the shape: + +```json +{"t":0.600, + "utterances":[{"speaker":0,"text":"mister Quilter is","start":0.4,"end":1.6,"conf":0.98}], + "words":[{"text":"mister","start":0.4,"end":0.6,"conf":0.99,"speaker":0}], + "speakers":[{"speaker":0,"start":0.0,"end":0.6}], + "sounds":[{"index":365,"label":"Chicken, rooster","start":24.0,"end":30.0,"peak":0.86}], + "active":{"speakers":[{"speaker":0,"start":0.6}], + "sounds":[{"index":365,"label":"Chicken, rooster","start":24.0,"end":26.0,"peak":0.7}]}} +``` + +`utterances`, `words` and `speakers` are the closed diarized ASR results for +this call (empty parts if the matching context is `NULL`); `sounds` are +sound-event segments that closed this call; `active` holds the speaker and +sound segments still open, with `end`/`peak` as of the current stream time. +`t` is the stream time consumed so far. `parakeet_capi_scene_stream_drain_scores_json` +returns the same shape `parakeet_capi_sound_stream_drain_scores_json` does +(`"[]"` without a tagger). The same rule applies: drain regularly, or set +`sound.top_k = 0` to keep no scores. + +```c +#include "parakeet_capi.h" +#include + +void run_scene(const char* asr_gguf, const char* diar_gguf, const char* ced_gguf, + const float* pcm, int n_samples) { + parakeet_ctx* asr = parakeet_capi_load(asr_gguf); + parakeet_ctx* diar = diar_gguf ? parakeet_capi_load(diar_gguf) : NULL; + parakeet_ctx* tagger = ced_gguf ? parakeet_capi_load(ced_gguf) : NULL; + + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + o.diar_latency = PARAKEET_DIAR_LATENCY_LOW; + + parakeet_scene_stream* s = parakeet_capi_scene_stream_begin(asr, diar, tagger, &o); + if (!s) { fprintf(stderr, "scene_stream_begin failed\n"); return; } + + const int chunk = 3200; // 200 ms at 16 kHz + for (int i = 0; i < n_samples; i += chunk) { + const int n = (i + chunk <= n_samples) ? chunk : n_samples - i; + const int is_last = (i + n >= n_samples); + + char* doc = parakeet_capi_scene_stream_feed_json(s, pcm + i, n, is_last); + if (!doc) { fprintf(stderr, "%s\n", parakeet_capi_scene_stream_last_error(s)); break; } + printf("%s\n", doc); + parakeet_capi_free_string(doc); + } + + parakeet_capi_scene_stream_free(s); + if (tagger) parakeet_capi_free(tagger); + if (diar) parakeet_capi_free(diar); + parakeet_capi_free(asr); +} +``` + +## Server: `--sound-model` + +`parakeet-server` accepts `--sound-model `, a local path to a CED +GGUF. With it set, a `verbose_json` transcription response gains a +`sound_events` array (`{"label","start","end","score"}` per event); `json` +and `text` responses are unchanged, and the sound pass only runs for +`verbose_json` requests. See `examples/server/README.md` for the full option +list. + +## GPU + +`test_ced_parity`, `test_sound_stream`, `test_sound_capi`, `test_scene_stream`, +`test_combined_offline`, `test_streaming_diarization`, `test_asr_committer` +and `parakeet-cli scene` (all three models, `--latency low`, on the demo clip +used above) were run on three GPU backends, staged and built through the `rc` +fleet (Vulkan on `strix:gpu0`, CUDA on `dgx:gpu0`) and over SSH (Metal on an +M4 Mac). All three matched the CPU transcript of the same build word for +word. Those runs predate the ASR changes that release non-speech audio, +which changed where the last transcript lines above break. Sound scores and +boundaries vary by low single hundredths and by hop-width ordering between +backends, as expected of independent floating-point runs. + +| Device | Backend | `-LE model` | sound/scene ctest set | `test_ced_parity` | Scene demo wall time | +|---|---|---|---|---|---| +| strix:gpu0 (AMD Radeon 8060S, Vulkan0) | Vulkan | 21/22 (`server_e2e` not run) | 8/8 | pass | 2.3 s | +| dgx:gpu0 (NVIDIA GB10, CUDA0) | CUDA | 21/22 (`server_e2e` not run) | 5/8 (3 known teardown crashes, see below) | pass | 2.8 s | +| Apple M4 (MTL0) | Metal | 22/22 | 5/8 (3 known teardown crashes, see below) | pass | 9.5 s | + +One known issue came out of these runs. It is not a regression in this +change: + +- **`test_combined_offline`, `test_streaming_diarization` and + `test_scene_stream` abort on process exit on CUDA and Metal, not Vulkan.** + Each test holds more than one GPU-backed context in one process (ASR + + diarization, or ASR + diarization + a CED tagger). All assertions print + PASS first; the abort comes later, during static teardown. The + process-global backend's allocator is freed by a static destructor after + the GPU context it belongs to is already gone: on CUDA the backtrace is + `ggml_gallocr_free -> ggml_backend_buffer_free -> cudaFree`, which fails + with `CUDA error: driver shutting down`; on Metal, + `ggml_metal_device_free` asserts `[rsets->data count] == 0`. + `test_combined_offline` and `test_streaming_diarization` predate the + sound-events work and abort the same way on the base branch, so this is + not new with the scene tests. `parakeet-cli` does not hit it: it calls + `pk::shutdown_backend()` before it returns from `main`, while the GPU + context is still alive. No test tolerance or code was changed for it. diff --git a/examples/cli/CMakeLists.txt b/examples/cli/CMakeLists.txt index ef14574..2b16184 100644 --- a/examples/cli/CMakeLists.txt +++ b/examples/cli/CMakeLists.txt @@ -5,3 +5,7 @@ target_include_directories(parakeet-cli PRIVATE ${CMAKE_SOURCE_DIR}/src) add_executable(diarize diarize.cpp) target_link_libraries(diarize PRIVATE parakeet) target_include_directories(diarize PRIVATE ${CMAKE_SOURCE_DIR}/include ${CMAKE_SOURCE_DIR}/src) + +add_executable(sound-window-eval sound_window_eval.cpp) +target_link_libraries(sound-window-eval PRIVATE parakeet) +target_include_directories(sound-window-eval PRIVATE ${CMAKE_SOURCE_DIR}/src) diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index 0db3f1a..f2fa537 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -18,6 +18,10 @@ #include "ggml.h" #include "gguf.h" #include "transcription_json.hpp" +#include "diarization.hpp" +#include "ced_tagger.hpp" +#include "scene_stream.hpp" +#include "scene_render.hpp" #include #include #include @@ -1335,6 +1339,163 @@ static int cmd_bench_decode(int argc, char** argv) { return 0; } +static const char* kSceneUsage = + "usage: parakeet-cli scene [--model ] [--diar ] " + "[--sound ] --input " + "[--latency model|low|very_low|ultra_low] [--chunk-ms N] " + "[--show-speech] [--json]\n"; + +// parakeet-cli scene [--model ] [--diar ] [--sound ] +// --input [--latency model|low|very_low|ultra_low] +// [--chunk-ms N] [--show-speech] [--json] +// 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; + bool json = false; + bool show_speech = false; + int chunk_ms = 200; + 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], "--diar") == 0 && i + 1 < argc) { + diar = argv[++i]; + } else if (std::strcmp(argv[i], "--sound") == 0 && i + 1 < argc) { + sound = argv[++i]; + } 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) { + latency_str = argv[++i]; + } else if (std::strcmp(argv[i], "--chunk-ms") == 0 && i + 1 < argc) { + chunk_ms = std::atoi(argv[++i]); + } else if (std::strcmp(argv[i], "--json") == 0) { + json = true; + } else if (std::strcmp(argv[i], "--show-speech") == 0) { + show_speech = true; + } else { + std::fprintf(stderr, "%s", kSceneUsage); + return 2; + } + } + + if ((model.empty() && diar.empty() && sound.empty()) || input.empty()) { + std::fprintf(stderr, "%s", kSceneUsage); + return 2; + } + if (chunk_ms <= 0) { + std::fprintf(stderr, "parakeet-cli scene: --chunk-ms must be > 0\n"); + return 2; + } + // chunk_ms * 16 (samples/ms at 16 kHz) must not overflow int; 60 s is far + // above any sane chunk size (the default is 200 ms) and leaves headroom. + if (chunk_ms > 60000) { + std::fprintf(stderr, "parakeet-cli scene: --chunk-ms must be <= 60000\n"); + return 2; + } + pk::DiarLatency latency = pk::DiarLatency::Model; + if (!latency_str.empty()) { + if (latency_str == "model") { + latency = pk::DiarLatency::Model; + } else if (latency_str == "low") { + latency = pk::DiarLatency::Low; + } else if (latency_str == "very_low") { + latency = pk::DiarLatency::VeryLow; + } else if (latency_str == "ultra_low") { + latency = pk::DiarLatency::UltraLow; + } else { + std::fprintf(stderr, + "parakeet-cli scene: unknown --latency '%s' (want model|low|very_low|ultra_low)\n", + latency_str.c_str()); + return 2; + } + } + if (!sound.empty() && !pk::CedTagger::available()) { + std::fprintf(stderr, "parakeet-cli: built without sound tagging (PARAKEET_WITH_CED=OFF)\n"); + return 2; + } + + std::unique_ptr asr_model; + if (!model.empty()) { + asr_model = pk::Model::load(model); + if (!asr_model) { + std::fprintf(stderr, "parakeet-cli scene: failed to load model %s\n", model.c_str()); + return 1; + } + } + std::unique_ptr diar_model; + if (!diar.empty()) { + diar_model = pk::DiarizationModel::load(diar); + if (!diar_model) { + std::fprintf(stderr, "parakeet-cli scene: failed to load diarization model %s\n", + diar.c_str()); + return 1; + } + } + std::unique_ptr tagger; + if (!sound.empty()) { + tagger = pk::CedTagger::load(sound); + if (!tagger) { + std::fprintf(stderr, "parakeet-cli scene: failed to load sound model %s\n", sound.c_str()); + return 1; + } + } + + pk::Audio audio; + if (!load_audio_arg_16k_mono(input, audio)) { + std::string display = input_display_name(input); + std::fprintf(stderr, "parakeet-cli scene: failed to load audio %s\n", display.c_str()); + return 1; + } + + pk::SceneParts parts; + parts.asr = asr_model.get(); + parts.diar = diar_model.get(); + parts.diar_latency = latency; + parts.tagger = tagger.get(); + + // scene_update_to_json's label(i) may return nullptr (emitted as ""); the + // same lambda drives the renderer's --json-less line formatting. + auto label = [&](int i) -> const char* { return tagger ? tagger->label(i) : nullptr; }; + + std::unique_ptr stream; + try { + stream.reset(new pk::SceneStream(parts)); + } catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-cli scene: %s\n", e.what()); + return 1; + } + + pk::SceneRenderer renderer(diar_model != nullptr, show_speech, label, asr_model != nullptr); + + const int chunk_samples = chunk_ms * 16; // 16 samples/ms at 16 kHz + const int n = (int)audio.samples.size(); + for (int lo = 0; lo < n || lo == 0; lo += chunk_samples) { + const int len = std::min(chunk_samples, n - lo); + const bool is_last = lo + len >= n; + pk::SceneUpdate u; + try { + u = stream->feed(audio.samples.data() + lo, len, is_last); + } catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-cli scene: feed failed: %s\n", e.what()); + return 1; + } + if (json) { + std::printf("%s\n", pk::scene_update_to_json(u, label).c_str()); + } else { + renderer.add(u); + for (const std::string& line : renderer.flush(u.safe_until)) + std::printf("%s\n", line.c_str()); + if (is_last) + for (const std::string& line : renderer.flush_all()) + std::printf("%s\n", line.c_str()); + } + if (is_last) break; + } + return 0; +} + // Run a subcommand, then free the process-global backend while the GPU driver is // still alive (the subcommand's local Model is already destroyed by the time it // returns, releasing its device weight buffer). Avoids the CUDA "driver shutting @@ -1363,6 +1524,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], "scene") == 0) + return run_and_shutdown(cmd_scene, argc - 2, argv + 2); std::fprintf(stderr, "usage:\n" " parakeet-cli info \n" @@ -1377,6 +1540,10 @@ int main(int argc, char** argv) { " parakeet-cli bench-batch --model --manifest " "[--decoder ctc|tdt] [--threads N] [--batch-sizes 1,4,8] [--json ]\n" " parakeet-cli bench-decode --model --audio " - "[--batch-sizes 1,4,8,16] [--threads N] [--reps R] [--json ]\n"); + "[--batch-sizes 1,4,8,16] [--threads N] [--reps R] [--json ]\n" + " parakeet-cli scene [--model ] [--diar ] " + "[--sound ] --input " + "[--latency model|low|very_low|ultra_low] [--chunk-ms N] " + "[--show-speech] [--json]\n"); return 2; } diff --git a/examples/cli/sound_window_eval.cpp b/examples/cli/sound_window_eval.cpp new file mode 100644 index 0000000..e0feb19 --- /dev/null +++ b/examples/cli/sound_window_eval.cpp @@ -0,0 +1,67 @@ +// One-off measurement: how much CED's top-1 tag changes when a clip is scored +// in short windows instead of one window of up to 10 s. For every WAV in the +// list, the reference is the top-1 class of the whole clip (<= 10 s); each +// window size W reports the share of W-second windows (hop W/2) whose top-1 +// equals the reference, and the share whose top-5 contains it. +// usage: sound-window-eval +#include "audio_io.hpp" +#include "ced_tagger.hpp" + +#include +#include +#include +#include +#include +#include + +// Load a WAV as 16 kHz mono (pk::load_audio_16k_mono downmixes and resamples). +static bool load16k(const std::string& path, std::vector& x) { + pk::Audio a; + if (!pk::load_audio_16k_mono(path, a)) return false; + x = std::move(a.samples); + return true; +} + +static std::vector topk(const std::vector& p, int k) { + std::vector i(p.size()); + std::iota(i.begin(), i.end(), 0); + std::partial_sort(i.begin(), i.begin() + k, i.end(), [&](int a, int b) { return p[a] > p[b]; }); + i.resize(k); + return i; +} + +int main(int argc, char** argv) { + if (argc < 3) { std::fprintf(stderr, "usage: sound-window-eval \n"); return 2; } + auto t = pk::CedTagger::load(argv[1]); + if (!t) { std::fprintf(stderr, "load failed\n"); return 1; } + auto score = t->scorer(); + const float sizes[] = {1.0f, 2.0f, 3.0f, 5.0f}; + std::vector top1(4, 0), top5(4, 0), total(4, 0); + std::ifstream list(argv[2]); + std::string path; + int clips = 0; + while (std::getline(list, path)) { + std::vector x; + if (path.empty() || !load16k(path, x)) continue; + x.resize(std::min(x.size(), 10 * 16000)); + std::vector p; + if (!score(x.data(), (int)x.size(), p)) continue; + const int ref = topk(p, 1)[0]; + ++clips; + for (int s = 0; s < 4; ++s) { + const int w = (int)(sizes[s] * 16000), h = w / 2; + for (int a = 0; a + w <= (int)x.size(); a += h) { + if (!score(x.data() + a, w, p)) continue; + auto k = topk(p, 5); + top1[s] += k[0] == ref; + top5[s] += std::find(k.begin(), k.end(), ref) != k.end(); + ++total[s]; + } + } + } + std::printf("clips: %d\n| window | windows | top-1 = clip top-1 | clip top-1 in top-5 |\n|---|--:|--:|--:|\n", clips); + for (int s = 0; s < 4; ++s) + std::printf("| %.0f s | %ld | %.1f%% | %.1f%% |\n", sizes[s], total[s], + 100.0 * top1[s] / std::max(1L, total[s]), 100.0 * top5[s] / std::max(1L, total[s])); + return 0; +} diff --git a/examples/server/README.md b/examples/server/README.md index 99ff275..d2900de 100644 --- a/examples/server/README.md +++ b/examples/server/README.md @@ -88,10 +88,21 @@ with open("audio.wav", "rb") as f: print(client.audio.transcriptions.create(model="parakeet", file=f).text) ``` +## Sound events + +Pass `--sound-model ` with a local path to a ced.cpp sound-event +tagger GGUF. When set, a `verbose_json` response gains a `sound_events` array, +one entry per detected event: `{"label","start","end","score"}`. `json` and +`text` responses are unchanged, and the sound pass only runs for +`verbose_json` requests, so other requests never pay for it. `--sound-model` +takes a local path only, it is not resolved through the alias/URL model +fetcher used by `--model`. + ## Supported - `response_format`: `json` (default), `text`, `verbose_json`. - `timestamp_granularities[]=word` adds a `words` array to `verbose_json`. +- `--sound-model` adds a `sound_events` array to `verbose_json`. ## Known simplifications diff --git a/examples/server/main.cpp b/examples/server/main.cpp index 0e4cbad..2ea0cca 100644 --- a/examples/server/main.cpp +++ b/examples/server/main.cpp @@ -7,6 +7,9 @@ #include "transcription.hpp" // pk::Transcription #include "ggml_graph.hpp" // pk::set_num_threads, pk::shutdown_backend #include "dr_wav.h" // declarations only; impl lives in libparakeet +#include "ced_tagger.hpp" // pk::CedTagger +#include "sound_stream.hpp" // pk::SoundStream, pk::SoundSegment +#include "audio_io.hpp" // pk::resample_linear #include #include @@ -49,15 +52,17 @@ void usage() { std::fprintf(stderr, "usage:\n" " parakeet-server --model [--host 127.0.0.1] " - "[--port 8080] [--threads N] [--cache-dir ]\n" + "[--port 8080] [--threads N] [--cache-dir ] [--sound-model ]\n" "\n" - "Serves POST /v1/audio/transcriptions (OpenAI-compatible) for one model.\n"); + "Serves POST /v1/audio/transcriptions (OpenAI-compatible) for one model.\n" + "--sound-model loads a ced.cpp sound-event tagger (local GGUF path); " + "verbose_json responses then include a \"sound_events\" array.\n"); } } // namespace int main(int argc, char** argv) { - std::string model_arg, host = "127.0.0.1", cache_dir; + std::string model_arg, host = "127.0.0.1", cache_dir, sound_model; int port = 8080, threads = 0; for (int i = 1; i < argc; ++i) { @@ -71,6 +76,7 @@ int main(int argc, char** argv) { else if (a == "--port") port = std::atoi(next("--port").c_str()); else if (a == "--threads") threads = std::atoi(next("--threads").c_str()); else if (a == "--cache-dir") cache_dir = next("--cache-dir"); + else if (a == "--sound-model") sound_model = next("--sound-model"); else if (a == "-h" || a == "--help") { usage(); return 0; } else if (a == "--version" || a == "-V") { std::printf("parakeet-server %s\n", parakeet_version()); @@ -106,6 +112,24 @@ int main(int argc, char** argv) { pk::shutdown_backend(); return 1; } + + // --sound-model is a local path only; it does not go through the ASR + // model's alias/URL resolver. + std::unique_ptr tagger; + if (!sound_model.empty()) { + if (!pk::CedTagger::available()) { + std::fprintf(stderr, "parakeet-server: built without sound tagging (PARAKEET_WITH_CED=OFF)\n"); + pk::shutdown_backend(); + return 2; + } + tagger = pk::CedTagger::load(sound_model); + if (!tagger) { + std::fprintf(stderr, "parakeet-server: failed to load sound model %s\n", sound_model.c_str()); + pk::shutdown_backend(); + return 1; + } + } + std::mutex infer_mu; httplib::Server svr; @@ -150,11 +174,39 @@ int main(int argc, char** argv) { try { pk::Transcription tr; + std::vector sound_events; + bool have_sounds = false; { std::lock_guard lock(infer_mu); tr = model->transcribe_with_timestamps(pcm, sr, pk::Decoder::kDefault); + + // Sound tagging is verbose_json-only so json/text requests never + // pay for it. Runs under the same lock as ASR: the tagger is not + // thread-safe. A failure here must not fail the transcription. + if (tagger && fmt == Format::kVerboseJson) { + try { + const std::vector* mono16k = &pcm; + std::vector resampled; + if (sr != 16000) { + resampled = pk::resample_linear(pcm, sr, 16000); + mono16k = &resampled; + } + pk::SoundStream stream(tagger->scorer(), tagger->n_classes(), pk::SoundOpts{}); + std::vector segs = + stream.feed(mono16k->data(), (int)mono16k->size(), /*is_last=*/true); + sound_events.reserve(segs.size()); + for (const pk::SoundSegment& s : segs) { + const char* label = tagger->label(s.cls); + sound_events.push_back({label ? label : "", s.start, s.end, s.peak}); + } + have_sounds = true; + } catch (const std::exception& e) { + std::fprintf(stderr, "parakeet-server: sound tagging error: %s\n", e.what()); + } + } } - Response out = format_transcription(tr, fmt, duration_sec, include_words); + Response out = format_transcription(tr, fmt, duration_sec, include_words, + have_sounds ? &sound_events : nullptr); res.set_content(out.body, out.content_type.c_str()); } catch (const std::exception& e) { std::fprintf(stderr, "parakeet-server: inference error: %s\n", e.what()); diff --git a/examples/server/openai_format.cpp b/examples/server/openai_format.cpp index 3c74a18..2e981bd 100644 --- a/examples/server/openai_format.cpp +++ b/examples/server/openai_format.cpp @@ -46,7 +46,8 @@ static std::string fixed(double v, int prec) { } Response format_transcription(const pk::Transcription& tr, Format fmt, - double duration_sec, bool include_words) { + double duration_sec, bool include_words, + const std::vector* sounds) { Response r; if (fmt == Format::kText) { r.body = tr.text; @@ -86,6 +87,18 @@ Response format_transcription(const pk::Transcription& tr, Format fmt, } b += "]"; } + if (sounds) { + b += ",\"sound_events\":["; + for (size_t i = 0; i < sounds->size(); ++i) { + const SoundEventOut& s = (*sounds)[i]; + if (i) b += ","; + b += "{\"label\":\"" + json_escape(s.label) + "\","; + b += "\"start\":" + fixed(s.start, 3) + ","; + b += "\"end\":" + fixed(s.end, 3) + ","; + b += "\"score\":" + fixed(s.score, 4) + "}"; + } + b += "]"; + } b += "}"; r.body = b; r.content_type = "application/json"; diff --git a/examples/server/openai_format.hpp b/examples/server/openai_format.hpp index 10f9ef3..3fce104 100644 --- a/examples/server/openai_format.hpp +++ b/examples/server/openai_format.hpp @@ -1,6 +1,7 @@ #pragma once #include "transcription.hpp" // pk::Transcription, pk::Word #include +#include namespace pkserver { @@ -15,12 +16,21 @@ struct Response { std::string content_type; }; +// One closed sound event, ready for JSON formatting. Mirrors pk::SoundSegment +// but with the class index already resolved to a label string. +struct SoundEventOut { + std::string label; + double start, end, score; +}; + // Build the response body and content type for a finished transcription. // duration_sec is the decoded audio length in seconds. include_words controls // whether the verbose_json "words" array is emitted (OpenAI gates it on -// timestamp_granularities[] containing "word"). +// timestamp_granularities[] containing "word"). When sounds is non-null and +// fmt is kVerboseJson, a "sound_events" array is appended to the body. Response format_transcription(const pk::Transcription& tr, Format fmt, - double duration_sec, bool include_words); + double duration_sec, bool include_words, + const std::vector* sounds = nullptr); // JSON-escape a UTF-8 string (quote, backslash, control chars). Exposed for the // error envelope and tests. diff --git a/include/parakeet_capi.h b/include/parakeet_capi.h index 789ef48..05a27a3 100644 --- a/include/parakeet_capi.h +++ b/include/parakeet_capi.h @@ -55,6 +55,9 @@ typedef struct parakeet_ctx parakeet_ctx; // Additive, same ABI: parakeet_capi_diarize_stream_begin_latency / // _time / _active and parakeet_capi_sas_stream_begin_latency (the model // card's 1.04 / 0.64 / 0.32 s streaming modes). +// 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. int parakeet_capi_abi_version(void); // Load a GGUF model. Returns an owning context, or NULL on failure. @@ -479,6 +482,96 @@ int parakeet_capi_sas_stream_feed(parakeet_sas_stream* s, const float* pcm, void parakeet_capi_sas_stream_free(parakeet_sas_stream* s); +// --- Sound events (ABI v8) -------------------------------------------------- +// A CED GGUF (ced.cpp) loads with parakeet_capi_load into a "tagger" context. + +typedef struct { + int size; // sizeof(parakeet_sound_opts), for versioning + float window_sec, hop_sec; + float on_threshold, off_threshold, min_duration_sec; + int top_k; // per-window scores kept for the drain +} parakeet_sound_opts; +void parakeet_capi_sound_opts_default(parakeet_sound_opts* o); + +typedef struct { + int class_index; + const char* label; // borrowed from the tagger ctx + float start, end, peak; // seconds from stream start +} parakeet_sound_segment; + +typedef struct parakeet_sound_stream parakeet_sound_stream; + +// NULL opts = defaults. NULL on error (last_error on tagger). +parakeet_sound_stream* parakeet_capi_sound_stream_begin(parakeet_ctx* tagger, + const parakeet_sound_opts* o); +// Segments that closed since the previous call; is_last closes all. +int parakeet_capi_sound_stream_feed(parakeet_sound_stream* s, const float* pcm, int n, + int is_last, parakeet_sound_segment** out, int* n_out); +// Still-open segments, end = current stream time. +int parakeet_capi_sound_stream_active(parakeet_sound_stream* s, + parakeet_sound_segment** out, int* n_out); +// [{"start":..,"end":..,"tags":[{"index":..,"label":..,"score":..}]}], windows +// since the previous drain. Free with parakeet_capi_free_string. The stream +// keeps one entry per hop until drained: drain regularly, or set top_k = 0 +// to keep no scores. +char* parakeet_capi_sound_stream_drain_scores_json(parakeet_sound_stream* s); +void parakeet_capi_free_sound_segments(parakeet_sound_segment* segs); +void parakeet_capi_sound_stream_free(parakeet_sound_stream* s); + +// Tagger introspection: -1 / NULL on a context that is not a tagger. +int parakeet_capi_num_classes(const parakeet_ctx* ctx); +const char* parakeet_capi_class_label(const parakeet_ctx* ctx, int index); + +// Which kind of model a context holds, so a caller loading through the same +// parakeet_capi_load can dispatch without probing individual entry points. +#define PARAKEET_MODEL_KIND_NONE 0 +#define PARAKEET_MODEL_KIND_ASR 1 +#define PARAKEET_MODEL_KIND_DIARIZATION 2 +#define PARAKEET_MODEL_KIND_SOUND 3 +int parakeet_capi_model_kind(const parakeet_ctx* ctx); + +// --- Combined scene stream (ABI v8) ----------------------------------------- +// ASR, diarization and sound-event tagging over one live 16 kHz mono PCM +// stream. Any of the three contexts may be NULL; at least one is required. +// Borrows the contexts it is given (same lifetime rule as sas_stream / +// sound_stream): free the scene stream first. + +typedef struct { + int size; // sizeof(parakeet_scene_opts), for versioning + 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 +} parakeet_scene_opts; +void parakeet_capi_scene_opts_default(parakeet_scene_opts* o); + +typedef struct parakeet_scene_stream parakeet_scene_stream; + +// Any context may be NULL; at least one must be given. NULL opts = defaults. +// NULL on error: with an all-NULL call there is no context to report on, so +// nothing is set; otherwise last_error is set on the context of the wrong +// kind, on the diar ctx for an unknown diar_latency ("unknown diarization +// latency mode"), or, on an internal failure, on the part that failed. +parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* asr, parakeet_ctx* diar, + parakeet_ctx* tagger, + const parakeet_scene_opts* o); + +// Everything finalized by this call, as one JSON document (see docs/sound.md +// for the shape: "t", "utterances", "words", "speakers", "sounds", "active"). +// NULL on error. Free with parakeet_capi_free_string. After an error, later +// timestamps may be misaligned (the parts that did not see the failed chunk +// lag behind), so end the stream instead of feeding it more. +char* parakeet_capi_scene_stream_feed_json(parakeet_scene_stream* s, const float* pcm, int n, + int is_last); + +// Same shape as parakeet_capi_sound_stream_drain_scores_json; "[]" without a +// tagger. Free with parakeet_capi_free_string. Scores are kept until drained: +// drain regularly, or set sound.top_k = 0 to keep no scores. +char* parakeet_capi_scene_stream_drain_scores_json(parakeet_scene_stream* s); + +// Last error of this stream, "" if none. Borrowed. +const char* parakeet_capi_scene_stream_last_error(parakeet_scene_stream* s); +void parakeet_capi_scene_stream_free(parakeet_scene_stream* s); + #ifdef __cplusplus } // extern "C" #endif diff --git a/src/asr_committer.cpp b/src/asr_committer.cpp new file mode 100644 index 0000000..71567bd --- /dev/null +++ b/src/asr_committer.cpp @@ -0,0 +1,95 @@ +#include "asr_committer.hpp" + +#include +#include +#include +#include +#include + +namespace pk { + +namespace { + +// Parakeet's word start times come from the frame where the first token is +// emitted, often one or two 80 ms encoder frames after the sound actually +// starts, so cutting exactly at a word start can permanently drop the tail +// end of the previous window's onset. Back off by this much when no word +// commits, so the next window re-hears that onset with full context. +constexpr double kOnsetMargin = 0.3; + +// Lowercase letters and digits only, for comparing a word heard twice. +std::string word_key(const std::string& w) { + std::string k; + for (unsigned char c : w) + if (std::isalnum(c) || c >= 0x80) k += (char)std::tolower(c); + return k; +} + +} // namespace + +AsrCommitter::AsrCommitter(Transcriber t, double min_window_sec, double right_context_sec) + : transcribe_(std::move(t)), min_window_sec_(min_window_sec), right_context_sec_(right_context_sec) {} + +void AsrCommitter::push(const float* pcm, int n) { + if (n > 0) audio_.insert(audio_.end(), pcm, pcm + n); +} + +// Uncommitted samples up to `until` (all of them with is_last). +size_t AsrCommitter::span(double until, bool is_last) const { + return is_last ? audio_.size() + : std::min(audio_.size(), (size_t)std::max(0.0, (until - commit_sec_) * 16000.0)); +} + +// Offline ASR on a short window loses words, so wait until enough +// uncommitted audio has built up (it bounds how often the text commits). +bool AsrCommitter::ready(double until, bool is_last) const { + return is_last || span(until, is_last) >= (size_t)(min_window_sec_ * 16000.0); +} + +std::vector AsrCommitter::commit(double until, bool is_last) { + if (!ready(until, is_last)) return {}; + const size_t span = this->span(until, is_last); + std::vector words; + if (span > 0) words = transcribe_(std::vector(audio_.begin(), audio_.begin() + span)); + // Commit only words that end right_context_sec before the cut: the ASR + // needs right context, and a word at the edge may be cut in half. + size_t keep = words.size(); + double next_commit = commit_sec_ + (double)span / 16000.0; + if (!is_last) { + const double limit = (double)span / 16000.0 - right_context_sec_; + keep = 0; + while (keep < words.size() && words[keep].end <= limit) ++keep; + if (keep > 0) { + // Resume right after the last committed word: audio the ASR + // skipped this time is heard again with more context. + next_commit = commit_sec_ + words[keep - 1].end; + } else { + // No committable word (silence, music, long non-speech). Release + // the audio before the first word heard, or before the right + // context when there is none, so the buffer and the cost of each + // transcription stay bounded. Back off by the onset margin so the + // first uncommitted word's onset is re-heard with full context + // next time, not cut at its (possibly late) reported start. + double first_start = limit; + if (!words.empty()) first_start = std::min(first_start, (double)words.front().start); + const double rel = std::max(0.0, first_start - kOnsetMargin); + next_commit = commit_sec_ + rel; + } + } + std::vector committed(words.begin(), words.begin() + keep); + for (auto& w : committed) { w.start += (float)commit_sec_; w.end += (float)commit_sec_; } + // ASR timestamps are only accurate to a frame or two, so the tail of the + // previously committed word can be heard again at the new start. + if (have_last_ && !committed.empty() && + word_key(committed.front().text) == word_key(last_.text) && + committed.front().start - last_.start < 0.5f) + committed.erase(committed.begin()); + if (!committed.empty()) { last_ = committed.back(); have_last_ = true; } + const size_t drop = std::min(audio_.size(), + (size_t)std::llround((next_commit - commit_sec_) * 16000.0)); + audio_.erase(audio_.begin(), audio_.begin() + drop); + commit_sec_ += (double)drop / 16000.0; + return committed; +} + +} // namespace pk diff --git a/src/asr_committer.hpp b/src/asr_committer.hpp new file mode 100644 index 0000000..77c2e67 --- /dev/null +++ b/src/asr_committer.hpp @@ -0,0 +1,46 @@ +#pragma once +#include "transcription.hpp" // pk::Word + +#include +#include + +namespace pk { + +// Offline transcription of one 16 kHz mono PCM window; word times are +// relative to the window start. +using Transcriber = std::function(const std::vector& pcm)>; + +// Streaming text commit over an offline transcriber: buffers uncommitted PCM, +// transcribes it once enough has built up, and commits only the words that +// have right context. The rest is transcribed again with the next window. +// When a window holds no committable word, the audio before the first word +// heard (or before the right context) is released, so non-speech does not +// build up. +class AsrCommitter { +public: + explicit AsrCommitter(Transcriber t, double min_window_sec = 4.0, double right_context_sec = 1.0); + void push(const float* pcm, int n); + // Words committed up to stream time `until` (absolute times). Empty when + // fewer than min_window_sec of uncommitted audio is available (unless is_last). + std::vector commit(double until, bool is_last); + // True when commit(until, is_last) would transcribe (the window is long + // enough, or is_last); false when it would return early. + bool ready(double until, bool is_last) const; + // Stream time of the first uncommitted sample. + double commit_sec() const { return commit_sec_; } + // Uncommitted samples held (starting at commit_sec()). + size_t buffered_samples() const { return audio_.size(); } + +private: + size_t span(double until, bool is_last) const; + + Transcriber transcribe_; + double min_window_sec_; + double right_context_sec_; + std::vector audio_; // uncommitted PCM, starting at commit_sec_ + double commit_sec_ = 0.0; // stream time of audio_[0] + Word last_; // last committed word (absolute times) + bool have_last_ = false; +}; + +} // namespace pk diff --git a/src/audio_io.cpp b/src/audio_io.cpp index aef6084..9d5ff9b 100644 --- a/src/audio_io.cpp +++ b/src/audio_io.cpp @@ -1,4 +1,3 @@ -#define DR_WAV_IMPLEMENTATION #include "dr_wav.h" #include "audio_io.hpp" #include "common.hpp" diff --git a/src/ced_tagger.cpp b/src/ced_tagger.cpp new file mode 100644 index 0000000..122d1fd --- /dev/null +++ b/src/ced_tagger.cpp @@ -0,0 +1,68 @@ +#include "ced_tagger.hpp" + +#include "gguf.h" + +#ifdef PARAKEET_WITH_CED +#include "ced_capi.h" +#endif + +namespace pk { + +bool gguf_is_ced(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 ced = id >= 0 && gguf_get_kv_type(g, id) == GGUF_TYPE_STRING && + std::string(gguf_get_val_str(g, id)) == "ced"; + gguf_free(g); + return ced; +} + +#ifdef PARAKEET_WITH_CED + +bool CedTagger::available() { return true; } + +std::unique_ptr CedTagger::load(const std::string& path) { + ced_ctx* c = ced_capi_load(path.c_str()); + if (!c) return nullptr; + std::unique_ptr t(new CedTagger()); + t->ctx_ = c; + return t; +} + +CedTagger::~CedTagger() { ced_capi_free(static_cast(ctx_)); } + +int CedTagger::n_classes() const { return ced_capi_num_classes(static_cast(ctx_)); } + +const char* CedTagger::label(int i) const { return ced_capi_label(static_cast(ctx_), i); } + +SoundScorer CedTagger::scorer() { + return [this](const float* pcm, int n, std::vector& probs) { + auto* c = static_cast(ctx_); + last_error_.clear(); // describes the latest call only + probs.assign((size_t)n_classes(), 0.0f); + const int w = ced_capi_classify_pcm_probs(c, pcm, n, 16000, probs.data(), (int)probs.size()); + if (w != (int)probs.size()) { + const char* m = ced_capi_last_error(c); + last_error_ = m ? m : ""; + return false; + } + return true; + }; +} + +#else // PARAKEET_WITH_CED + +bool CedTagger::available() { return false; } +std::unique_ptr CedTagger::load(const std::string&) { return nullptr; } +CedTagger::~CedTagger() = default; +int CedTagger::n_classes() const { return 0; } +const char* CedTagger::label(int) const { return nullptr; } +SoundScorer CedTagger::scorer() { + return [](const float*, int, std::vector&) { return false; }; +} + +#endif + +} // namespace pk diff --git a/src/ced_tagger.hpp b/src/ced_tagger.hpp new file mode 100644 index 0000000..b932916 --- /dev/null +++ b/src/ced_tagger.hpp @@ -0,0 +1,40 @@ +#pragma once +#include +#include +#include +#include + +namespace pk { + +// Scores one PCM window (16 kHz mono): fills `probs` with one score per class +// in class-index order and returns true, or returns false on failure. +using SoundScorer = std::function& probs)>; + +// A loaded ced.cpp model (CED AudioSet tagger). The only parakeet code that +// talks to ced.cpp, and only through ced_capi.h. Not thread-safe: one stream +// at a time per tagger, like the other contexts. +class CedTagger { +public: + // False when parakeet was built with PARAKEET_WITH_CED=OFF. + static bool available(); + // nullptr on failure (or when unavailable). + static std::unique_ptr load(const std::string& gguf_path); + ~CedTagger(); + CedTagger(const CedTagger&) = delete; + CedTagger& operator=(const CedTagger&) = delete; + + int n_classes() const; + const char* label(int index) const; + SoundScorer scorer(); + const std::string& last_error() const { return last_error_; } + +private: + CedTagger() = default; + void* ctx_ = nullptr; // ced_ctx* + std::string last_error_; +}; + +// True when the GGUF's general.architecture is "ced". Reads only the header. +bool gguf_is_ced(const std::string& gguf_path); + +} // namespace pk diff --git a/src/diar_pcm_stream.cpp b/src/diar_pcm_stream.cpp new file mode 100644 index 0000000..9655f69 --- /dev/null +++ b/src/diar_pcm_stream.cpp @@ -0,0 +1,62 @@ +#include "diar_pcm_stream.hpp" + +#include "mel.hpp" // pk::StreamingMel + +#include + +namespace pk { + +DiarPcmStream::DiarPcmStream(const DiarizationModel& m, DiarLatency latency) + : hop_length_((int)m.config().hop_length) { + const ModelLoader& ml = m.loader(); + sd_ = std::make_unique(ml, diar_stream_config(latency, ml.config())); + mel_ = std::make_unique(ml); +} + +DiarPcmStream::~DiarPcmStream() = default; + +long long DiarPcmStream::feed(const float* pcm, int n, bool is_last, + std::vector& closed) { + const int n_mels = sd_->n_mels(); + const long long done_before = sd_->frames_done(); + std::vector mel; + int nf = 0; + if (n > 0) { + mel = mel_->feed(pcm, n, nf); + samples_in_ += n; + } + if (is_last) { + int nt = 0; + std::vector tail = mel_->finalize(nt); + // Join the two feat-major blocks, then keep floor(S / hop) frames in + // total like NeMo (the centered STFT emits one more). + const long long valid = samples_in_ / (long long)hop_length_; + const int keep = (int)std::max(0LL, std::min(nf + nt, valid - sd_->frames_in())); + std::vector joined((size_t)n_mels * keep); + for (int m = 0; m < n_mels; ++m) + for (int t = 0; t < keep; ++t) + joined[(size_t)m * keep + t] = t < nf ? mel[(size_t)m * nf + t] + : tail[(size_t)m * nt + (t - nf)]; + mel.swap(joined); + nf = keep; + finished_ = true; + } + auto segs = sd_->feed_mel(mel, n_mels, nf, is_last); + closed.insert(closed.end(), segs.begin(), segs.end()); + return sd_->frames_done() - done_before; +} + +double DiarPcmStream::diarized_until() const { + const double hop_sec = (double)hop_length_ / 16000.0; + return sd_->frames_done() * hop_sec; +} + +std::vector DiarPcmStream::open_segments() const { + return sd_->open_segments(); +} + +int DiarPcmStream::chunk_samples() const { + return sd_->latency_mel_frames() * hop_length_; +} + +} // namespace pk diff --git a/src/diar_pcm_stream.hpp b/src/diar_pcm_stream.hpp new file mode 100644 index 0000000..ae6ee73 --- /dev/null +++ b/src/diar_pcm_stream.hpp @@ -0,0 +1,38 @@ +#pragma once +#include "diarization.hpp" // pk::DiarizationModel +#include "diarization_streaming.hpp" // pk::StreamingDiarization, pk::DiarLatency + +#include +#include + +namespace pk { + +class StreamingMel; + +// PCM -> incremental mel -> StreamingDiarization. The diarizer runs every +// chunk whose look-ahead has arrived (and, with is_last, the rest). +class DiarPcmStream { +public: + DiarPcmStream(const DiarizationModel& m, DiarLatency latency); + ~DiarPcmStream(); + DiarPcmStream(const DiarPcmStream&) = delete; + DiarPcmStream& operator=(const DiarPcmStream&) = delete; + + // Returns mel frames diarized by this call; appends closed segments. + long long feed(const float* pcm, int n, bool is_last, std::vector& closed); + double diarized_until() const; // seconds + std::vector open_segments() const; + bool finished() const { return finished_; } + StreamingDiarization& sd() { return *sd_; } + const StreamingDiarization& sd() const { return *sd_; } + int chunk_samples() const; + +private: + int hop_length_; + std::unique_ptr sd_; + std::unique_ptr mel_; + long long samples_in_ = 0; // PCM samples fed so far + bool finished_ = false; +}; + +} // namespace pk diff --git a/src/dr_wav_impl.cpp b/src/dr_wav_impl.cpp new file mode 100644 index 0000000..1c99e3a --- /dev/null +++ b/src/dr_wav_impl.cpp @@ -0,0 +1,12 @@ +// The single translation unit that provides the dr_wav implementation for +// the whole parakeet build. Both libparakeet (src/audio_io.cpp) and, when +// PARAKEET_WITH_CED is on, ced.cpp (built with CED_EXTERNAL_DR_WAV) declare +// dr_wav's functions via #include "dr_wav.h" without defining +// DR_WAV_IMPLEMENTATION themselves, and link against this object instead. +// That keeps there being exactly one copy of dr_wav's symbols in the final +// binary (a second DR_WAV_IMPLEMENTATION define anywhere else would be a +// multiple-definition link error) and, unlike having libparakeet itself own +// the implementation, does not create a link-time dependency cycle between +// the parakeet and ced static libraries. +#define DR_WAV_IMPLEMENTATION +#include "dr_wav.h" diff --git a/src/parakeet_capi.cpp b/src/parakeet_capi.cpp index 510a701..8998f4b 100644 --- a/src/parakeet_capi.cpp +++ b/src/parakeet_capi.cpp @@ -4,8 +4,12 @@ #include "diarization.hpp" // pk::DiarizationModel #include "diarization_streaming.hpp" // pk::StreamingDiarization #include "streaming.hpp" // pk::StreamingSession +#include "ced_tagger.hpp" // pk::CedTagger +#include "sound_stream.hpp" // pk::SoundStream #include "mel.hpp" // pk::MelFrontend #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 "transcription.hpp" // pk::Transcription, pk::Word #include "transcription_json.hpp" @@ -13,6 +17,7 @@ #include #include #include +#include #include #include #include @@ -44,14 +49,17 @@ // v7: speaker diarization (diarize_*), speaker-attributed ASR // (transcribe_and_diarize*, sas_stream_*) and streaming diarization // (diarize_stream_*). A context holds either an ASR or a diarization model. -#define PARAKEET_CAPI_ABI_VERSION 7 +// v8: sound-event detection (CED), sound_stream_*, scene_stream_*; additive. +#define PARAKEET_CAPI_ABI_VERSION 8 // The opaque context: a loaded model plus a buffer for the last error message. -// Exactly one of `model` / `diar` is non-null: ASR models use `model`, -// diarization models (Sortformer) use `diar`. +// Exactly one of `model` / `diar` / `tagger` is non-null: ASR models use +// `model`, diarization models (Sortformer) use `diar`, CED sound-event +// taggers use `tagger`. struct parakeet_ctx { std::unique_ptr model; std::unique_ptr diar; + std::unique_ptr tagger; std::string last_error; }; @@ -187,6 +195,15 @@ extern "C" parakeet_ctx* parakeet_capi_load(const char* gguf_path) { auto* ctx = new (std::nothrow) parakeet_ctx(); if (!ctx) return nullptr; + // A CED GGUF (general.architecture "ced") is a sound tagger. Check the + // header first: the ASR loader would misread it. + if (pk::gguf_is_ced(gguf_path)) { + ctx->tagger = pk::CedTagger::load(gguf_path); + if (ctx->tagger) 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. @@ -921,9 +938,9 @@ namespace { bool require_diar(parakeet_ctx* ctx) { if (!ctx) return false; if (!ctx->diar) { - ctx->last_error = ctx->model - ? "context holds an ASR model; diarize_* needs a diarization model" - : "context has no loaded model"; + 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" + : "context has no loaded model"; return false; } return true; @@ -932,9 +949,23 @@ bool require_diar(parakeet_ctx* ctx) { bool require_asr(parakeet_ctx* ctx) { if (!ctx) return false; if (!ctx->model) { - ctx->last_error = ctx->diar - ? "context holds a diarization model; an ASR model is needed here" - : "context has no loaded model"; + 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" + : "context has no loaded model"; + return false; + } + return true; +} + +constexpr const char* kNoCed = "built without sound tagging (PARAKEET_WITH_CED=OFF)"; + +bool require_tagger(parakeet_ctx* ctx) { + if (!ctx) return false; + if (!pk::CedTagger::available()) { ctx->last_error = kNoCed; return false; } + 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" + : "context has no loaded model"; return false; } return true; @@ -1122,48 +1153,11 @@ extern "C" char* parakeet_capi_transcribe_and_diarize_json(parakeet_ctx* asr_ctx struct parakeet_diar_stream { parakeet_ctx* ctx = nullptr; - std::unique_ptr sd; - std::unique_ptr mel; - long long samples_in = 0; // PCM samples fed so far - bool finished = false; + std::unique_ptr ds; }; namespace { -// Feed PCM to the stream's mel front end and the diarizer, which runs every -// chunk whose look-ahead has arrived (and, with is_last, the rest). Closed -// segments are appended to `segs`. Returns the mel frames diarized by this call. -long long diar_stream_advance(parakeet_diar_stream* s, const float* pcm, int n, bool is_last, - std::vector& segs) { - const int n_mels = s->sd->n_mels(); - const long long done_before = s->sd->frames_done(); - std::vector mel; - int nf = 0; - if (n > 0) { - mel = s->mel->feed(pcm, n, nf); - s->samples_in += n; - } - if (is_last) { - int nt = 0; - std::vector tail = s->mel->finalize(nt); - // Join the two feat-major blocks, then keep floor(S / hop) frames in - // total like NeMo (the centered STFT emits one more). - const long long valid = s->samples_in / (long long)s->ctx->diar->config().hop_length; - const int keep = (int)std::max(0LL, std::min(nf + nt, valid - s->sd->frames_in())); - std::vector joined((size_t)n_mels * keep); - for (int m = 0; m < n_mels; ++m) - for (int t = 0; t < keep; ++t) - joined[(size_t)m * keep + t] = t < nf ? mel[(size_t)m * nf + t] - : tail[(size_t)m * nt + (t - nf)]; - mel.swap(joined); - nf = keep; - s->finished = true; - } - auto closed = s->sd->feed_mel(mel, n_mels, nf, is_last); - segs.insert(segs.end(), closed.begin(), closed.end()); - return s->sd->frames_done() - done_before; -} - pk::DiarLatency latency_from_int(int latency) { switch (latency) { case PARAKEET_DIAR_LATENCY_LOW: return pk::DiarLatency::Low; @@ -1197,12 +1191,10 @@ extern "C" parakeet_diar_stream* parakeet_capi_diarize_stream_begin_latency(para return nullptr; } try { - const pk::ModelLoader& ml = diar_ctx->diar->loader(); + auto ds = std::make_unique(*diar_ctx->diar, latency_from_int(latency)); auto* s = new parakeet_diar_stream(); s->ctx = diar_ctx; - s->sd = std::make_unique( - ml, pk::diar_stream_config(latency_from_int(latency), ml.config())); - s->mel = std::make_unique(ml); + s->ds = std::move(ds); diar_ctx->last_error.clear(); return s; } catch (const std::exception& e) { @@ -1219,18 +1211,18 @@ extern "C" parakeet_diar_stream* parakeet_capi_diarize_stream_begin(parakeet_ctx extern "C" int parakeet_capi_diarize_stream_chunk_samples(parakeet_diar_stream* s) { if (!s) return 0; - return s->sd->latency_mel_frames() * (int)s->ctx->diar->config().hop_length; + return s->ds->chunk_samples(); } extern "C" float parakeet_capi_diarize_stream_time(parakeet_diar_stream* s) { if (!s) return 0.0f; - return (float)(s->sd->frames_done() * s->sd->frame_sec()); + return (float)(s->ds->sd().frames_done() * s->ds->sd().frame_sec()); } extern "C" int parakeet_capi_diarize_stream_active(parakeet_diar_stream* s, parakeet_diar_segment** out, int* n_out) { if (!s || !out || !n_out) return 1; - if (!to_c_segments(s->sd->open_segments(), out, n_out)) { + if (!to_c_segments(s->ds->open_segments(), out, n_out)) { s->ctx->last_error = "out of memory"; return 1; } @@ -1244,10 +1236,10 @@ extern "C" int parakeet_capi_diarize_stream_feed(parakeet_diar_stream* s, const *out = nullptr; *n_out = 0; if ((!pcm && n_samples > 0) || n_samples < 0) { s->ctx->last_error = "invalid samples buffer"; return 1; } - if (s->finished) { s->ctx->last_error = "stream already finished"; return 1; } + if (s->ds->finished()) { s->ctx->last_error = "stream already finished"; return 1; } try { std::vector segs; - diar_stream_advance(s, pcm, n_samples, is_last != 0, segs); + s->ds->feed(pcm, n_samples, is_last != 0, segs); if (!to_c_segments(segs, out, n_out)) { s->ctx->last_error = "out of memory"; return 1; } s->ctx->last_error.clear(); return 0; @@ -1269,39 +1261,39 @@ extern "C" void parakeet_capi_diarize_stream_free(parakeet_diar_stream* s) { // --- Streaming speaker-attributed ASR --------------------------------------- +// A pk::SceneStream over an ASR and a diarization context (no sound part). struct parakeet_sas_stream { parakeet_ctx* asr = nullptr; - parakeet_diar_stream* diar = nullptr; - std::vector audio; // uncommitted PCM, starting at commit_sec - double commit_sec = 0.0; // stream time of audio[0] - std::vector segs; // closed diarization segments - pk::Word last_word; // last committed word (absolute times) - bool have_last_word = false; + parakeet_ctx* diar = nullptr; + std::unique_ptr scene; }; -namespace { - -// Lowercase letters and digits only, for comparing a word heard twice. -std::string word_key(const std::string& w) { - std::string k; - for (unsigned char c : w) - if (std::isalnum(c) || c >= 0x80) k += (char)std::tolower(c); - return k; -} - -} // namespace - extern "C" parakeet_sas_stream* parakeet_capi_sas_stream_begin_latency(parakeet_ctx* asr_ctx, parakeet_ctx* diar_ctx, int latency) { if (!require_asr(asr_ctx) || !require_diar(diar_ctx)) return nullptr; - parakeet_diar_stream* d = parakeet_capi_diarize_stream_begin_latency(diar_ctx, latency); - if (!d) return nullptr; - auto* s = new (std::nothrow) parakeet_sas_stream(); - if (!s) { parakeet_capi_diarize_stream_free(d); return nullptr; } - s->asr = asr_ctx; - s->diar = d; - return s; + if (latency < PARAKEET_DIAR_LATENCY_MODEL || latency > PARAKEET_DIAR_LATENCY_ULTRA_LOW) { + diar_ctx->last_error = "unknown diarization latency mode"; + return nullptr; + } + try { + pk::SceneParts parts; + parts.asr = asr_ctx->model.get(); + parts.diar = diar_ctx->diar.get(); + parts.diar_latency = latency_from_int(latency); + auto scene = std::make_unique(parts); + auto* s = new parakeet_sas_stream(); + s->asr = asr_ctx; + s->diar = diar_ctx; + s->scene = std::move(scene); + diar_ctx->last_error.clear(); + return s; + } catch (const std::exception& e) { + diar_ctx->last_error = e.what(); + } catch (...) { + diar_ctx->last_error = "unknown error"; + } + return nullptr; } extern "C" parakeet_sas_stream* parakeet_capi_sas_stream_begin(parakeet_ctx* asr_ctx, @@ -1316,88 +1308,307 @@ extern "C" int parakeet_capi_sas_stream_feed(parakeet_sas_stream* s, const float *out = nullptr; *n_out = 0; if ((!pcm && n_samples > 0) || n_samples < 0) { s->asr->last_error = "invalid samples buffer"; return 1; } - if (s->diar->finished) { s->asr->last_error = "stream already finished"; return 1; } - parakeet_ctx* failed = s->diar->ctx; + if (s->scene->finished()) { s->asr->last_error = "stream already finished"; return 1; } try { - std::vector closed; - const long long advanced = diar_stream_advance(s->diar, pcm, n_samples, is_last != 0, closed); - for (const auto& c : closed) s->segs.push_back({c.speaker, c.start, c.end}); - if (n_samples > 0) s->audio.insert(s->audio.end(), pcm, pcm + n_samples); - if (advanced == 0 && !is_last) return 0; - - // Offline ASR on a short window loses words, so wait until enough - // uncommitted, diarized audio has built up (it bounds how often the - // text commits, not the speaker latency). - constexpr double kSasMinWindowSec = 4.0; - - // Diarized audio ends at frames_done; transcribe the uncommitted span. - const double hop_sec = (double)s->diar->ctx->diar->config().hop_length / 16000.0; - const double diar_end = s->diar->sd->frames_done() * hop_sec; - size_t span = is_last ? s->audio.size() - : std::min(s->audio.size(), - (size_t)std::max(0.0, (diar_end - s->commit_sec) * 16000.0)); - if (!is_last && span < (size_t)(kSasMinWindowSec * 16000.0)) return 0; - failed = s->asr; - std::vector words; - if (span > 0) { - const std::vector seg(s->audio.begin(), s->audio.begin() + span); - words = s->asr->model->transcribe_with_timestamps(seg, 16000).words; - } - // Commit only words that end kSasRightContextSec before the cut: the - // ASR needs right context, and a word at the edge may be cut in half. - // The rest is transcribed again with the next chunk. - constexpr double kSasRightContextSec = 1.0; - size_t keep = words.size(); - double next_commit = s->commit_sec + (double)span / 16000.0; - if (!is_last) { - const double limit = (double)span / 16000.0 - kSasRightContextSec; - keep = 0; - while (keep < words.size() && words[keep].end <= limit) ++keep; - // Resume right after the last committed word: audio the ASR - // skipped this time is heard again with more context. - next_commit = s->commit_sec + (keep > 0 ? words[keep - 1].end : 0.0); - } - std::vector committed(words.begin(), words.begin() + keep); - for (auto& w : committed) { w.start += (float)s->commit_sec; w.end += (float)s->commit_sec; } - // ASR timestamps are only accurate to a frame or two, so the tail of the - // previously committed word can be heard again at the new start. - if (s->have_last_word && !committed.empty() && - word_key(committed.front().text) == word_key(s->last_word.text) && - committed.front().start - s->last_word.start < 0.5f) { - committed.erase(committed.begin()); - } - if (!committed.empty()) { - s->last_word = committed.back(); - s->have_last_word = true; - } - - // Speaker segments known so far: closed ones plus those still open. - std::vector segs = s->segs; - for (const auto& o : s->diar->sd->open_segments()) segs.push_back({o.speaker, o.start, o.end}); - const auto utts = pk::group_speaker_words(pk::merge_asr_diarization(committed, segs)); - - const size_t drop = std::min(s->audio.size(), - (size_t)std::llround((next_commit - s->commit_sec) * 16000.0)); - s->audio.erase(s->audio.begin(), s->audio.begin() + drop); - s->commit_sec += (double)drop / 16000.0; - // Segments that ended before the commit point can no longer match a word. - s->segs.erase(std::remove_if(s->segs.begin(), s->segs.end(), - [&](const pk::SpeakerSegment& g) { return g.end < s->commit_sec; }), - s->segs.end()); - - if (!to_c_results(utts, out, n_out)) { s->asr->last_error = "out of memory"; return 1; } + const pk::SceneUpdate u = s->scene->feed(pcm, n_samples, is_last != 0); + if (!to_c_results(u.utterances, out, n_out)) { s->asr->last_error = "out of memory"; return 1; } + // Any successful feed clears the ASR ctx's last error, also one + // that commits nothing. s->asr->last_error.clear(); return 0; } catch (const std::exception& e) { - failed->last_error = e.what(); + (s->scene->failed_part() == pk::ScenePart::Asr ? s->asr : s->diar)->last_error = e.what(); } catch (...) { - failed->last_error = "unknown error"; + (s->scene->failed_part() == pk::ScenePart::Asr ? s->asr : s->diar)->last_error = "unknown error"; } return 1; } extern "C" void parakeet_capi_sas_stream_free(parakeet_sas_stream* s) { - if (!s) return; - parakeet_capi_diarize_stream_free(s->diar); delete s; } + +// --------------------------------------------------------------------------- +// Sound events (ABI v8) +// --------------------------------------------------------------------------- + +namespace { + +pk::SoundOpts to_sound_opts(const parakeet_sound_opts* o) { + pk::SoundOpts s; + if (!o) return s; + // Only read fields the caller's struct has (size versioning). + auto has = [&](size_t end) { return o->size >= (int)end; }; + if (has(offsetof(parakeet_sound_opts, hop_sec) + sizeof(float))) { + s.window_sec = o->window_sec; + s.hop_sec = o->hop_sec; + } + if (has(offsetof(parakeet_sound_opts, min_duration_sec) + sizeof(float))) { + s.on_threshold = o->on_threshold; + s.off_threshold = o->off_threshold; + s.min_duration_sec = o->min_duration_sec; + } + if (has(offsetof(parakeet_sound_opts, top_k) + sizeof(int))) s.top_k = o->top_k; + return s; +} + +bool to_c_sound_segments(const std::vector& in, const pk::CedTagger& t, + parakeet_sound_segment** out, int* n_out) { + *out = nullptr; + *n_out = 0; + if (in.empty()) return true; + auto* a = static_cast(std::malloc(in.size() * sizeof(parakeet_sound_segment))); + if (!a) return false; + for (size_t i = 0; i < in.size(); ++i) { + const char* l = t.label(in[i].cls); + a[i] = {in[i].cls, l ? l : "", in[i].start, in[i].end, in[i].peak}; + } + *out = a; + *n_out = (int)in.size(); + return true; +} + +} // namespace + +struct parakeet_sound_stream { + parakeet_ctx* ctx = nullptr; + std::unique_ptr ss; +}; + +extern "C" void parakeet_capi_sound_opts_default(parakeet_sound_opts* o) { + if (!o) return; + const pk::SoundOpts d; + *o = {(int)sizeof(*o), d.window_sec, d.hop_sec, d.on_threshold, d.off_threshold, + d.min_duration_sec, d.top_k}; +} + +extern "C" parakeet_sound_stream* parakeet_capi_sound_stream_begin(parakeet_ctx* tagger, + const parakeet_sound_opts* o) { + if (!require_tagger(tagger)) return nullptr; + try { + const pk::SoundOpts so = to_sound_opts(o); + const std::string err = pk::validate_sound_opts(so, tagger->tagger->n_classes()); + if (!err.empty()) { tagger->last_error = "invalid sound options: " + err; return nullptr; } + auto* s = new parakeet_sound_stream(); + s->ctx = tagger; + s->ss = std::make_unique(tagger->tagger->scorer(), + tagger->tagger->n_classes(), so); + tagger->last_error.clear(); + return s; + } catch (const std::exception& e) { + tagger->last_error = e.what(); + } catch (...) { + tagger->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" int parakeet_capi_sound_stream_feed(parakeet_sound_stream* s, const float* pcm, int n, + int is_last, parakeet_sound_segment** out, int* n_out) { + if (!s || !out || !n_out) return 1; + *out = nullptr; + *n_out = 0; + if ((!pcm && n > 0) || n < 0) { s->ctx->last_error = "invalid samples buffer"; return 1; } + if (s->ss->finished()) { s->ctx->last_error = "stream already finished"; return 1; } + try { + auto closed = s->ss->feed(pcm, n, is_last != 0); + if (!to_c_sound_segments(closed, *s->ctx->tagger, out, n_out)) { + s->ctx->last_error = "out of memory"; + return 1; + } + s->ctx->last_error.clear(); + return 0; + } catch (const std::exception& e) { + s->ctx->last_error = e.what(); + const std::string& detail = s->ctx->tagger->last_error(); + if (!detail.empty()) s->ctx->last_error += ": " + detail; + } catch (...) { + s->ctx->last_error = "unknown error"; + } + return 1; +} + +extern "C" int parakeet_capi_sound_stream_active(parakeet_sound_stream* s, + parakeet_sound_segment** out, int* n_out) { + if (!s || !out || !n_out) return 1; + *out = nullptr; + *n_out = 0; + try { + if (!to_c_sound_segments(s->ss->open_segments(), *s->ctx->tagger, out, n_out)) { + s->ctx->last_error = "out of memory"; + return 1; + } + s->ctx->last_error.clear(); + return 0; + } catch (const std::exception& e) { + s->ctx->last_error = e.what(); + } catch (...) { + s->ctx->last_error = "unknown error"; + } + return 1; +} + +extern "C" char* parakeet_capi_sound_stream_drain_scores_json(parakeet_sound_stream* s) { + if (!s) return nullptr; + try { + const pk::CedTagger& t = *s->ctx->tagger; + char* out = dup_to_c(pk::sound_windows_to_json(s->ss->drain_windows(), + [&](int i) { return t.label(i); })); + s->ctx->last_error.clear(); + return out; + } catch (const std::exception& e) { + s->ctx->last_error = e.what(); + } catch (...) { + s->ctx->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" void parakeet_capi_free_sound_segments(parakeet_sound_segment* segs) { std::free(segs); } + +extern "C" void parakeet_capi_sound_stream_free(parakeet_sound_stream* s) { delete s; } + +// --------------------------------------------------------------------------- +// Sound events (ABI v8): CED tagger introspection +// --------------------------------------------------------------------------- + +extern "C" int parakeet_capi_num_classes(const parakeet_ctx* ctx) { + return (ctx && ctx->tagger) ? ctx->tagger->n_classes() : -1; +} + +extern "C" const char* parakeet_capi_class_label(const parakeet_ctx* ctx, int index) { + return (ctx && ctx->tagger) ? ctx->tagger->label(index) : nullptr; +} + +extern "C" int parakeet_capi_model_kind(const parakeet_ctx* ctx) { + if (!ctx) return PARAKEET_MODEL_KIND_NONE; + if (ctx->model) return PARAKEET_MODEL_KIND_ASR; + if (ctx->diar) return PARAKEET_MODEL_KIND_DIARIZATION; + if (ctx->tagger) return PARAKEET_MODEL_KIND_SOUND; + return PARAKEET_MODEL_KIND_NONE; +} + +// --------------------------------------------------------------------------- +// Combined scene stream (ABI v8) +// --------------------------------------------------------------------------- + +// A pk::SceneStream over up to three contexts (ASR, diarization, tagger). +struct parakeet_scene_stream { + parakeet_ctx* asr_ctx = nullptr; + parakeet_ctx* diar_ctx = nullptr; + parakeet_ctx* tagger_ctx = nullptr; + std::unique_ptr scene; + std::string last_error; +}; + +namespace { + +// Which context was running when the scene stream last threw, mirroring the +// sas_stream wrapper's diar/asr attribution, extended with the tagger. +parakeet_ctx* scene_failed_ctx(parakeet_scene_stream* s) { + switch (s->scene->failed_part()) { + 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; + default: return s->asr_ctx ? s->asr_ctx + : s->diar_ctx ? s->diar_ctx : s->tagger_ctx; + } +} + +} // namespace + +extern "C" void parakeet_capi_scene_opts_default(parakeet_scene_opts* o) { + if (!o) return; + o->size = (int)sizeof(*o); + o->diar_latency = PARAKEET_DIAR_LATENCY_MODEL; + parakeet_capi_sound_opts_default(&o->sound); + o->flags = 0; +} + +extern "C" parakeet_scene_stream* parakeet_capi_scene_stream_begin(parakeet_ctx* asr, parakeet_ctx* diar, + parakeet_ctx* tagger, + const parakeet_scene_opts* o) { + if (!asr && !diar && !tagger) return nullptr; + if ((asr && !require_asr(asr)) || (diar && !require_diar(diar)) || (tagger && !require_tagger(tagger))) + return nullptr; + parakeet_scene_opts def; + parakeet_capi_scene_opts_default(&def); + if (!o) o = &def; + if (o->size >= (int)(offsetof(parakeet_scene_opts, flags) + sizeof(int)) && o->flags != 0) { + (asr ? asr : diar ? diar : tagger)->last_error = "scene flags must be 0"; + return nullptr; + } + if (diar && (o->diar_latency < PARAKEET_DIAR_LATENCY_MODEL || + o->diar_latency > PARAKEET_DIAR_LATENCY_ULTRA_LOW)) { + diar->last_error = "unknown diarization latency mode"; + return nullptr; + } + try { + pk::SceneParts p; + p.asr = asr ? asr->model.get() : nullptr; + p.diar = diar ? diar->diar.get() : nullptr; + p.diar_latency = latency_from_int(o->diar_latency); + p.tagger = tagger ? tagger->tagger.get() : nullptr; + p.sound = to_sound_opts(&o->sound); + if (tagger) { + 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* s = new parakeet_scene_stream(); + s->asr_ctx = asr; + s->diar_ctx = diar; + s->tagger_ctx = tagger; + s->scene = std::make_unique(p); + if (asr) asr->last_error.clear(); + if (diar) diar->last_error.clear(); + if (tagger) tagger->last_error.clear(); + return s; + } catch (const std::exception& e) { + (asr ? asr : diar ? diar : tagger)->last_error = e.what(); + } catch (...) { + (asr ? asr : diar ? diar : tagger)->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" char* parakeet_capi_scene_stream_feed_json(parakeet_scene_stream* s, const float* pcm, + int n, int is_last) { + if (!s) return nullptr; + if ((!pcm && n > 0) || n < 0) { s->last_error = "invalid samples buffer"; return nullptr; } + if (s->scene->finished()) { s->last_error = "stream already finished"; return nullptr; } + try { + const pk::SceneUpdate u = s->scene->feed(pcm, n, is_last != 0); + const pk::CedTagger* t = s->tagger_ctx ? s->tagger_ctx->tagger.get() : nullptr; + s->last_error.clear(); + return dup_to_c(pk::scene_update_to_json(u, [t](int i) { return t ? t->label(i) : nullptr; })); + } catch (const std::exception& e) { + s->last_error = e.what(); + if (parakeet_ctx* c = scene_failed_ctx(s)) c->last_error = e.what(); + } catch (...) { + s->last_error = "unknown error"; + if (parakeet_ctx* c = scene_failed_ctx(s)) c->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" char* parakeet_capi_scene_stream_drain_scores_json(parakeet_scene_stream* s) { + if (!s) return nullptr; + try { + const pk::CedTagger* t = s->tagger_ctx ? s->tagger_ctx->tagger.get() : nullptr; + char* out = dup_to_c(pk::sound_windows_to_json(s->scene->drain_windows(), + [t](int i) { return t ? t->label(i) : nullptr; })); + s->last_error.clear(); + return out; + } catch (const std::exception& e) { + s->last_error = e.what(); + } catch (...) { + s->last_error = "unknown error"; + } + return nullptr; +} + +extern "C" const char* parakeet_capi_scene_stream_last_error(parakeet_scene_stream* s) { + return s ? s->last_error.c_str() : ""; +} + +extern "C" void parakeet_capi_scene_stream_free(parakeet_scene_stream* s) { delete s; } diff --git a/src/scene_render.cpp b/src/scene_render.cpp new file mode 100644 index 0000000..c3a1830 --- /dev/null +++ b/src/scene_render.cpp @@ -0,0 +1,126 @@ +#include "scene_render.hpp" + +#include +#include +#include + +namespace pk { + +namespace { + +// CED labels the renderer treats as "just speech": already carried by the +// ASR/diarization transcript, so redundant on screen unless asked for. +bool speech_label_set(const std::string& label) { + static const char* kSpeech[] = { + "Speech", + "Male speech, man speaking", + "Female speech, woman speaking", + "Child speech, kid speaking", + "Conversation", + "Narration, monologue", + "Speech synthesizer", + }; + for (const char* s : kSpeech) + if (label == s) return true; + return false; +} + +} // namespace + +bool is_speech_label(const std::string& label) { return speech_label_set(label); } + +std::string format_span(double start, double end) { + auto mmss = [](double x) { + // Tenths truncated, not rounded. Callers pass timestamps that + // started life as float32 (SoundSegment/SpeakerUtterance), so the + // double here can sit a few ULPs under the intended value (10.9f + // widens to ~10.899999...). Round to hundredths first to absorb + // that noise, keeping genuine truncation (10.04 -> 10.0) intact. + if (!std::isfinite(x)) x = 0.0; // NaN/inf: never let the cast below hit UB. + double hundredths = std::round(x * 100.0) / 100.0; + long long total_tenths = (long long)std::floor(hundredths * 10.0 + 1e-9); + if (total_tenths < 0) total_tenths = 0; + long long minutes = total_tenths / 600; + long long rem = total_tenths % 600; + long long seconds = rem / 10; + long long tenth = rem % 10; + char buf[32]; + std::snprintf(buf, sizeof(buf), "%02lld:%02lld.%lld", minutes, seconds, tenth); + return std::string(buf); + }; + return "[" + mmss(start) + " - " + mmss(end) + "]"; +} + +SceneRenderer::SceneRenderer(bool has_diar, bool show_speech, std::function label, + bool has_asr) + : has_diar_(has_diar), + show_speech_(show_speech), + speaker_lines_(has_diar && !has_asr), + label_(std::move(label)) {} + +void SceneRenderer::add(const SceneUpdate& u) { + if (speaker_lines_) { + for (const SpeakerSegment& g : u.speakers) { + pending_.push_back({(double)g.start, + format_span(g.start, g.end) + " Speaker " + std::to_string(g.speaker)}); + diarized_ = std::max(diarized_, (double)g.end); + } + // A closed segment ends at or before the diarized time, and an open + // one reports it as its end. A segment that opens later starts at or + // after it; one that is open now keeps its start. + for (const StreamingSpeakerSegment& o : u.active_speakers) + diarized_ = std::max(diarized_, (double)o.end); + speaker_bound_ = diarized_; + for (const StreamingSpeakerSegment& o : u.active_speakers) + speaker_bound_ = std::min(speaker_bound_, (double)o.start); + } + for (const SpeakerUtterance& utt : u.utterances) { + std::string line = format_span(utt.start, utt.end) + " "; + if (has_diar_) { + if (utt.speaker >= 0) + line += "Speaker " + std::to_string(utt.speaker) + ": "; + else + line += "Speaker ?: "; + } + line += utt.text; + pending_.push_back({(double)utt.start, std::move(line)}); + } + for (const SoundSegment& s : u.sounds) { + const char* raw = label_ ? label_(s.cls) : nullptr; + std::string label_str = raw ? raw : ""; + if (!show_speech_ && is_speech_label(label_str)) continue; + char peak[16]; + std::snprintf(peak, sizeof(peak), "%.2f", (double)s.peak); + std::string line = format_span(s.start, s.end) + " (" + label_str + " " + peak + ")"; + pending_.push_back({(double)s.start, std::move(line)}); + } +} + +std::vector SceneRenderer::flush(double safe_until) { + if (speaker_lines_) safe_until = std::min(safe_until, speaker_bound_); + std::stable_sort(pending_.begin(), pending_.end(), + [](const Item& a, const Item& b) { return a.start < b.start; }); + std::vector out; + std::vector remain; + remain.reserve(pending_.size()); + for (Item& it : pending_) { + if (it.start < safe_until) + out.push_back(std::move(it.line)); + else + remain.push_back(std::move(it)); + } + pending_ = std::move(remain); + return out; +} + +std::vector SceneRenderer::flush_all() { + std::stable_sort(pending_.begin(), pending_.end(), + [](const Item& a, const Item& b) { return a.start < b.start; }); + std::vector out; + out.reserve(pending_.size()); + for (Item& it : pending_) out.push_back(std::move(it.line)); + pending_.clear(); + return out; +} + +} // namespace pk diff --git a/src/scene_render.hpp b/src/scene_render.hpp new file mode 100644 index 0000000..114f9e7 --- /dev/null +++ b/src/scene_render.hpp @@ -0,0 +1,58 @@ +#pragma once +#include "scene_stream.hpp" // pk::SceneUpdate + +#include +#include +#include + +namespace pk { + +// True for a CED label the scene renderer treats as plain speech (already +// carried by the ASR/diarization transcript, so it is redundant on screen +// unless the caller asks to see it). +bool is_speech_label(const std::string& label); + +// "[mm:ss.s - mm:ss.s]", tenths truncated (not rounded). +std::string format_span(double start, double end); + +// Turns SceneUpdate pieces into printable transcript lines, time-ordered. +// Not thread-safe; one scene stream at a time. +// +// With diarization but no ASR (has_diar && !has_asr) there are no +// utterances, so the closed speaker segments are printed instead, as +// "[mm:ss.s - mm:ss.s] Speaker N". With ASR they are not printed (the +// utterances already carry the speaker). +class SceneRenderer { +public: + SceneRenderer(bool has_diar, bool show_speech, std::function label, + bool has_asr = true); + + // Queues the lines for one SceneUpdate's utterances, sounds and (without + // ASR) speaker segments. + void add(const SceneUpdate& u); + + // Returns (and removes) queued lines with start < safe_until, time-ordered. + // SceneUpdate::safe_until does not cover speaker segments, so when they + // are printed the bound is also capped at the earliest start a later + // speaker segment can have (the start of the earliest still-open segment, + // or the diarized time when none is open). + std::vector flush(double safe_until); + // Returns (and removes) every remaining queued line, time-ordered. + std::vector flush_all(); + +private: + struct Item { + double start; + std::string line; + }; + + bool has_diar_; + bool show_speech_; + bool speaker_lines_; // has_diar && !has_asr + double diarized_ = 0.0; // lower bound on the diarized time seen so far + double speaker_bound_ = 0.0; // no later speaker segment starts before this + std::function label_; + std::vector pending_; +}; + +} // namespace pk diff --git a/src/scene_stream.cpp b/src/scene_stream.cpp new file mode 100644 index 0000000..47028a8 --- /dev/null +++ b/src/scene_stream.cpp @@ -0,0 +1,167 @@ +#include "scene_stream.hpp" + +#include "ced_tagger.hpp" // pk::CedTagger +#include "model.hpp" // pk::Model +#include "transcription_json.hpp" + +#include +#include + +namespace pk { + +SceneStream::SceneStream(const SceneParts& p) { + if (!p.asr && !p.diar && !p.tagger) + throw std::invalid_argument("scene stream needs at least one model"); + if (p.diar) diar_ = std::make_unique(*p.diar, p.diar_latency); + if (p.asr) { + const Model* m = p.asr; + asr_ = std::make_unique( + [m](const std::vector& x) { return m->transcribe_with_timestamps(x, 16000).words; }); + } + if (p.tagger) + sound_ = std::make_unique(p.tagger->scorer(), p.tagger->n_classes(), p.sound); +} + +SceneStream::~SceneStream() = default; + +SceneUpdate SceneStream::feed(const float* pcm, int n, bool is_last) { + SceneUpdate u; + std::vector closed; + long long advanced = 0; + // Until the ASR transcribes, a failure is charged to diarization (as the + // SAS C-API always did), including the ASR audio push. + part_ = diar_ ? ScenePart::Diarization : ScenePart::Asr; + if (diar_) { + advanced = diar_->feed(pcm, n, is_last, closed); + for (const auto& c : closed) { + segs_.push_back({c.speaker, c.start, c.end}); + u.speakers.push_back({c.speaker, c.start, c.end}); + } + } + if (asr_) { + asr_->push(pcm, n); + // With diarization, ASR follows how far diarization has got, and only + // when it advanced (the SAS behavior). Without it, the audio received. + // A commit that the min window skips changes nothing, not even the + // segment pruning below (the SAS behavior). + const bool tick = !diar_ || advanced > 0 || is_last; + const double until = diar_ ? diar_->diarized_until() : t_ + (n > 0 ? n : 0) / 16000.0; + if (tick && asr_->ready(until, is_last)) { + part_ = ScenePart::Asr; + auto committed = asr_->commit(until, is_last); + // Speaker segments known so far: closed ones plus those still open. + std::vector segs = segs_; + 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); + 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(); + segs_.erase(std::remove_if(segs_.begin(), segs_.end(), + [&](const SpeakerSegment& g) { return g.end < commit_sec; }), + segs_.end()); + } + } + if (sound_) { + part_ = ScenePart::Sound; + u.sounds = sound_->feed(pcm, n, is_last); + u.active_sounds = sound_->open_segments(); + } + t_ += (n > 0 ? n : 0) / 16000.0; + if (is_last) finished_ = true; + u.t = t_; + if (diar_) { + part_ = ScenePart::Diarization; + u.active_speakers = diar_->open_segments(); + } + if (finished_) { + u.safe_until = t_; + } else if (sound_ && asr_) { + u.safe_until = std::min(asr_->commit_sec(), sound_->safe_until()); + } else if (sound_) { + u.safe_until = sound_->safe_until(); + } else if (asr_) { + u.safe_until = asr_->commit_sec(); + } else { + u.safe_until = t_; + } + part_ = ScenePart::None; + return u; +} + +std::vector SceneStream::drain_windows() { + return sound_ ? sound_->drain_windows() : std::vector{}; +} + +namespace { + +void append_speaker_segment(std::string& out, const SpeakerSegment& s) { + out += "{\"speaker\":"; append_json_int(out, s.speaker); + 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) { + out += "{\"speaker\":"; append_json_int(out, s.speaker); + out += ",\"start\":"; append_json_float(out, "%.3f", s.start); + out += "}"; +} + +std::string utterances_to_json(const std::vector& utts) { + std::string out = "["; + for (size_t i = 0; i < utts.size(); ++i) { + if (i) out += ","; + out += "{\"speaker\":"; append_json_int(out, utts[i].speaker); + 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); + out += ",\"conf\":"; append_json_float(out, "%.4f", utts[i].conf); + out += "}"; + } + return out + "]"; +} + +std::string words_to_json(const std::vector& words) { + std::string out = "["; + for (size_t i = 0; i < words.size(); ++i) { + if (i) out += ","; + out += "{\"text\":"; append_json_string(out, words[i].text); + out += ",\"start\":"; append_json_float(out, "%.3f", words[i].start); + 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); + out += "}"; + } + return out + "]"; +} + +std::string speakers_to_json(const std::vector& segs) { + std::string out = "["; + for (size_t i = 0; i < segs.size(); ++i) { + if (i) out += ","; + append_speaker_segment(out, segs[i]); + } + return out + "]"; +} + +} // namespace + +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); + 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]); + } + out += "],\"sounds\":" + sound_segments_to_json(u.active_sounds, label); + out += "}}"; + return out; +} + +} // namespace pk diff --git a/src/scene_stream.hpp b/src/scene_stream.hpp new file mode 100644 index 0000000..b123390 --- /dev/null +++ b/src/scene_stream.hpp @@ -0,0 +1,99 @@ +#pragma once +#include "asr_committer.hpp" +#include "diar_pcm_stream.hpp" +#include "sas_merge.hpp" // pk::SpeakerWord, pk::SpeakerUtterance +#include "sound_stream.hpp" // pk::SoundOpts, pk::SoundSegment, pk::SoundWindow + +#include +#include +#include +#include + +namespace pk { + +class Model; +class CedTagger; + +// The models a scene stream runs over. Any may be null; at least one is needed. +struct SceneParts { + const Model* asr = nullptr; + const DiarizationModel* diar = nullptr; + DiarLatency diar_latency = DiarLatency::Model; + CedTagger* tagger = nullptr; // sound events + SoundOpts sound; // sound events +}; + +// What one feed finalized. All times are seconds on the stream clock. +struct SceneUpdate { + double t = 0; // stream time consumed + // Renderer promise: no utterance, word or sound segment returned by a + // later feed() can start before this. It does NOT cover `speakers`: an + // already-open speaker segment can still close later with a start + // earlier than safe_until (diarization does not give that bound, and + // speaker segments are not the thing a renderer commits to the screen). + double safe_until = 0; + std::vector utterances; + std::vector words; + std::vector speakers; // closed this call + std::vector sounds; // closed this call + std::vector active_speakers; + std::vector active_sounds; +}; + +// The part running when feed() threw, so a caller can attribute the error. +enum class ScenePart { None, Diarization, Asr, Sound }; + +// 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 +// 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 +// deferred: it does not see that audio again). If a part throws on the +// is_last feed while diarization is present, finished() is already true +// (diarization takes the is_last chunk before ASR or sound run), so the +// stream ends without flushing whatever the throwing part (or anything +// after it) would otherwise have flushed on that final call. +// +// 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 + ~SceneStream(); + SceneStream(const SceneStream&) = delete; + SceneStream& operator=(const SceneStream&) = delete; + + SceneUpdate feed(const float* pcm, int n, bool is_last); + std::vector drain_windows(); // empty without a tagger + // True after an is_last feed. With diarization it turns true as soon as + // the diarizer takes the is_last chunk, so an is_last feed that throws + // later (diarizer or transcriber) still ends the stream and is not re-run. + bool finished() const { return finished_ || (diar_ && diar_->finished()); } + const DiarPcmStream* diar() const { return diar_.get(); } + // The part that was running when the last feed() threw. + ScenePart failed_part() const { return part_; } + +private: + std::unique_ptr diar_; + std::unique_ptr asr_; + std::unique_ptr sound_; + std::vector segs_; // closed diarization segments not yet behind the commit point + double t_ = 0.0; // stream time consumed + bool finished_ = false; + ScenePart part_ = ScenePart::None; +}; + +// Serialize a SceneUpdate to the scene stream's JSON document shape: +// {"t","utterances","words","speakers","sounds", +// "active":{"speakers","sounds"}}. `label(i)` may return nullptr (emitted +// as ""); pass a function that always returns nullptr when there is no +// tagger (sounds are then always empty, so it is never called). +std::string scene_update_to_json(const SceneUpdate& u, const std::function& label); + +} // namespace pk diff --git a/src/sound_stream.cpp b/src/sound_stream.cpp new file mode 100644 index 0000000..d2ffb3f --- /dev/null +++ b/src/sound_stream.cpp @@ -0,0 +1,183 @@ +#include "sound_stream.hpp" + +#include "transcription_json.hpp" + +#include +#include +#include +#include +#include + +namespace pk { + +namespace { +constexpr float kMinScoredSec = 0.16f; // one CED patch (16 mel frames) +constexpr float kMaxWindowSec = 10.0f; // stay within one CED chunk (10.12 s) +constexpr float kMinHopSec = 0.2f; +} + +std::string validate_sound_opts(const SoundOpts& o, int n_classes) { + if (!(o.hop_sec >= kMinHopSec)) return "hop_sec must be >= 0.2"; + if (!(o.window_sec >= o.hop_sec)) return "window_sec must be >= hop_sec"; + if (!(o.window_sec <= kMaxWindowSec)) return "window_sec must be <= 10"; + if (!(o.off_threshold >= 0.0f && o.off_threshold <= o.on_threshold && o.on_threshold <= 1.0f)) + return "thresholds must satisfy 0 <= off_threshold <= on_threshold <= 1"; + if (!(o.min_duration_sec >= 0.0f)) return "min_duration_sec must be >= 0"; + if (o.top_k < 0 || o.top_k > n_classes) return "top_k must be in [0, number of classes]"; + return ""; +} + +SoundStream::SoundStream(SoundScorer scorer, int n_classes, const SoundOpts& o) + : scorer_(std::move(scorer)), n_classes_(n_classes), o_(o) { + const std::string err = validate_sound_opts(o, n_classes); + if (!err.empty()) throw std::invalid_argument(err); + hop_n_ = std::llround(o.hop_sec * kRate); + win_n_ = std::llround(o.window_sec * kRate); + next_end_ = hop_n_; + open_.assign(n_classes, 0); + open_start_.assign(n_classes, 0.0f); + open_peak_.assign(n_classes, 0.0f); +} + +void SoundStream::close(int c, float end, std::vector& closed) { + open_[c] = 0; + end = std::max(end, open_start_[c]); + if (end - open_start_[c] >= o_.min_duration_sec) + closed.push_back({c, open_start_[c], end, open_peak_[c]}); +} + +void SoundStream::score_window(long long ws, long long we, std::vector& closed) { + if (ws < buf_start_) + throw std::logic_error("sound_stream: window start precedes the buffered range"); + assert(ws >= buf_start_); + const float* p = buf_.data() + (ws - buf_start_); + if (!scorer_(p, (int)(we - ws), probs_) || (int)probs_.size() != n_classes_) + throw std::runtime_error("sound scorer failed"); + const float newest_hop_start = (float)std::max(ws, we - hop_n_) / kRate; + const float oldest_hop_end = (float)std::min(ws + hop_n_, we) / kRate; + for (int c = 0; c < n_classes_; ++c) { + const float s = probs_[c]; + if (!open_[c]) { + if (s >= o_.on_threshold) { + open_[c] = 1; + open_start_[c] = newest_hop_start; + open_peak_[c] = s; + } + } else if (s < o_.off_threshold) { + close(c, oldest_hop_end, closed); + } else { + open_peak_[c] = std::max(open_peak_[c], s); + } + } + if (o_.top_k > 0) { + std::vector idx(n_classes_); + std::iota(idx.begin(), idx.end(), 0); + std::partial_sort(idx.begin(), idx.begin() + o_.top_k, idx.end(), + [&](int a, int b) { return probs_[a] > probs_[b]; }); + SoundWindow w{(float)ws / kRate, (float)we / kRate, {}}; + for (int i = 0; i < o_.top_k; ++i) w.top.push_back({idx[i], probs_[idx[i]]}); + windows_.push_back(std::move(w)); + } + scored_end_ = we; +} + +std::vector SoundStream::feed(const float* pcm, int n, bool is_last) { + std::vector closed; + if (finished_) return closed; + if (n > 0 && pcm) { + buf_.insert(buf_.end(), pcm, pcm + n); + samples_in_ += n; + } + while (samples_in_ >= next_end_) { + score_window(std::max(0LL, next_end_ - win_n_), next_end_, closed); + next_end_ += hop_n_; + } + if (is_last) { + const long long tail = samples_in_ - scored_end_; + if (tail >= (long long)std::llround(kMinScoredSec * kRate)) + score_window(std::max(0LL, samples_in_ - win_n_), samples_in_, closed); + for (int c = 0; c < n_classes_; ++c) + if (open_[c]) close(c, (float)time(), closed); + finished_ = true; + } + // Keep only what the next window needs. Trimming to `next_end_ - win_n_` + // alone looks ahead to a hop that has not arrived yet, which over-trims + // whenever the stream ends between two hops: the is_last tail window + // above starts at `samples_in_ - win_n_`, earlier than that, and would + // read before the buffer start. Clamp the lookahead to what has + // actually streamed. + const long long keep_from = std::max(0LL, std::min(next_end_, samples_in_) - win_n_); + if (keep_from > buf_start_) { + const long long drop = std::min(keep_from - buf_start_, (long long)buf_.size()); + buf_.erase(buf_.begin(), buf_.begin() + drop); + buf_start_ += drop; + } + return closed; +} + +std::vector SoundStream::open_segments() const { + std::vector out; + for (int c = 0; c < n_classes_; ++c) + if (open_[c]) out.push_back({c, open_start_[c], (float)time(), open_peak_[c]}); + return out; +} + +std::vector SoundStream::drain_windows() { + std::vector out; + out.swap(windows_); + return out; +} + +double SoundStream::safe_until() const { + if (finished_) return time(); + // Two ways a not-yet-returned segment could still start earlier than + // expected: a regular window scored after scored_end_ (its newest hop + // begins at scored_end_), or an is_last tail window scored on the next + // feed() call, whose newest hop begins at max(0, samples_in_ - hop_n_) + // (a tail window is shorter than win_n_ but never shorter than hop_n_). + // An already-open class keeps its (earlier) start regardless. + const long long tail_hop_start = std::max(0LL, samples_in_ - hop_n_); + double t = (double)std::min(scored_end_, tail_hop_start) / kRate; + for (int c = 0; c < n_classes_; ++c) + if (open_[c]) t = std::min(t, (double)open_start_[c]); + return t; +} + +std::string sound_segments_to_json(const std::vector& s, + const std::function& label) { + std::string out = "["; + for (size_t i = 0; i < s.size(); ++i) { + if (i) out += ","; + const char* l = label(s[i].cls); + out += "{\"index\":"; append_json_int(out, s[i].cls); + out += ",\"label\":"; append_json_string(out, l ? l : ""); + out += ",\"start\":"; append_json_float(out, "%.3f", s[i].start); + out += ",\"end\":"; append_json_float(out, "%.3f", s[i].end); + out += ",\"peak\":"; append_json_float(out, "%.4f", s[i].peak); + out += "}"; + } + return out + "]"; +} + +std::string sound_windows_to_json(const std::vector& w, + const std::function& label) { + std::string out = "["; + for (size_t i = 0; i < w.size(); ++i) { + if (i) out += ","; + out += "{\"start\":"; append_json_float(out, "%.3f", w[i].start); + out += ",\"end\":"; append_json_float(out, "%.3f", w[i].end); + out += ",\"tags\":["; + for (size_t j = 0; j < w[i].top.size(); ++j) { + if (j) out += ","; + const char* l = label(w[i].top[j].first); + out += "{\"index\":"; append_json_int(out, w[i].top[j].first); + out += ",\"label\":"; append_json_string(out, l ? l : ""); + out += ",\"score\":"; append_json_float(out, "%.4f", w[i].top[j].second); + out += "}"; + } + out += "]}"; + } + return out + "]"; +} + +} // namespace pk diff --git a/src/sound_stream.hpp b/src/sound_stream.hpp new file mode 100644 index 0000000..fdbafc8 --- /dev/null +++ b/src/sound_stream.hpp @@ -0,0 +1,97 @@ +#pragma once +#include "ced_tagger.hpp" // SoundScorer + +#include +#include +#include +#include + +namespace pk { + +struct SoundOpts { + float window_sec = 3.0f; // audio scored per window (CED clip length) + float hop_sec = 1.0f; // a window ends every hop + float on_threshold = 0.4f; // a class opens at score >= on + float off_threshold = 0.3f; // and closes at score < off + float min_duration_sec = 0.3f; // shorter segments are dropped + int top_k = 5; // per-window scores kept for drain_windows +}; + +// "" when valid, otherwise a message naming the bad field. +std::string validate_sound_opts(const SoundOpts& o, int n_classes); + +// A closed (or, from open_segments, still running) sound event. +struct SoundSegment { + int cls; // class index + float start, end; // seconds from stream start + float peak; // highest score while open +}; + +// One scored window: its span and the top-k (class, score), score-descending. +struct SoundWindow { + float start, end; + std::vector> top; +}; + +// Sliding-window sound-event detection over live 16 kHz PCM. See +// docs/sound.md for the timing rules. Not thread-safe. +class SoundStream { +public: + // Throws std::invalid_argument when validate_sound_opts fails. + SoundStream(SoundScorer scorer, int n_classes, const SoundOpts& o); + + // Append PCM and score every window whose end has arrived. Returns the + // segments that closed in this call. `is_last` scores the tail and closes + // everything. Throws std::runtime_error when the scorer fails. + std::vector feed(const float* pcm, int n, bool is_last); + + // Classes open right now, with end = time(). + std::vector open_segments() const; + + // Windows scored since the previous drain. The queue grows by one entry + // per hop until drained: drain regularly, or set top_k = 0 to keep no + // scores. + std::vector drain_windows(); + + double time() const { return (double)samples_in_ / kRate; } + // No segment returned by a later feed() call can start before this + // time. While the stream is still open this is a lower bound over: any + // class already open (its start won't move), the last window actually + // scored (scored_end_), and the earliest a still-unscored is_last tail + // window could open a brand-new class (its newest hop, which starts at + // max(0, samples_in_ - hop_n_): a tail window can be shorter than a + // full window but is never shorter than one hop). Once finished(), it + // equals time() exactly. + double safe_until() const; + bool finished() const { return finished_; } + +private: + static constexpr int kRate = 16000; + void score_window(long long win_start, long long win_end, std::vector& closed); + void close(int cls, float end, std::vector& closed); + + SoundScorer scorer_; + int n_classes_; + SoundOpts o_; + long long hop_n_, win_n_; + std::vector buf_; // PCM from sample buf_start_ + long long buf_start_ = 0; + long long samples_in_ = 0; + long long next_end_; // sample index where the next window ends + long long scored_end_ = 0; // end of the last scored window + std::vector open_; // per class + std::vector open_start_, open_peak_; + std::vector windows_; // scored windows not yet drained (empty with top_k = 0) + std::vector probs_; + bool finished_ = false; +}; + +// JSON arrays for the C-API. `label(i)` may return nullptr (emitted as ""). +// segments: [{"index":..,"label":..,"start":..,"end":..,"peak":..}] +// windows: [{"start":..,"end":..,"tags":[{"index":..,"label":..,"score":..}]}] +std::string sound_segments_to_json(const std::vector& s, + const std::function& label); +std::string sound_windows_to_json(const std::vector& w, + const std::function& label); + +} // namespace pk diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 96c92e9..b143508 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -73,8 +73,21 @@ pk_add_test(test_capi_ctc_logits) pk_add_test(test_diarization) pk_add_test(test_diarization_accuracy) pk_add_test(test_sas_merge) +pk_add_test(test_asr_committer) pk_add_test(test_combined_offline) pk_add_test(test_streaming_diarization) +pk_add_test(test_sound_stream) +pk_add_test(test_ced_parity) +set_tests_properties(test_ced_parity PROPERTIES LABELS "model") +pk_add_test(test_sound_capi) +target_compile_definitions(test_sound_capi PRIVATE PK_CED_SOURCE_DIR="${CMAKE_SOURCE_DIR}/third_party/ced.cpp") +set_tests_properties(test_sound_capi PROPERTIES LABELS "model") +pk_add_test(test_scene_stream) +target_compile_definitions(test_scene_stream PRIVATE + PK_SOURCE_DIR="${CMAKE_SOURCE_DIR}" + PK_CED_SOURCE_DIR="${CMAKE_SOURCE_DIR}/third_party/ced.cpp") +set_tests_properties(test_scene_stream PROPERTIES LABELS "model") +pk_add_test(test_scene_render) if(TARGET parakeet-cli) add_test(NAME cli_version_long COMMAND $ --version) diff --git a/tests/test_asr_committer.cpp b/tests/test_asr_committer.cpp new file mode 100644 index 0000000..f5f4bf9 --- /dev/null +++ b/tests/test_asr_committer.cpp @@ -0,0 +1,169 @@ +// Unit test for pk::AsrCommitter with a fake transcriber (no model). +#include "asr_committer.hpp" + +#include +#include +#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) + +// The PCM value at each sample is its absolute time in seconds, so the fake +// knows where the buffer starts. +// Words start at `speech_from` seconds (silence before it). +static Transcriber fake(double speech_from = 0.0) { + return [speech_from](const std::vector& pcm) { + std::vector w; + if (pcm.empty()) return w; + const double t0 = pcm[0]; + const double dur = pcm.size() / 16000.0; + const int first = (int)std::ceil(std::max(t0, speech_from) / 0.5 - 1e-6); + for (int i = first;; ++i) { + const double s = i * 0.5 + 0.05 - t0, e = i * 0.5 + 0.45 - t0; + if (e > dur) break; + w.push_back({"w" + std::to_string(i), (float)s, (float)e, 0.9f}); + } + return w; + }; +} +static std::vector timeline(double a, double b) { + std::vector x; + for (long i = std::lround(a * 16000); i < std::lround(b * 16000); ++i) x.push_back((float)(i / 16000.0)); + return x; +} + +static void test_min_window() { + AsrCommitter c(fake()); + auto x = timeline(0, 3); + c.push(x.data(), (int)x.size()); + CHECK(c.commit(3.0, false).empty()); // 3 s < 4 s window + CHECK(c.commit_sec() == 0.0); +} + +static void test_right_context_and_resume() { + AsrCommitter c(fake()); + auto x = timeline(0, 5); + c.push(x.data(), (int)x.size()); + auto w = c.commit(5.0, false); + // words must end <= 5 - 1 = 4 s: w0..w7 (w7 ends 3.95) + CHECK(w.size() == 8); + if (w.size() == 8) { CHECK(w.front().text == "w0"); CHECK(w.back().text == "w7"); } + CHECK(std::fabs(c.commit_sec() - 3.95) < 1e-3); // resumes after the last word + auto y = timeline(5, 10); + c.push(y.data(), (int)y.size()); + auto w2 = c.commit(10.0, false); + CHECK(!w2.empty() && w2.front().text == "w8"); // no repeat, no gap + CHECK(std::fabs(w2.front().start - 4.05f) < 1e-3); // absolute time +} + +static void test_is_last_commits_all() { + AsrCommitter c(fake()); + auto x = timeline(0, 2); + c.push(x.data(), (int)x.size()); + auto w = c.commit(2.0, true); + CHECK(w.size() == 4); // w0..w3, no right-context hold +} + +static void test_duplicate_word_dropped() { + // A transcriber that hears the last committed word again at the new start. + int calls = 0; + AsrCommitter c([&](const std::vector& pcm) { + ++calls; + std::vector w; + if (calls == 1) { w = {{"hello", 0.5f, 1.0f, 0.9f}, {"there", 1.1f, 1.5f, 0.9f}}; } + else { w = {{"there", 0.0f, 0.2f, 0.9f}, {"friend", 0.4f, 0.9f, 0.9f}}; } + (void)pcm; + return w; + }); + std::vector x(16000 * 5, 0.0f); + c.push(x.data(), (int)x.size()); + auto a = c.commit(5.0, false); + CHECK(a.size() == 2); + c.push(x.data(), (int)x.size()); + auto b = c.commit(10.0, true); + CHECK(b.size() == 1 && b[0].text == "friend"); +} + +// A word whose reported span straddles the cut (its end lands past the right +// context, so it does not commit) must not have the cut land exactly on its +// start: the onset margin must back off from it, and the word must come back +// intact, at its correct absolute times, once more audio gives it context. +static void test_onset_margin_backoff() { + AsrCommitter c([](const std::vector& pcm) { + std::vector w; + if (pcm.empty()) return w; + const double t0 = pcm[0]; + const double dur = pcm.size() / 16000.0; + const double abs_start = 3.8, abs_end = 4.4; + if (abs_start >= t0 && abs_end <= t0 + dur) + w.push_back({"onset", (float)(abs_start - t0), (float)(abs_end - t0), 0.9f}); + return w; + }); + auto x = timeline(0, 5); + c.push(x.data(), (int)x.size()); + auto w = c.commit(5.0, false); + CHECK(w.empty()); // straddles limit (4.0), so keep == 0 + CHECK(std::fabs(c.commit_sec() - 3.5) < 1e-3); // 3.8 - kOnsetMargin(0.3), not 3.8 + // More audio gives the word its right context; it must return intact. + auto y = timeline(5, 10); + c.push(y.data(), (int)y.size()); + auto w2 = c.commit(10.0, true); + CHECK(w2.size() == 1); + if (w2.size() == 1) { + CHECK(w2[0].text == "onset"); + CHECK(std::fabs(w2[0].start - 3.8f) < 1e-3); + CHECK(std::fabs(w2[0].end - 4.4f) < 1e-3); + } +} + +// Non-speech: no words at all. The commit point must still advance and the +// buffer must not grow with the stream. +static void test_silence_released() { + AsrCommitter c([](const std::vector&) { return std::vector{}; }); + size_t max_buffered = 0; + for (int s = 0; s < 20; ++s) { + auto x = timeline(s, s + 1); + c.push(x.data(), (int)x.size()); + CHECK(c.commit(s + 1.0, false).empty()); + max_buffered = std::max(max_buffered, c.buffered_samples()); + } + CHECK(c.commit_sec() >= 15.0); // only the right context is held back + CHECK(max_buffered <= (size_t)(5 * 16000)); // min window + one piece + CHECK(c.buffered_samples() <= (size_t)(5 * 16000)); +} + +// Speech after a silent stretch commits every word once, at absolute times. +static void test_speech_after_silence() { + AsrCommitter c(fake(10.0)); + std::vector all; + for (int s = 0; s < 20; ++s) { + auto x = timeline(s, s + 1); + c.push(x.data(), (int)x.size()); + auto w = c.commit(s + 1.0, s == 19); + all.insert(all.end(), w.begin(), w.end()); + CHECK(c.buffered_samples() <= (size_t)(5 * 16000)); + } + CHECK(all.size() == 20); // w20..w39 + for (size_t i = 0; i < all.size(); ++i) { + const int k = 20 + (int)i; + CHECK(all[i].text == "w" + std::to_string(k)); + CHECK(std::fabs(all[i].start - (k * 0.5 + 0.05)) < 1e-3); + CHECK(std::fabs(all[i].end - (k * 0.5 + 0.45)) < 1e-3); + } +} + +int main() { + test_silence_released(); + test_speech_after_silence(); + test_min_window(); + test_right_context_and_resume(); + test_is_last_commits_all(); + test_duplicate_word_dropped(); + test_onset_margin_backoff(); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/tests/test_audio_io.cpp b/tests/test_audio_io.cpp index ec9b671..9ec1705 100644 --- a/tests/test_audio_io.cpp +++ b/tests/test_audio_io.cpp @@ -8,7 +8,6 @@ #include // dr_wav writer is only needed in the test; include without implementation -// (DR_WAV_IMPLEMENTATION lives in audio_io.cpp, linked via parakeet) #include "dr_wav.h" static void write_sine(const char* path, int sr, int n, float freq) { diff --git a/tests/test_capi.cpp b/tests/test_capi.cpp index 43b81e4..0f3ff2c 100644 --- a/tests/test_capi.cpp +++ b/tests/test_capi.cpp @@ -38,6 +38,12 @@ int main() { return 1; } + // A NULL context reports PARAKEET_MODEL_KIND_NONE. + if (parakeet_capi_model_kind(nullptr) != PARAKEET_MODEL_KIND_NONE) { + std::fprintf(stderr, "test_capi: model_kind(NULL) != NONE\n"); + return 1; + } + // The 110m anchor (PARAKEET_TEST_GGUF) and the prompt/multilingual model // (PARAKEET_TEST_GGUF_NEMOTRON) are independent: each block runs only when // its env var is set. If NEITHER is set the test skips (77). @@ -52,6 +58,12 @@ int main() { return 1; } + if (parakeet_capi_model_kind(ctx) != PARAKEET_MODEL_KIND_ASR) { + std::fprintf(stderr, "test_capi: model_kind(ASR ctx) != ASR\n"); + parakeet_capi_free(ctx); + return 1; + } + // decoder == 2 -> TDT/transducer head. char* text = parakeet_capi_transcribe_path(ctx, "tests/fixtures/speech.wav", 2); if (!text) { diff --git a/tests/test_ced_parity.cpp b/tests/test_ced_parity.cpp new file mode 100644 index 0000000..86e1308 --- /dev/null +++ b/tests/test_ced_parity.cpp @@ -0,0 +1,66 @@ +// In-tree check that ced.cpp built against parakeet's ggml (with its patches) +// gives the same output as standalone ced.cpp: probabilities vs ced's own +// PyTorch baseline, at the level standalone ced reaches on CPU f32 (1.7e-7). +// Needs PARAKEET_TEST_CED_GGUF and PARAKEET_TEST_CED_BASELINE. Exit 77 when +// either is unset (the baseline .baseline.gguf fixtures are gitignored in +// ced.cpp, so there is no in-tree default path). +#include "ced_tagger.hpp" +#include "ggml.h" +#include "gguf.h" + +#include +#include +#include +#include +#include +#include + +static bool load_f32(const std::string& path, const char* name, std::vector& out) { + ggml_context* ctx = nullptr; + gguf_init_params p{false, &ctx}; + gguf_context* g = gguf_init_from_file(path.c_str(), p); + if (!g) return false; + ggml_tensor* t = ggml_get_tensor(ctx, name); + bool ok = t != nullptr; + if (ok) { + out.resize((size_t)ggml_nelements(t)); + std::memcpy(out.data(), t->data, out.size() * sizeof(float)); + } + gguf_free(g); + ggml_free(ctx); + return ok; +} + +int main() { + const char* gguf = std::getenv("PARAKEET_TEST_CED_GGUF"); + const char* base_env = std::getenv("PARAKEET_TEST_CED_BASELINE"); + if (!gguf || !base_env) { + std::fprintf(stderr, "SKIP: PARAKEET_TEST_CED_GGUF / PARAKEET_TEST_CED_BASELINE unset\n"); + return 77; + } + if (!pk::CedTagger::available()) { std::fprintf(stderr, "SKIP: built without CED\n"); return 77; } + const std::string base = base_env; + std::vector wav, ref; + if (!load_f32(base, "audio_waveform", wav) || !load_f32(base, "probs", ref)) { + std::fprintf(stderr, "FAIL: baseline %s\n", base.c_str()); + return 1; + } + if (!pk::gguf_is_ced(gguf)) { std::fprintf(stderr, "FAIL: gguf_is_ced\n"); return 1; } + auto tagger = pk::CedTagger::load(gguf); + if (!tagger) { std::fprintf(stderr, "FAIL: load\n"); return 1; } + std::vector probs; + if (!tagger->scorer()(wav.data(), (int)wav.size(), probs)) { std::fprintf(stderr, "FAIL: score\n"); return 1; } + if (probs.size() != ref.size()) { std::fprintf(stderr, "FAIL: size\n"); return 1; } + double md = 0; + for (size_t i = 0; i < ref.size(); ++i) md = std::fmax(md, std::fabs((double)probs[i] - ref[i])); + std::fprintf(stderr, "max|d| = %.3e (standalone ced f32 CPU: 1.7e-7)\n", md); + // Off-CPU the GPU matmuls move probs by ~1e-4 (ced.cpp PR #2), so only the + // CPU gets the tight gate. + const char* dev = std::getenv("CED_DEVICE"); + const double tol = (dev && std::strcmp(dev, "cpu") == 0) ? 1e-6 : 1e-3; + if (md > tol) { std::fprintf(stderr, "FAIL: tol %.0e\n", tol); return 1; } + const char* label0 = tagger->label(0); + if (!label0 || std::strcmp(label0, "Speech") != 0) { std::fprintf(stderr, "FAIL: label 0\n"); return 1; } + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/tests/test_combined_offline.cpp b/tests/test_combined_offline.cpp index 3311d3e..19c4b6c 100644 --- a/tests/test_combined_offline.cpp +++ b/tests/test_combined_offline.cpp @@ -205,6 +205,13 @@ int main() { parakeet_capi_free(diar); return 1; } + if (parakeet_capi_model_kind(asr) != PARAKEET_MODEL_KIND_ASR || + parakeet_capi_model_kind(diar) != PARAKEET_MODEL_KIND_DIARIZATION) { + std::fprintf(stderr, "test_combined_offline: model_kind mismatch\n"); + parakeet_capi_free(asr); + parakeet_capi_free(diar); + return 1; + } pk::Audio audio; if (!pk::load_audio_16k_mono("tests/fixtures/two_speakers.wav", audio)) { std::fprintf(stderr, "test_combined_offline: cannot read the fixture\n"); diff --git a/tests/test_scene_render.cpp b/tests/test_scene_render.cpp new file mode 100644 index 0000000..b9bb681 --- /dev/null +++ b/tests/test_scene_render.cpp @@ -0,0 +1,82 @@ +#include "scene_render.hpp" +#include "scene_stream.hpp" + +#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) + +int main() { + auto label = [](int i) -> const char* { + return i == 0 ? "Speech" : i == 359 ? "Knock" : i == 42 ? "Speech synthesizer" : "Other"; + }; + CHECK(format_span(9.0, 10.04) == "[00:09.0 - 00:10.0]"); + CHECK(format_span(75.25, 80.0) == "[01:15.2 - 01:20.0]"); + CHECK(is_speech_label("Speech") && is_speech_label("Male speech, man speaking") && !is_speech_label("Knock")); + CHECK(is_speech_label("Speech synthesizer")); + + SceneRenderer r(/*has_diar=*/true, /*show_speech=*/false, label); + SceneUpdate u1; + u1.sounds = {{359, 9.0f, 10.0f, 0.81f}, {0, 10.5f, 11.8f, 0.9f}}; + u1.utterances = {{0, "who is it", 10.9f, 11.6f, 0.93f}}; + r.add(u1); + auto early = r.flush(10.0); + CHECK(early.size() == 1 && early[0] == "[00:09.0 - 00:10.0] (Knock 0.81)"); + auto rest = r.flush_all(); + CHECK(rest.size() == 1 && rest[0] == "[00:10.9 - 00:11.6] Speaker 0: who is it"); // Speech hidden + + SceneRenderer nd(/*has_diar=*/false, /*show_speech=*/true, label); + SceneUpdate u2; + u2.utterances = {{-1, "hello", 1.0f, 1.5f, 0.9f}}; + u2.sounds = {{0, 0.5f, 2.0f, 0.9f}}; + nd.add(u2); + auto lines = nd.flush_all(); + CHECK(lines.size() == 2 && lines[0] == "[00:00.5 - 00:02.0] (Speech 0.90)" && + lines[1] == "[00:01.0 - 00:01.5] hello"); + + // "Speech synthesizer" is a child of Speech in the AudioSet ontology; CED + // emits it over clean narration, so it hides the same way "Speech" does. + SceneRenderer hide_synth(/*has_diar=*/false, /*show_speech=*/false, label); + SceneUpdate u3; + u3.sounds = {{42, 2.0f, 4.0f, 0.8f}}; + hide_synth.add(u3); + CHECK(hide_synth.flush_all().empty()); + + SceneRenderer show_synth(/*has_diar=*/false, /*show_speech=*/true, label); + show_synth.add(u3); + auto synth_lines = show_synth.flush_all(); + CHECK(synth_lines.size() == 1 && synth_lines[0] == "[00:02.0 - 00:04.0] (Speech synthesizer 0.80)"); + + // Diarization without ASR: closed speaker segments become lines, ordered + // with sounds, and held while an earlier-starting segment is still open. + SceneRenderer dz(/*has_diar=*/true, /*show_speech=*/false, label, /*has_asr=*/false); + SceneUpdate d1; + d1.speakers = {{1, 2.0f, 4.0f}}; + d1.active_speakers = {{0, 1.0f, 5.0f}}; // speaker 0 open since 1.0 s + d1.sounds = {{359, 3.0f, 3.5f, 0.7f}}; + dz.add(d1); + CHECK(dz.flush(5.0).empty()); // speaker 0 may still close with start 1.0 + SceneUpdate d2; + d2.speakers = {{0, 1.0f, 6.0f}}; + d2.active_speakers = {{1, 6.5f, 7.0f}}; + dz.add(d2); + auto dl = dz.flush(7.0); + CHECK(dl.size() == 3 && dl[0] == "[00:01.0 - 00:06.0] Speaker 0" && + dl[1] == "[00:02.0 - 00:04.0] Speaker 1" && dl[2] == "[00:03.0 - 00:03.5] (Knock 0.70)"); + SceneUpdate d3; + d3.speakers = {{1, 6.5f, 8.0f}}; + dz.add(d3); + auto dr = dz.flush_all(); + CHECK(dr.size() == 1 && dr[0] == "[00:06.5 - 00:08.0] Speaker 1"); + + // With ASR, speaker segments are not printed (utterances carry the speaker). + SceneRenderer withasr(/*has_diar=*/true, /*show_speech=*/false, label); + withasr.add(d1); + auto wl = withasr.flush_all(); + CHECK(wl.size() == 1 && wl[0] == "[00:03.0 - 00:03.5] (Knock 0.70)"); + + if (failures) return 1; + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/tests/test_scene_stream.cpp b/tests/test_scene_stream.cpp new file mode 100644 index 0000000..9034af8 --- /dev/null +++ b/tests/test_scene_stream.cpp @@ -0,0 +1,215 @@ +// Scene stream on real models: composition changes nothing. +// 1. asr+diar+tagger: utterances == sas_stream on the same PCM and pieces; +// sounds == sound_stream alone. +// 2. scene_sound_only: tagger alone returns sounds and empty word arrays. +// 3. asr+tagger (no diar): words have speaker -1. +// 4. scene_wrong_kinds: a tagger passed as ASR is rejected with a message, +// and so is an unknown diarization latency mode. +// Needs PARAKEET_TEST_GGUF, PARAKEET_TEST_DIAR_GGUF, PARAKEET_TEST_CED_GGUF. +#include "parakeet_capi.h" +#include "audio_io.hpp" +#include "ced_tagger.hpp" + +#include +#include +#include +#include +#include + +// Load a WAV as 16 kHz mono (pk::load_audio_16k_mono downmixes and resamples). +static bool load16k(const std::string& path, std::vector& x) { + pk::Audio a; + if (!pk::load_audio_16k_mono(path, a)) return false; + x = std::move(a.samples); + return true; +} + +static int fails = 0; +#define CHECK(c, ...) do { if (!(c)) { std::fprintf(stderr, "FAIL: " __VA_ARGS__); std::fprintf(stderr, "\n"); ++fails; } } while (0) + +static std::string feed_all_scene(parakeet_scene_stream* s, const std::vector& pcm, int piece) { + std::string all; + for (size_t i = 0; i < pcm.size(); i += piece) { + const int n = (int)std::min(piece, pcm.size() - i); + char* j = parakeet_capi_scene_stream_feed_json(s, pcm.data() + i, n, i + piece >= pcm.size()); + if (!j) { std::fprintf(stderr, "feed: %s\n", parakeet_capi_scene_stream_last_error(s)); ++fails; return all; } + all += j; all += "\n"; + parakeet_capi_free_string(j); + } + return all; +} + +// One "speaker|text|start|end;" record, %.3f on the times, comparable +// directly against the same formatting applied to a parakeet_sas_result. +static std::string fmt_utt(int speaker, const std::string& text, float start, float end) { + char b[512]; + std::snprintf(b, sizeof(b), "%d|%s|%.3f|%.3f;", speaker, text.c_str(), start, end); + return b; +} + +// Pull full utterance records ("utterances" array only, not "words", which +// also has a "text" field) from the JSON lines, in order. Un-escapes \" and +// \\ so a quote inside a word does not truncate early. +static std::string scene_utterances(const std::string& jl) { + std::string out; + size_t p = 0; + while ((p = jl.find("\"utterances\":[", p)) != std::string::npos) { + size_t q = p + 14; + const size_t end = jl.find("],\"words\"", q); + while ((q = jl.find("{\"speaker\":", q)) != std::string::npos && q < end) { + const int speaker = std::atoi(jl.c_str() + q + 11); + size_t r = jl.find("\"text\":\"", q); + if (r == std::string::npos || r >= end) break; + r += 8; + std::string text; + while (r < jl.size() && jl[r] != '"') { + if (jl[r] == '\\' && r + 1 < jl.size()) { text += jl[r + 1]; r += 2; } + else { text += jl[r]; ++r; } + } + float start = 0, e = 0; + std::sscanf(jl.c_str() + r, "\",\"start\":%f,\"end\":%f", &start, &e); + out += fmt_utt(speaker, text, start, e); + q = r; + } + p = end; + } + return out; +} + +// Only the top-level "sounds" array (closed this call) counts here; the +// "active" object has its own "sounds" array (still-open segments, growing +// "end" on every call) that would otherwise match the same "\"sounds\":[" +// needle and pollute the comparison. Returns both the "index start end;" +// tuples (for the sound_stream comparison) and the raw slices (to check a +// label appears among CLOSED segments only, not merely opened-and-still-open +// ones in "active"). +static void scene_closed_sounds(const std::string& jl, std::string& tuples, std::string& raw) { + size_t line_start = 0; + while (line_start < jl.size()) { + size_t line_end = jl.find('\n', line_start); + if (line_end == std::string::npos) line_end = jl.size(); + const size_t active_pos = jl.find("\"active\":", line_start); + const size_t scan_end = (active_pos != std::string::npos && active_pos < line_end) ? active_pos : line_end; + const size_t p = jl.find("\"sounds\":[", line_start); + if (p != std::string::npos && p < scan_end) { + size_t q = p, end = jl.find(']', p); + raw += jl.substr(p, end - p); + while ((q = jl.find("{\"index\":", q)) != std::string::npos && q < end) { + int idx; float s0, s1; + std::sscanf(jl.c_str() + q, "{\"index\":%d,\"label\":\"%*[^\"]\",\"start\":%f,\"end\":%f", &idx, &s0, &s1); + char b[128]; + std::snprintf(b, sizeof(b), "%d %.3f %.3f;", idx, s0, s1); + tuples += b; + ++q; + } + } + line_start = line_end + 1; + } +} + +int main() { + const char* a = std::getenv("PARAKEET_TEST_GGUF"); + const char* d = std::getenv("PARAKEET_TEST_DIAR_GGUF"); + const char* c = std::getenv("PARAKEET_TEST_CED_GGUF"); + if (!a || !d || !c) { std::fprintf(stderr, "SKIP: needs ASR, diar and CED GGUFs\n"); return 77; } + if (!pk::CedTagger::available()) { std::fprintf(stderr, "SKIP: built without CED\n"); return 77; } + parakeet_ctx* asr = parakeet_capi_load(a); + parakeet_ctx* diar = parakeet_capi_load(d); + parakeet_ctx* tag = parakeet_capi_load(c); + if (!asr || !diar || !tag) { std::fprintf(stderr, "FAIL: load\n"); return 1; } + + // Speech, then a rooster, then speech again. + std::vector pcm, x; + load16k(std::string(PK_SOURCE_DIR) + "/tests/fixtures/two_speakers.wav", x); + pcm.insert(pcm.end(), x.begin(), x.end()); + load16k(std::string(PK_CED_SOURCE_DIR) + "/benchmarks/demo/clips/rooster.wav", x); + pcm.insert(pcm.end(), x.begin(), x.end()); + load16k(std::string(PK_SOURCE_DIR) + "/tests/fixtures/speech.wav", x); + pcm.insert(pcm.end(), x.begin(), x.end()); + const int piece = 8000; // 0.5 s, like test_combined_offline + + // 1a. reference utterances from sas_stream + std::string ref; + { + parakeet_sas_stream* ss = parakeet_capi_sas_stream_begin_latency(asr, diar, PARAKEET_DIAR_LATENCY_LOW); + for (size_t i = 0; i < pcm.size(); i += piece) { + parakeet_sas_result* r = nullptr; int nr = 0; + const int n = (int)std::min(piece, pcm.size() - i); + CHECK(parakeet_capi_sas_stream_feed(ss, pcm.data() + i, n, i + piece >= pcm.size(), &r, &nr) == 0, + "sas reference feed: %s", parakeet_capi_last_error(asr)); + for (int k = 0; k < nr; ++k) ref += fmt_utt(r[k].speaker, r[k].text, r[k].start, r[k].end); + parakeet_capi_free_sas_results(r, nr); + } + parakeet_capi_sas_stream_free(ss); + } + CHECK(!ref.empty(), "sas reference produced no utterances"); + // 1b. reference sounds from sound_stream + std::string ref_sounds; + { + parakeet_sound_stream* s = parakeet_capi_sound_stream_begin(tag, nullptr); + for (size_t i = 0; i < pcm.size(); i += piece) { + parakeet_sound_segment* o = nullptr; int no = 0; + const int n = (int)std::min(piece, pcm.size() - i); + CHECK(parakeet_capi_sound_stream_feed(s, pcm.data() + i, n, i + piece >= pcm.size(), &o, &no) == 0, + "sound reference feed: %s", parakeet_capi_last_error(tag)); + for (int k = 0; k < no; ++k) { + char b[128]; + std::snprintf(b, sizeof(b), "%d %.3f %.3f;", o[k].class_index, o[k].start, o[k].end); + ref_sounds += b; + } + parakeet_capi_free_sound_segments(o); + } + parakeet_capi_sound_stream_free(s); + } + CHECK(!ref_sounds.empty(), "sound reference produced no segments"); + // 1c. the scene with all three + parakeet_scene_opts o; + parakeet_capi_scene_opts_default(&o); + o.diar_latency = PARAKEET_DIAR_LATENCY_LOW; + parakeet_scene_stream* sc = parakeet_capi_scene_stream_begin(asr, diar, tag, &o); + CHECK(sc != nullptr, "scene begin"); + const std::string jl = feed_all_scene(sc, pcm, piece); + parakeet_capi_scene_stream_free(sc); + const std::string got_utts = scene_utterances(jl); + CHECK(got_utts == ref, "scene utterances != sas_stream\n got: %s\n ref: %s", got_utts.c_str(), ref.c_str()); + std::string got_sounds, closed_sounds_raw; + scene_closed_sounds(jl, got_sounds, closed_sounds_raw); + CHECK(got_sounds == ref_sounds, "scene sounds != sound_stream\n %s\n %s", got_sounds.c_str(), ref_sounds.c_str()); + CHECK(closed_sounds_raw.find("Chicken, rooster") != std::string::npos, "rooster in a closed scene sound segment"); + + // 2. scene_sound_only + parakeet_scene_stream* so = parakeet_capi_scene_stream_begin(nullptr, nullptr, tag, nullptr); + CHECK(so != nullptr, "sound-only begin"); + const std::string jso = feed_all_scene(so, pcm, piece); + CHECK(jso.find("\"words\":[{") == std::string::npos, "sound-only has words"); + CHECK(jso.find("\"sounds\":[{") != std::string::npos, "sound-only has no sounds"); + parakeet_capi_scene_stream_free(so); + + // 3. asr+tagger without diarization: speaker -1 + parakeet_scene_stream* at = parakeet_capi_scene_stream_begin(asr, nullptr, tag, nullptr); + const std::string jat = feed_all_scene(at, pcm, piece); + CHECK(jat.find("\"speaker\":-1") != std::string::npos, "no-diar words carry speaker -1"); + CHECK(jat.find("\"speaker\":0") == std::string::npos, "no-diar words carry a speaker"); + parakeet_capi_scene_stream_free(at); + + // 4. scene_wrong_kinds + CHECK(parakeet_capi_scene_stream_begin(tag, nullptr, nullptr, nullptr) == nullptr, "tagger as ASR accepted"); + CHECK(std::strstr(parakeet_capi_last_error(tag), "CED sound model") != nullptr, "message: %s", + parakeet_capi_last_error(tag)); + CHECK(parakeet_capi_scene_stream_begin(nullptr, nullptr, nullptr, nullptr) == nullptr, "no parts accepted"); + parakeet_scene_opts bad; + parakeet_capi_scene_opts_default(&bad); + bad.diar_latency = 42; + CHECK(parakeet_capi_scene_stream_begin(nullptr, diar, nullptr, &bad) == nullptr, "latency 42 accepted"); + CHECK(std::strcmp(parakeet_capi_last_error(diar), "unknown diarization latency mode") == 0, "message: %s", + parakeet_capi_last_error(diar)); + // Without a diar ctx the latency is unused, so it is not checked. + parakeet_scene_stream* nolat = parakeet_capi_scene_stream_begin(nullptr, nullptr, tag, &bad); + CHECK(nolat != nullptr, "latency checked without a diar ctx"); + parakeet_capi_scene_stream_free(nolat); + + parakeet_capi_free(asr); parakeet_capi_free(diar); parakeet_capi_free(tag); + if (fails) return 1; + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/tests/test_server_format.cpp b/tests/test_server_format.cpp index 565c6bc..cd12212 100644 --- a/tests/test_server_format.cpp +++ b/tests/test_server_format.cpp @@ -63,6 +63,17 @@ int main() { check(contains(e, "\"message\":\"bad\""), "error message"); check(contains(e, "\"type\":\"invalid_request_error\""), "error type"); + { + std::vector se = {{"Knock", 9.0, 10.0, 0.81}}; + Response v = format_transcription(tr, Format::kVerboseJson, 12.0, false, &se); + check(contains(v.body, "\"sound_events\":[{\"label\":\"Knock\",\"start\":9.000,\"end\":10.000,\"score\":0.8100}]"), + "verbose_json sound_events"); + Response j = format_transcription(tr, Format::kJson, 12.0, false, &se); + check(!contains(j.body, "sound_events"), "json has no sound_events"); + Response none = format_transcription(tr, Format::kVerboseJson, 12.0, false); + check(!contains(none.body, "sound_events"), "no sounds -> no field"); + } + if (fails) { std::fprintf(stderr, "%d checks failed\n", fails); return 1; } std::printf("test_server_format: OK\n"); return 0; diff --git a/tests/test_sound_capi.cpp b/tests/test_sound_capi.cpp new file mode 100644 index 0000000..4e7ea65 --- /dev/null +++ b/tests/test_sound_capi.cpp @@ -0,0 +1,155 @@ +// Sound stream C-API on real audio: ced.cpp's public-domain demo clips +// (rooster, thunder, guitar) back to back, 6 s each. Needs +// PARAKEET_TEST_CED_GGUF (any CED GGUF). Exit 77 when unset. +#include "parakeet_capi.h" +#include "audio_io.hpp" +#include "ced_tagger.hpp" +#include "sound_stream.hpp" + +#include +#include +#include +#include +#include +#include +#include + +// Load a WAV as 16 kHz mono (pk::load_audio_16k_mono downmixes and resamples). +static bool load16k(const std::string& path, std::vector& x) { + pk::Audio a; + if (!pk::load_audio_16k_mono(path, a)) return false; + x = std::move(a.samples); + return true; +} + +static int fails = 0; +#define CHECK(c, ...) do { if (!(c)) { std::fprintf(stderr, "FAIL: " __VA_ARGS__); std::fprintf(stderr, "\n"); ++fails; } } while (0) + +int main() { + const char* gguf = std::getenv("PARAKEET_TEST_CED_GGUF"); + if (!gguf) { std::fprintf(stderr, "SKIP: PARAKEET_TEST_CED_GGUF unset\n"); return 77; } + if (!pk::CedTagger::available()) { std::fprintf(stderr, "SKIP: built without CED\n"); return 77; } + parakeet_ctx* tag = parakeet_capi_load(gguf); + if (!tag) { std::fprintf(stderr, "FAIL: load\n"); return 1; } + + // Build the clip. + const char* names[] = {"rooster", "thunder", "guitar"}; + std::vector pcm; + for (const char* n : names) { + std::vector x; + const std::string p = std::string(PK_CED_SOURCE_DIR) + "/benchmarks/demo/clips/" + n + ".wav"; + CHECK(load16k(p, x), "read %s", p.c_str()); + pcm.insert(pcm.end(), x.begin(), x.end()); + } + + // opts_validation + parakeet_sound_opts o; + parakeet_capi_sound_opts_default(&o); + CHECK(o.size == (int)sizeof(o) && o.hop_sec == 1.0f && o.window_sec == 3.0f, "defaults"); + parakeet_sound_opts bad = o; + bad.hop_sec = 5.0f; + CHECK(parakeet_capi_sound_stream_begin(tag, &bad) == nullptr, "bad opts accepted"); + CHECK(std::strstr(parakeet_capi_last_error(tag), "hop_sec") != nullptr, "bad opts message: %s", + parakeet_capi_last_error(tag)); + + // Stream it in 0.25 s pieces. + parakeet_sound_stream* s = parakeet_capi_sound_stream_begin(tag, nullptr); + CHECK(s != nullptr, "begin: %s", parakeet_capi_last_error(tag)); + std::vector segs; + std::string scores; + bool checked_active = false; + for (size_t i = 0; i < pcm.size(); i += 4000) { + const int n = (int)std::min(4000, pcm.size() - i); + parakeet_sound_segment* out = nullptr; + int nout = 0; + CHECK(parakeet_capi_sound_stream_feed(s, pcm.data() + i, n, i + 4000 >= pcm.size(), &out, &nout) == 0, + "feed: %s", parakeet_capi_last_error(tag)); + for (int k = 0; k < nout; ++k) segs.push_back(out[k]); + parakeet_capi_free_sound_segments(out); + char* j = parakeet_capi_sound_stream_drain_scores_json(s); + CHECK(j != nullptr, "drain"); + if (j) { scores += j; parakeet_capi_free_string(j); } + + // Around 3.5 s in (inside the rooster slot), check the still-open + // segments: whatever comes back must be internally consistent, even + // when nothing happens to be open at this exact instant. + if (!checked_active && i + n >= 3.5 * 16000) { + checked_active = true; + parakeet_sound_segment* act = nullptr; + int nact = 0; + CHECK(parakeet_capi_sound_stream_active(s, &act, &nact) == 0, + "active: %s", parakeet_capi_last_error(tag)); + const double t = (i + n) / 16000.0; + for (int k = 0; k < nact; ++k) { + CHECK(act[k].start <= act[k].end, "active start <= end"); + CHECK(std::fabs(act[k].end - t) < 1e-3, "active end == stream time (%.3f vs %.3f)", + act[k].end, t); + } + parakeet_capi_free_sound_segments(act); + } + } + for (const auto& g : segs) + std::fprintf(stderr, " %-24s %6.2f %6.2f %.3f\n", g.label, g.start, g.end, g.peak); + + // Each clip's main label appears, inside its own 6 s slot (one hop slack). + auto found = [&](const char* label, float lo, float hi) { + for (const auto& g : segs) + if (std::strcmp(g.label, label) == 0 && g.start >= lo - 1.0f && g.start < hi) return true; + return false; + }; + CHECK(found("Chicken, rooster", 0.0f, 6.0f), "rooster in [0, 6)"); + CHECK(found("Thunder", 6.0f, 12.0f) || found("Thunderstorm", 6.0f, 12.0f), "thunder in [6, 12)"); + CHECK(found("Guitar", 12.0f, 18.0f) || found("Acoustic guitar", 12.0f, 18.0f), "guitar in [12, 18)"); + CHECK(scores.find("\"tags\":[{\"index\":") != std::string::npos, "score json shape"); + + // Window scores are exactly CedTagger::scorer()'s numbers for the same + // samples (test_ced_parity checks those against the PyTorch baseline). + { + auto t = pk::CedTagger::load(gguf); + std::vector direct; + CHECK(t && t->scorer()(pcm.data(), 16000, direct), "direct score"); + pk::SoundOpts so; so.window_sec = 1.0f; so.hop_sec = 1.0f; so.top_k = 1; + pk::SoundStream ss(t->scorer(), t->n_classes(), so); + ss.feed(pcm.data(), 16000, false); + auto w = ss.drain_windows(); + const float best = *std::max_element(direct.begin(), direct.end()); + CHECK(w.size() == 1 && w[0].top[0].second == best, "window score != direct score"); + } + // The first window of a 1 s / 1 s stream spans [0, 1). + { + parakeet_sound_opts one = o; + one.window_sec = 1.0f; one.hop_sec = 1.0f; one.top_k = 1; + parakeet_sound_stream* s1 = parakeet_capi_sound_stream_begin(tag, &one); + parakeet_sound_segment* out = nullptr; int nout = 0; + parakeet_capi_sound_stream_feed(s1, pcm.data(), 16000, 0, &out, &nout); + parakeet_capi_free_sound_segments(out); + char* j = parakeet_capi_sound_stream_drain_scores_json(s1); + std::fprintf(stderr, "first window: %s\n", j ? j : "(null)"); + CHECK(j && std::strstr(j, "\"start\":0.000,\"end\":1.000") != nullptr, "first window span"); + parakeet_capi_free_string(j); + parakeet_capi_sound_stream_free(s1); + } + + // Finished stream refuses more audio. + parakeet_sound_segment* out = nullptr; int nout = 0; + CHECK(parakeet_capi_sound_stream_feed(s, pcm.data(), 10, 0, &out, &nout) != 0, "feed after is_last"); + parakeet_capi_sound_stream_free(s); + + // wrong_ctx_kind: an ASR context is rejected with a clear message. + if (const char* asr_gguf = std::getenv("PARAKEET_TEST_GGUF")) { + parakeet_ctx* asr = parakeet_capi_load(asr_gguf); + CHECK(parakeet_capi_sound_stream_begin(asr, nullptr) == nullptr, "ASR ctx accepted as tagger"); + CHECK(std::strstr(parakeet_capi_last_error(asr), "ASR model") != nullptr, "message: %s", + parakeet_capi_last_error(asr)); + CHECK(parakeet_capi_num_classes(asr) == -1, "num_classes on ASR ctx"); + CHECK(parakeet_capi_model_kind(asr) == PARAKEET_MODEL_KIND_ASR, "model_kind on ASR ctx"); + parakeet_capi_free(asr); + } + CHECK(parakeet_capi_num_classes(tag) == 527, "num_classes"); + CHECK(parakeet_capi_model_kind(tag) == PARAKEET_MODEL_KIND_SOUND, "model_kind on tagger ctx"); + CHECK(parakeet_capi_model_kind(nullptr) == PARAKEET_MODEL_KIND_NONE, "model_kind on NULL"); + parakeet_capi_free(tag); + if (fails) return 1; + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/tests/test_sound_stream.cpp b/tests/test_sound_stream.cpp new file mode 100644 index 0000000..cb48258 --- /dev/null +++ b/tests/test_sound_stream.cpp @@ -0,0 +1,239 @@ +// Unit test for pk::SoundStream with a fake scorer (no model). +#include "sound_stream.hpp" + +#include +#include +#include +#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) +static bool near(float a, float b) { return std::fabs(a - b) < 1e-3f; } + +// Class 0 score = fraction of samples == 1.0 (like CED's mean pooling). +static SoundScorer fake() { + return [](const float* pcm, int n, std::vector& p) { + int on = 0; + for (int i = 0; i < n; ++i) on += pcm[i] == 1.0f; + p = {n > 0 ? (float)on / n : 0.0f, 0.0f}; + return true; + }; +} +// The default SoundOpts::top_k is 5, but validate_sound_opts requires +// top_k <= n_classes, and every fake-scorer test below uses n_classes = 2 +// (fake() returns two scores), so SoundOpts{} would make the constructor +// throw std::invalid_argument. top_k only feeds drain_windows()'s top-k +// queue, never the open/close hysteresis, so capping it to n_classes +// changes no segment expectation in this file. +static SoundOpts opts2() { SoundOpts o; o.top_k = 2; return o; } +// `sec` seconds of audio with the "sound" (1.0) in [a, b) seconds. +static std::vector clip(float sec, float a, float b) { + std::vector x((size_t)(sec * 16000), 0.0f); + for (size_t i = (size_t)(a * 16000); i < (size_t)(b * 16000) && i < x.size(); ++i) x[i] = 1.0f; + return x; +} +static std::vector run_all(SoundStream& s, const std::vector& x, int piece) { + std::vector all; + for (size_t i = 0; i < x.size(); i += piece) { + const int n = (int)std::min(piece, x.size() - i); + const bool last = i + piece >= x.size(); + auto got = s.feed(x.data() + i, n, last); + all.insert(all.end(), got.begin(), got.end()); + } + return all; +} + +// Sound in [5, 8): windows (3 s, hop 1 s) ending at 7 score 2/3 -> open at 6. +// Window ending 10 is [7, 10) = 1/3, still >= off 0.3. Window ending 11 is +// [8, 11) = 0 -> close at 8 + 1 = 9. Expected segment [6, 9), peak 1.0. +static void test_single_sound() { + SoundStream s(fake(), 2, opts2()); + auto segs = run_all(s, clip(14.0f, 5.0f, 8.0f), 1600); // 0.1 s pieces + CHECK(segs.size() == 1); + if (segs.size() == 1) { + CHECK(segs[0].cls == 0); + CHECK(near(segs[0].start, 6.0f)); + CHECK(near(segs[0].end, 9.0f)); + CHECK(near(segs[0].peak, 1.0f)); + } +} + +// Piece size must not change the result. +static void test_piece_size_invariant() { + const auto x = clip(14.0f, 5.0f, 8.0f); + SoundStream a(fake(), 2, opts2()), b(fake(), 2, opts2()); + auto sa = run_all(a, x, 1600), sb = run_all(b, x, 16000 * 14); + CHECK(sa.size() == sb.size()); + for (size_t i = 0; i < sa.size() && i < sb.size(); ++i) + CHECK(near(sa[i].start, sb[i].start) && near(sa[i].end, sb[i].end)); +} + +// Hysteresis: a score that dips to 0.33 (above off 0.3) does not close. +static void test_hysteresis_holds() { + SoundStream s(fake(), 2, opts2()); + auto x = clip(20.0f, 5.0f, 8.0f); + for (size_t i = 9 * 16000; i < 10 * 16000; ++i) x[i] = 1.0f; // second burst [9, 10) + auto segs = run_all(s, x, 1600); + CHECK(segs.size() == 1); // one segment, not two +} + +// Sound still present at the end: is_last closes it at the stream end. +static void test_is_last_closes_open() { + SoundStream s(fake(), 2, opts2()); + auto segs = run_all(s, clip(8.0f, 5.0f, 8.0f), 1600); + CHECK(segs.size() == 1); + if (!segs.empty()) { CHECK(near(segs[0].start, 6.0f)); CHECK(near(segs[0].end, 8.0f)); } + CHECK(s.finished()); +} + +// Short sounds are dropped by min_duration. +static void test_min_duration() { + SoundOpts o = opts2(); o.min_duration_sec = 5.0f; + SoundStream s(fake(), 2, o); + CHECK(run_all(s, clip(14.0f, 5.0f, 8.0f), 1600).empty()); +} + +// Growing first windows: sound from 0 opens in the first window at 0. +static void test_growing_first_window() { + SoundStream s(fake(), 2, opts2()); + auto segs = run_all(s, clip(10.0f, 0.0f, 2.0f), 1600); + CHECK(segs.size() == 1); + if (!segs.empty()) CHECK(near(segs[0].start, 0.0f)); +} + +// Very short input and empty feeds: no segments, no throw. +static void test_short_and_empty() { + SoundStream s(fake(), 2, opts2()); + CHECK(s.feed(nullptr, 0, false).empty()); + auto x = clip(0.1f, 0.0f, 0.1f); + CHECK(s.feed(x.data(), (int)x.size(), true).empty()); + CHECK(s.finished()); +} + +// Score queue: one window per hop, top-k sorted, drained once. +static void test_drain_windows() { + SoundOpts o; o.top_k = 1; + SoundStream s(fake(), 2, o); + auto x = clip(4.0f, 0.0f, 4.0f); + s.feed(x.data(), (int)x.size(), false); + auto w = s.drain_windows(); + CHECK(w.size() == 4); // windows ending at 1, 2, 3, 4 s + if (w.size() == 4) { + CHECK(near(w[0].start, 0.0f) && near(w[0].end, 1.0f)); + CHECK(near(w[3].start, 1.0f) && near(w[3].end, 4.0f)); + CHECK(w[3].top.size() == 1 && w[3].top[0].first == 0 && near(w[3].top[0].second, 1.0f)); + } + CHECK(s.drain_windows().empty()); +} + +// open_segments reports a running class with end = time(). +static void test_open_segments() { + SoundStream s(fake(), 2, opts2()); + auto x = clip(7.0f, 5.0f, 7.0f); + s.feed(x.data(), (int)x.size(), false); + auto open = s.open_segments(); + CHECK(open.size() == 1); + if (!open.empty()) CHECK(near(open[0].end, 7.0f)); + CHECK(s.safe_until() <= 6.0 + 1e-6); // the open segment started at 6 +} + +static void test_opts_validation() { + SoundOpts o; + CHECK(validate_sound_opts(o, 527).empty()); + o = SoundOpts{}; o.hop_sec = 4.0f; CHECK(!validate_sound_opts(o, 527).empty()); + o = SoundOpts{}; o.off_threshold = 0.5f; CHECK(!validate_sound_opts(o, 527).empty()); + o = SoundOpts{}; o.hop_sec = 0.1f; CHECK(!validate_sound_opts(o, 527).empty()); + o = SoundOpts{}; o.window_sec = 12.0f; CHECK(!validate_sound_opts(o, 527).empty()); + o = SoundOpts{}; o.top_k = 600; CHECK(!validate_sound_opts(o, 527).empty()); + bool threw = false; + try { SoundStream bad(fake(), 2, o); } catch (const std::invalid_argument&) { threw = true; } + CHECK(threw); +} + +// A failing scorer surfaces as std::runtime_error. +static void test_scorer_failure() { + SoundStream s([](const float*, int, std::vector&) { return false; }, 2, opts2()); + auto x = clip(2.0f, 0.0f, 0.0f); + bool threw = false; + try { s.feed(x.data(), (int)x.size(), false); } catch (const std::runtime_error&) { threw = true; } + CHECK(threw); +} + +// The buffer trim must key off +// min(next_end_, samples_in_), not next_end_ alone, or an is_last tail +// window that ends between two hops reads before the buffer start. +// clip(8.5, 5.0, 8.5) fed in 0.1 s (1600-sample) pieces, is_last on the +// last piece: the last full hop window ends at 8 s (unscored tail 0.5 s +// < the next hop at 9 s), so at is_last the tail window is +// [8 - 3, 8.5) intersected with what has streamed = [5.5, 8.5) (0.16 s+ +// tail, scored). The sound (still on at 8.5) is present throughout, so it +// stays open through 6 s (opened by the window ending at 7 s, 2/3 on) and +// is closed by is_last at the stream end, 8.5 s. +static void test_is_last_tail_window_after_hop_gap() { + SoundStream s(fake(), 2, opts2()); + auto segs = run_all(s, clip(8.5f, 5.0f, 8.5f), 1600); // 0.1 s pieces + CHECK(segs.size() == 1); + if (segs.size() == 1) { + CHECK(segs[0].cls == 0); + CHECK(near(segs[0].start, 6.0f)); + CHECK(near(segs[0].end, 8.5f)); + } + CHECK(s.finished()); +} + +// safe_until() must be a true lower bound on every future segment's start, +// including across the is_last tail window (while streaming, +// safe_until() has to account for a tail window that could still open a +// class at max(0, samples_in_ - hop_n_), not just the last window actually +// scored). With the default 3 s window / 1 s hop, sound in [7.5, 8.5) only +// scores 0.5/3 = 0.167 in the regular window ending at 8 s, below +// on_threshold, so use a lower on_threshold (0.3) to open a class in the +// is_last tail window [5.5, 8.5) (score 1/3 = 0.333). Right before the +// final 0.1 s piece, samples_in_ = 8.4 s and scored_end_ = 8.0 s. A bound +// of scored_end_ alone would say 8.0 s, but that final feed() call returns +// a segment starting at 7.5 s (the tail window's newest hop, +// max(0, samples_in_ - hop_n_) at samples_in_ = 8.5 s), 0.5 s earlier. +// safe_until() says min(8.0, 8.4 - 1.0) = 7.4 s at that point, which +// 7.5 s honors. +static void test_safe_until_bounds_tail_open() { + SoundOpts o = opts2(); o.on_threshold = 0.3f; o.off_threshold = 0.2f; + SoundStream s(fake(), 2, o); + auto x = clip(8.5f, 7.5f, 8.5f); + const int piece = 1600; // 0.1 s + bool saw_segment = false; + for (size_t i = 0; i < x.size(); i += piece) { + const int n = (int)std::min(piece, x.size() - i); + const bool last = i + piece >= x.size(); + const double safe = s.safe_until(); + auto got = s.feed(x.data() + i, n, last); + for (const auto& seg : got) { + saw_segment = true; + CHECK(seg.start >= safe - 1e-6); + } + } + CHECK(saw_segment); // otherwise the invariant above is checked vacuously + CHECK(s.finished()); +} + +int main() { + test_single_sound(); + test_piece_size_invariant(); + test_hysteresis_holds(); + test_is_last_closes_open(); + test_min_duration(); + test_growing_first_window(); + test_short_and_empty(); + test_drain_windows(); + test_open_segments(); + test_opts_validation(); + test_scorer_failure(); + test_is_last_tail_window_after_hop_gap(); + test_safe_until_bounds_tail_open(); + if (failures) { std::fprintf(stderr, "%d failure(s)\n", failures); return 1; } + std::fprintf(stderr, "PASS\n"); + return 0; +} diff --git a/third_party/ced.cpp b/third_party/ced.cpp new file mode 160000 index 0000000..e1a3cfa --- /dev/null +++ b/third_party/ced.cpp @@ -0,0 +1 @@ +Subproject commit e1a3cfa365feeb429cfd0373abb689647552c1ed