diff --git a/src/diarization_encoder.cpp b/src/diarization_encoder.cpp index 8339dd2..ded4439 100644 --- a/src/diarization_encoder.cpp +++ b/src/diarization_encoder.cpp @@ -4,6 +4,7 @@ #include "ggml_graph.hpp" #include "ggml.h" +#include #include #include #include diff --git a/src/tdt.cpp b/src/tdt.cpp index 31ebc0f..37a1f65 100644 --- a/src/tdt.cpp +++ b/src/tdt.cpp @@ -373,11 +373,12 @@ std::vector tdt_beam_search( const int duration = durations[pair.duration_idx]; BeamState child = best; child.hyp.score += pair.score; - if (duration == 0 && - !(child.hyp.score < best.hyp.score)) { + // A near-certain label's log-prob can round to no change in + // the running score, so only an increase is an error. + if (duration == 0 && child.hyp.score > best.hyp.score) { throw std::runtime_error( "tdt_beam_search: zero-duration expansion " - "did not reduce score"); + "increased score"); } child.hyp.tokens.push_back(TdtBeamToken{ (int32_t)pair.token, (int32_t)time_idx, (int32_t)duration}); diff --git a/tests/test_transcribe_0_6b.cpp b/tests/test_transcribe_0_6b.cpp index 97d91b6..2be69aa 100644 --- a/tests/test_transcribe_0_6b.cpp +++ b/tests/test_transcribe_0_6b.cpp @@ -1,7 +1,9 @@ -#include "parakeet.h" +#include "model.hpp" #include #include #include +#include +#include #include // North-star end-to-end TDT transcription test for the real @@ -11,7 +13,10 @@ // code must honour: d_model=1024 / 24 layers / 128 mels, FastConformer linears // and conv convolutions configured with bias=False, and a STACKED 2-layer // prediction LSTM (pred_rnn_layers=2). This test asserts the C++ TDT path -// reproduces NeMo's transcript of tests/fixtures/speech.wav word-for-word. +// reproduces NeMo's transcript of tests/fixtures/speech.wav word-for-word, +// and that TDT beam search (beam 4) completes on it with the same top +// hypothesis: this clip produces zero-duration labels whose log-prob rounds to +// no change in the float32 running score, which beam search used to reject. // // The model GGUF is a ~2.4GB download not present in CI, so the test skips // cleanly (exit 77) unless PARAKEET_TEST_GGUF_06B points to a converted GGUF. @@ -45,10 +50,16 @@ int main() { : std::string(kRefV2); std::string got; + std::string got_beam; try { - got = pk::transcribe(gguf, "tests/fixtures/speech.wav"); + std::unique_ptr model = pk::Model::load(gguf); + if (!model) throw std::runtime_error(std::string("failed to load ") + gguf); + got = model->transcribe_path("tests/fixtures/speech.wav"); + got_beam = model->transcribe_path_nbest("tests/fixtures/speech.wav", + /*beam_size=*/4, /*nbest=*/1) + .at(0).text; } catch (const std::exception& e) { - std::fprintf(stderr, "test_transcribe_0_6b: pk::transcribe threw: %s\n", e.what()); + std::fprintf(stderr, "test_transcribe_0_6b: transcribe threw: %s\n", e.what()); return 1; } std::fprintf(stderr, "test_transcribe_0_6b: got = %s\n", got.c_str()); @@ -63,6 +74,15 @@ int main() { return 1; } + if (got_beam != expected) { + std::fprintf(stderr, + "test_transcribe_0_6b: beam 4 MISMATCH\n" + " got: %s\n" + " expected: %s\n", + got_beam.c_str(), expected.c_str()); + return 1; + } + std::fprintf(stderr, "test_transcribe_0_6b: PASS (word-for-word match with NeMo TDT)\n"); return 0; }