Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/diarization_encoder.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
#include "ggml_graph.hpp"
#include "ggml.h"

#include <algorithm>
#include <cmath>
#include <stdexcept>
#include <string>
Expand Down
7 changes: 4 additions & 3 deletions src/tdt.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -373,11 +373,12 @@ std::vector<TdtBeamHypothesis> 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});
Expand Down
28 changes: 24 additions & 4 deletions tests/test_transcribe_0_6b.cpp
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
#include "parakeet.h"
#include "model.hpp"
#include <cstdio>
#include <cstdlib>
#include <exception>
#include <memory>
#include <stdexcept>
#include <string>

// North-star end-to-end TDT transcription test for the real
Expand All @@ -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.
Expand Down Expand Up @@ -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<pk::Model> 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());
Expand All @@ -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;
}