From 1f7ebfd7d0ce08a27cb01ca0c021940303ddfeba Mon Sep 17 00:00:00 2001 From: CJ Pais Date: Mon, 21 Sep 2026 19:16:23 +0800 Subject: [PATCH 1/4] repetition loop --- docs/environment-variables.md | 1 + docs/input-limits.md | 14 +++ include/transcribe.h | 4 +- src/arch/canary/model.cpp | 32 +++--- src/arch/canary_qwen/model.cpp | 29 +++-- src/arch/cohere/model.cpp | 26 ++--- src/arch/funasr_nano/model.cpp | 29 +++-- src/arch/granite/model.cpp | 28 ++--- src/arch/moonshine/model.cpp | 40 +++---- src/arch/moonshine_streaming/model.cpp | 34 +++--- src/arch/moss/model.cpp | 35 ++---- src/arch/qwen3_asr/model.cpp | 29 +++-- src/arch/voxtral/model.cpp | 29 +++-- src/causal_lm/causal_lm.cpp | 8 +- src/causal_lm/causal_lm.h | 3 +- src/transcribe-batch-util.cpp | 53 ++++++++- src/transcribe-batch-util.h | 25 ++++- src/transcribe-repetition-guard.h | 82 ++++++++++++++ src/transcribe.cpp | 34 +----- tests/CMakeLists.txt | 15 +++ tests/repetition_guard_unit.cpp | 144 +++++++++++++++++++++++++ tests/run_dispatch_unit.cpp | 98 +++++++++++++++++ 22 files changed, 572 insertions(+), 220 deletions(-) create mode 100644 src/transcribe-repetition-guard.h create mode 100644 tests/repetition_guard_unit.cpp diff --git a/docs/environment-variables.md b/docs/environment-variables.md index 5d6927fa..1beb8e73 100644 --- a/docs/environment-variables.md +++ b/docs/environment-variables.md @@ -28,6 +28,7 @@ of tests. | `TRANSCRIBE_FORCE_FLASH` | Force flash attention on. Wins over `TRANSCRIBE_NO_FLASH` if both are set. | | `TRANSCRIBE_CONV_DIRECT_DW` / `TRANSCRIBE_CONV_NO_DIRECT_DW` | Force the depthwise-conv dispatch to the direct `conv_2d_dw` path / the im2col path, overriding the per-family backend default. | | `TRANSCRIBE_CONV_DIRECT_PW` / `TRANSCRIBE_CONV_NO_DIRECT_PW` | Force the pointwise-conv dispatch to direct `mul_mat` / im2col, overriding the backend default. | +| `TRANSCRIBE_NO_REPETITION_GUARD` | Turn off the stop for greedy decodes that start repeating themselves (see [`input-limits.md`](input-limits.md)). For byte-exact reference parity; a looping decode then runs to its budget. Read once per process. | | `TRANSCRIBE_DUMP_DIR=` | Enable the per-stage tensor dumper; writes `.f32` + `.json` per dumped tensor into ``. The basis for the numerical-comparison harness (`scripts/compare_tensors.py`). | | `TRANSCRIBE_PERF_DEBUG` | Print a per-stage timing breakdown to stderr (DEBUG log) on the families that profile (`cohere`, `granite`, `canary`, `canary_qwen`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `whisper`). For whisper, a value containing `cpu` or `all` additionally prints the CPU sub-section breakdown. | | `TRANSCRIBE_VOXTRAL_REALTIME_STREAM_TIMING` | Print a per-component streaming wall-time breakdown at stream finalize (voxtral_realtime). | diff --git a/docs/input-limits.md b/docs/input-limits.md index e53123f8..feced3a3 100644 --- a/docs/input-limits.md +++ b/docs/input-limits.md @@ -94,6 +94,19 @@ more likely. For `canary` and `cohere`, input and output have separate encoder and decoder limits; `max_audio_ms` reports the encoder limit, not a recommended chunk size. +A greedy decode can also fall into repeating one phrase until the budget runs +out. The greedy families (`canary`, `canary_qwen`, `cohere`, `funasr_nano`, +`granite`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `voxtral`) +stop as soon as a block of up to 64 tokens has repeated at least 4 times and +the copies cover at least 32 tokens. The repeats are dropped, leaving one copy. +Whatever the audio said after the loop was never decoded, so the run reports +the same `OUTPUT_TRUNCATED` status and `WARN` as a budget stop. `whisper` is +excluded because it recovers from loops with its own temperature fallback. +`voxtral_realtime` is excluded because it emits one token per audio frame, so +its padding tokens repeat through any silence. Set +`TRANSCRIBE_NO_REPETITION_GUARD=1` to turn the guard off, e.g. for +byte-exact reference parity. + ### 3. Soft window — warn and proceed | Families | Window | Behavior | @@ -165,6 +178,7 @@ with `TRANSCRIBE_ERR_INPUT_TOO_LONG` (one-shot and batch) or surfaced via | Input within limit and decode completes | `TRANSCRIBE_OK` | — | full transcript | | Over-length, hard-cap family | `TRANSCRIBE_ERR_INPUT_TOO_LONG` | `ERROR` via callback | no transcript (rejected before the decode) | | Generation ran long mid-decode | `TRANSCRIBE_ERR_OUTPUT_TRUNCATED` | `WARN` via callback | partial transcript readable; `transcribe_was_truncated() == true` | +| Greedy decode started repeating | `TRANSCRIBE_ERR_OUTPUT_TRUNCATED` | `WARN` via callback | partial transcript readable, repeats dropped; `transcribe_was_truncated() == true` | | Over-window, soft-window family | `TRANSCRIBE_OK` | `WARN` via callback | full transcript (accuracy may be degraded) | | Chunked / unbounded family | `TRANSCRIBE_OK` | — | full transcript | | Cache/graph allocation failed | `TRANSCRIBE_ERR_OOM` | `ERROR` via callback | no transcript (no silent context shrink) | diff --git a/include/transcribe.h b/include/transcribe.h index 6d74a75f..552c1c7a 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -269,7 +269,9 @@ typedef enum { /* * Returned by transcribe_run when the decode stopped because it hit * the model's context / generation budget BEFORE the model emitted - * end-of-stream — i.e. the transcript is incomplete. This is the + * end-of-stream — i.e. the transcript is incomplete. A greedy decode + * that falls into repeating itself is also stopped early and reported + * here, with the repeats dropped from the partial. This is the * "started, couldn't finish" counterpart to INPUT_TOO_LONG, and it is * a hard non-OK status by design: a truncated transcript must not be * mistaken for a complete one. diff --git a/src/arch/canary/model.cpp b/src/arch/canary/model.cpp index 775c1995..e824a35e 100644 --- a/src/arch/canary/model.cpp +++ b/src/arch/canary/model.cpp @@ -22,6 +22,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -1194,6 +1195,7 @@ transcribe_status run(transcribe_session * session, const bool primary_is_gpu = cm->plan.primary_kind != transcribe::BackendKind::Cpu && cm->plan.primary_kind != transcribe::BackendKind::Accel && cm->plan.primary_kind != transcribe::BackendKind::Unknown; + bool repeating = false; if (primary_is_gpu) { // Static-graph step path (GPU). max_n_kv: pad to next power of two @@ -1277,6 +1279,11 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "canary run")) { + cc->was_truncated = true; + repeating = true; + break; + } } } } else { @@ -1339,6 +1346,11 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "canary run")) { + cc->was_truncated = true; + repeating = true; + break; + } } } } @@ -1347,7 +1359,8 @@ transcribe_status run(transcribe_session * session, // KV-full / a compute break without end-of-stream: flag truncation and // WARN rather than silently shortening. (Abort paths return early and // intentionally do NOT set the flag — abort is not a length truncation.) - if (next_token != eos_id) { + // A repetition stop has already flagged and logged itself. + if (!repeating && next_token != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "canary run: output truncated at %d tokens — decode reached the " @@ -1476,21 +1489,8 @@ transcribe_status run_batch_serial(CanarySession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, diff --git a/src/arch/canary_qwen/model.cpp b/src/arch/canary_qwen/model.cpp index 386fd9da..fdd61d5f 100644 --- a/src/arch/canary_qwen/model.cpp +++ b/src/arch/canary_qwen/model.cpp @@ -36,6 +36,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -1154,6 +1155,7 @@ transcribe_status run(transcribe_session * context, per_step_compute_us.reserve(64); } + bool repeating = false; while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new && cur_past + 1 <= max_n_kv) { const int64_t t_in0 = profile_decode ? ggml_time_us() : 0; ggml_backend_tensor_set(sb.input_id_in, &next_tok, 0, sizeof(int32_t)); @@ -1198,11 +1200,17 @@ transcribe_status run(transcribe_session * context, cur_past += 1; n_steps += 1; + if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "canary_qwen run")) { + cc->was_truncated = true; + repeating = true; + break; + } } // Decode stopped at EOS (complete) or the generation budget / context width - // (truncated). Surface the latter via transcribe_was_truncated() + WARN. - if (next_tok != eos_id) { + // (truncated). Surface the latter via transcribe_was_truncated() + WARN; a + // repetition stop has already done both. + if (!repeating && next_tok != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "canary_qwen run: output truncated at %d tokens — decode reached " @@ -1353,21 +1361,8 @@ transcribe_status run_batch_serial(CanaryQwenSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } } // namespace diff --git a/src/arch/cohere/model.cpp b/src/arch/cohere/model.cpp index bdf8050e..92ec38fb 100644 --- a/src/arch/cohere/model.cpp +++ b/src/arch/cohere/model.cpp @@ -23,6 +23,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -1217,6 +1218,10 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "cohere run")) { + cc->was_truncated = true; + break; + } } } } else { @@ -1288,6 +1293,10 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "cohere run")) { + cc->was_truncated = true; + break; + } } } } @@ -1436,21 +1445,8 @@ transcribe_status run_batch_serial(CohereSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, diff --git a/src/arch/funasr_nano/model.cpp b/src/arch/funasr_nano/model.cpp index 751d6ede..fd6628bc 100644 --- a/src/arch/funasr_nano/model.cpp +++ b/src/arch/funasr_nano/model.cpp @@ -20,6 +20,7 @@ #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -866,6 +867,7 @@ transcribe_status run(transcribe_session * session, // (prefill = 1st call, iter K = (K+2)th call), so dump when n_steps == 7. const int gen_dump_step = 7; int n_steps = 0; + bool repeating = false; while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new && cur_past + 1 <= max_n_kv) { ggml_backend_tensor_set(sb.input_id_in, &next_tok, 0, sizeof(int32_t)); const int32_t pos_val = cur_past; @@ -897,14 +899,20 @@ transcribe_status run(transcribe_session * session, cur_past += 1; n_steps += 1; + if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "funasr_nano run")) { + cc->was_truncated = true; + repeating = true; + break; + } } (void) n_steps; // The decode stopped at EOS (complete) or at the generation budget / // context width (truncated). Surface the latter via // transcribe_was_truncated() and a WARN rather than returning a silently - // shortened transcript. See docs/input-limits.md. - if (next_tok != eos_id) { + // shortened transcript; a repetition stop has already done both. See + // docs/input-limits.md. + if (!repeating && next_tok != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "funasr_nano run: output truncated at %d tokens — decode reached " @@ -1042,21 +1050,8 @@ transcribe_status run_batch_serial(FunAsrNanoSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } } // namespace diff --git a/src/arch/granite/model.cpp b/src/arch/granite/model.cpp index 959c314a..9ebd549d 100644 --- a/src/arch/granite/model.cpp +++ b/src/arch/granite/model.cpp @@ -19,6 +19,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -1234,11 +1235,17 @@ transcribe_status run(transcribe_session * ctx_base, // valid positions get zeroed per step. std::vector step_mask(max_n_kv, 0xFC00); + bool repeating = false; for (int step_i = 0; step_i < max_steps; ++step_i) { if (next_id == eos_id) { break; } gen_ids.push_back(next_id); + if (transcribe::stop_on_repetition(gen_ids, "granite run")) { + cc->was_truncated = true; + repeating = true; + break; + } const int32_t pos = T_prompt + step_i; // RoPE position const int64_t kv_idx = pos; // KV write row @@ -1274,8 +1281,8 @@ transcribe_status run(transcribe_session * ctx_base, // The decode stopped either at EOS (complete) or at the generation // budget / context ceiling (truncated). Surface the latter via // transcribe_was_truncated() and a WARN rather than handing back a - // silently shortened transcript. - if (next_id != eos_id) { + // silently shortened transcript; a repetition stop has already done both. + if (!repeating && next_id != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "granite run: output truncated at %d tokens — decode reached the " @@ -1450,21 +1457,8 @@ transcribe_status run_batch_serial(GraniteSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } } // namespace diff --git a/src/arch/moonshine/model.cpp b/src/arch/moonshine/model.cpp index 9f3de4a4..8d6ffc31 100644 --- a/src/arch/moonshine/model.cpp +++ b/src/arch/moonshine/model.cpp @@ -21,6 +21,7 @@ #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -630,6 +631,7 @@ transcribe_status run(transcribe_session * session, cm->plan.primary_kind != transcribe::BackendKind::Accel && cm->plan.primary_kind != transcribe::BackendKind::Unknown; const bool use_step_graph = primary_is_gpu && !transcribe::debug::enabled(); + bool repeating = false; if (use_step_graph) { // ---------- Static-graph step path (GPU) ---------- @@ -708,6 +710,11 @@ transcribe_status run(transcribe_session * session, if (next_token != eos) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "moonshine run")) { + cc->was_truncated = true; + repeating = true; + break; + } } } } else { @@ -733,14 +740,19 @@ transcribe_status run(transcribe_session * session, } if (next_token != eos) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "moonshine run")) { + cc->was_truncated = true; + repeating = true; + break; + } } } } // A non-eos last token means the decode hit the position cap before // end-of-stream (see the input-length contract above): flag truncation - // and WARN. - if (next_token != eos) { + // and WARN. A repetition stop has already done both. + if (!repeating && next_token != eos) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "moonshine run: output truncated at %d tokens — decode reached the " @@ -856,21 +868,8 @@ transcribe_status run_batch_serial(MoonshineSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, @@ -1056,7 +1055,8 @@ transcribe_status run_batch(transcribe_session * session, const int64_t dec_us = ggml_time_us() - t_dec0; // Batched truncation: the shared step loop marks each valid row that hit - // the output cap before end-of-stream. Mirror the serial path (WARN + flag). + // the output cap, or started repeating, before end-of-stream. Mirror the + // serial path (WARN + flag). { int n_truncated = 0; for (int b = 0; b < n; ++b) { @@ -1068,8 +1068,8 @@ transcribe_status run_batch(transcribe_session * session, cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "moonshine run_batch: %d of %d utterances truncated — decode " - "reached the position cap (%d) before end-of-stream; those " - "transcripts may be incomplete. This model is intended for " + "reached the position cap (%d) or began repeating before " + "end-of-stream; those transcripts may be incomplete. This model is intended for " "short utterances. See transcribe_capabilities.max_audio_ms.", n_truncated, n, max_pos); } diff --git a/src/arch/moonshine_streaming/model.cpp b/src/arch/moonshine_streaming/model.cpp index 524dd31a..48d1edfe 100644 --- a/src/arch/moonshine_streaming/model.cpp +++ b/src/arch/moonshine_streaming/model.cpp @@ -36,6 +36,7 @@ #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "transcribe/moonshine_streaming.h" #include "weights.h" @@ -977,6 +978,7 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, // dec.logits_raw.gen20 dumps the logits that predict the 20th // emitted token (n_past == 20 at that step). Matches moonshine. constexpr int k_mid_gen_step = 20; + bool repeating = false; while (next_token != eos && n_past < gen_cap) { if (cc->poll_abort()) { return TRANSCRIBE_ERR_ABORTED; @@ -994,12 +996,17 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, } if (next_token != eos) { generated_ids.push_back(next_token); + if (transcribe::stop_on_repetition(generated_ids, "moonshine_streaming run")) { + cc->was_truncated = true; + repeating = true; + break; + } } } // Non-EOS after the loop means gen_cap stopped decode before EOS. gen_cap // is either the position cap or the tighter duration budget. - if (next_token != eos) { + if (!repeating && next_token != eos) { cc->was_truncated = true; const bool hit_duration_budget = (gen_cap < max_pos) || (max_pos <= 0); transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, @@ -1801,21 +1808,8 @@ transcribe_status run_batch_serial(MoonshineStreamingSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, @@ -2011,6 +2005,7 @@ transcribe_status run_batch(transcribe_session * session, std::vector tok_buf(n, 0), pos_buf(n, 0), argmax_buf(n, 0); std::vector kvidx_buf(n, 0); std::vector finished(n, 0); + std::vector repeating(n, 0); std::vector> generated(n); std::vector next_tok(n, 0); for (int b = 0; b < n; ++b) { @@ -2112,6 +2107,11 @@ transcribe_status run_batch(transcribe_session * session, finished[b] = 1; } else { generated[b].push_back(next_tok[b]); + if (transcribe::stop_on_repetition(generated[b], "moonshine_streaming run_batch")) { + cc->was_truncated = true; + finished[b] = 1; + repeating[b] = 1; + } } } } @@ -2160,7 +2160,7 @@ transcribe_status run_batch(transcribe_session * session, rs.result_kind = TRANSCRIBE_TIMESTAMPS_NONE; rs.has_result = true; // Per-utterance truncation parity (offline run_batch, not streaming). - rs.status = !finished[b] ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + rs.status = (!finished[b] || repeating[b]) ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; rs.t_mel_us = 0; rs.t_encode_us = enc_us / valid_count; rs.t_decode_us = dec_us / valid_count; diff --git a/src/arch/moss/model.cpp b/src/arch/moss/model.cpp index 5133a957..605ae14b 100644 --- a/src/arch/moss/model.cpp +++ b/src/arch/moss/model.cpp @@ -24,6 +24,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -920,6 +921,7 @@ transcribe_status run(transcribe_session * session, per_step_us.reserve(512); } + bool repeating = false; while (next_tok != eos_id && static_cast(generated_ids.size()) < gen_budget && cur_past + 1 <= max_n_kv) { const int64_t t_i0 = perf_debug ? ggml_time_us() : 0; ggml_backend_tensor_set(sb.input_id_in, &next_tok, 0, sizeof(int32_t)); @@ -968,12 +970,17 @@ transcribe_status run(transcribe_session * session, cur_past += 1; cc->kv_cache.n = cur_past + 1; cc->kv_cache.head = cur_past + 1; + if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "moss run")) { + cc->was_truncated = true; + repeating = true; + break; + } if (cc->poll_abort()) { return TRANSCRIBE_ERR_ABORTED; } } - if (next_tok != eos_id) { + if (!repeating && next_tok != eos_id) { cc->was_truncated = true; log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "moss run: output truncated at %d tokens", static_cast(generated_ids.size())); @@ -1052,30 +1059,8 @@ transcribe_status run_batch_serial(MossSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - bool any_truncated = false; - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - cc->clear_result(); - cc->t_mel_us = 0; - cc->t_encode_us = 0; - cc->t_decode_us = 0; - cc->was_truncated = false; - - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - any_truncated = any_truncated || st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED; - if (st == TRANSCRIBE_OK || cc->has_result) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - cc->was_truncated = any_truncated; - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, diff --git a/src/arch/qwen3_asr/model.cpp b/src/arch/qwen3_asr/model.cpp index 250bef3e..fb1a02b3 100644 --- a/src/arch/qwen3_asr/model.cpp +++ b/src/arch/qwen3_asr/model.cpp @@ -18,6 +18,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "weights.h" #include @@ -931,6 +932,7 @@ transcribe_status run(transcribe_session * session, int64_t t_step_comp_us = 0; int64_t t_step_get_us = 0; const int64_t t_step_loop_start = ggml_time_us(); + bool repeating = false; while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new && cur_past + 1 <= max_n_kv) { const int64_t t_set0 = ggml_time_us(); @@ -968,13 +970,19 @@ transcribe_status run(transcribe_session * session, cc->kv_cache.n = cur_past + 1; cc->kv_cache.head = cur_past + 1; t_step_get_us += ggml_time_us() - t_comp1; + if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "qwen3_asr run")) { + cc->was_truncated = true; + repeating = true; + break; + } } t_step_loop_us = ggml_time_us() - t_step_loop_start; n_steps = static_cast(generated_ids.size()) - 1; // Decode stopped at EOS (complete) or the generation budget / context width - // (truncated). Surface the latter via transcribe_was_truncated() + WARN. - if (next_tok != eos_id) { + // (truncated). Surface the latter via transcribe_was_truncated() + WARN; a + // repetition stop has already done both. + if (!repeating && next_tok != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "qwen3_asr run: output truncated at %d tokens — decode reached the " @@ -1470,21 +1478,8 @@ transcribe_status run_batch_serial(QwenAsrSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } } // namespace diff --git a/src/arch/voxtral/model.cpp b/src/arch/voxtral/model.cpp index aba6fa73..9739df9f 100644 --- a/src/arch/voxtral/model.cpp +++ b/src/arch/voxtral/model.cpp @@ -26,6 +26,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-repetition-guard.h" #include "voxtral.h" #include "weights.h" @@ -945,6 +946,7 @@ transcribe_status run(transcribe_session * session, const ggml_fp16_t mz = ggml_fp32_to_fp16(0.0f); const ggml_fp16_t mn = ggml_fp32_to_fp16(-INFINITY); std::vector step_mask(max_n_kv, mn); + bool repeating = false; while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new && cur_past + 1 <= max_n_kv) { if (cc->poll_abort()) { return TRANSCRIBE_ERR_ABORTED; @@ -978,12 +980,18 @@ transcribe_status run(transcribe_session * session, cur_past += 1; cc->kv_cache.n = cur_past + 1; cc->kv_cache.head = cur_past + 1; + if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "voxtral run")) { + cc->was_truncated = true; + repeating = true; + break; + } } cc->t_decode_us = ggml_time_us() - t_dec_start; // Decode stopped at EOS (complete) or at the generation budget / context - // width (truncated). Surface the latter via transcribe_was_truncated() + WARN. - if (next_tok != eos_id) { + // width (truncated). Surface the latter via transcribe_was_truncated() + WARN; + // a repetition stop has already done both. + if (!repeating && next_tok != eos_id) { cc->was_truncated = true; transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "voxtral run: output truncated at %d tokens — decode reached the " @@ -1037,21 +1045,8 @@ transcribe_status run_batch_serial(VoxtralSession * cc, const int * n_samples, int n, const transcribe_run_params * params) { - for (int i = 0; i < n; ++i) { - if (cc->poll_abort()) { - return TRANSCRIBE_ERR_ABORTED; - } - const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : - run(cc, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - cc->batch_results.push_back(cc->capture_result(st)); - } else { - transcribe_session::ResultSet rs; - rs.status = st; - cc->batch_results.push_back(std::move(rs)); - } - } - return TRANSCRIBE_OK; + return transcribe::run_batch_serial(cc, pcm, n_samples, n, + [&](const float * p, int ns) { return run(cc, p, ns, params); }); } transcribe_status run_batch(transcribe_session * session, diff --git a/src/causal_lm/causal_lm.cpp b/src/causal_lm/causal_lm.cpp index f8ecb2da..001d9e58 100644 --- a/src/causal_lm/causal_lm.cpp +++ b/src/causal_lm/causal_lm.cpp @@ -8,6 +8,7 @@ #include "transcribe-backend.h" #include "transcribe-env.h" #include "transcribe-log.h" +#include "transcribe-repetition-guard.h" #include "transcribe-session.h" #include @@ -963,7 +964,8 @@ transcribe_status run_batched_step_loop(transcribe_session * sess if (n_past[b] < max_n_kv) { mask_buf[base + n_past[b]] = mz; } - if (tok == eos_id || static_cast(generated[b].size()) >= max_new || n_past[b] + 1 > max_n_kv) { + if (tok == eos_id || stop_on_repetition(generated[b], "batched decode") || + static_cast(generated[b].size()) >= max_new || n_past[b] + 1 > max_n_kv) { finished[b] = 1; } else { all_done = false; @@ -977,8 +979,8 @@ transcribe_status run_batched_step_loop(transcribe_session * sess } // A valid row was truncated if it stopped for a reason OTHER than eos - // (generation budget or KV window). `finished` is set on every stop - // reason, so it can't discriminate; the signal is the last sampled token: + // (generation budget, KV window, repetition guard). `finished` is set on + // every stop reason, so it can't discriminate; the signal is the last sampled token: // `next_tok[b] != eos_id` means the row was cut off mid-transcript (it is // frozen once the row finishes). See docs/input-limits.md. if (truncated_out != nullptr) { diff --git a/src/causal_lm/causal_lm.h b/src/causal_lm/causal_lm.h index 46999d84..b5f477c8 100644 --- a/src/causal_lm/causal_lm.h +++ b/src/causal_lm/causal_lm.h @@ -333,7 +333,8 @@ struct StepLoopStats { }; // Run the lockstep batched greedy decode. Each row steps until it emits -// `eos_id`, accumulates `max_new` generated tokens, or fills the KV window; +// `eos_id`, starts repeating (transcribe-repetition-guard.h), accumulates +// `max_new` generated tokens, or fills the KV window; // each emitted token is appended to generated[b]. Finished / invalid rows keep // stepping into their own KV slab (a no-op for live rows). Polls // session->poll_abort() once per step. The step graph must already be built diff --git a/src/transcribe-batch-util.cpp b/src/transcribe-batch-util.cpp index 11372a9a..3901e61b 100644 --- a/src/transcribe-batch-util.cpp +++ b/src/transcribe-batch-util.cpp @@ -5,6 +5,7 @@ #include "ggml-backend.h" #include "ggml.h" #include "transcribe-log.h" +#include "transcribe-repetition-guard.h" #include "transcribe-session.h" #include @@ -197,6 +198,44 @@ transcribe_status decode_batch_slices(transcribe_session * session, return TRANSCRIBE_OK; } +transcribe_status run_batch_serial(transcribe_session * session, + const float * const * pcm, + const int * n_samples, + int n, + const RunOneFn & run_one) { + bool any_truncated = false; + for (int i = 0; i < n; ++i) { + if (session->poll_abort()) { + session->was_truncated = any_truncated; + return TRANSCRIBE_ERR_ABORTED; + } + session->clear_result(); + session->t_mel_us = 0; + session->t_encode_us = 0; + session->t_decode_us = 0; + session->was_truncated = false; + + const transcribe_status st = + (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : run_one(pcm[i], n_samples[i]); + any_truncated = any_truncated || st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + // The slot was cleared above, so has_result means this utterance wrote + // it: keep partials (truncated, aborted), never a stale snapshot. + if (st == TRANSCRIBE_OK || session->has_result) { + session->batch_results.push_back(session->capture_result(st)); + } else { + transcribe_session::ResultSet rs; + rs.status = st; + session->batch_results.push_back(std::move(rs)); + } + if (st == TRANSCRIBE_ERR_ABORTED) { + session->was_truncated = any_truncated; + return TRANSCRIBE_ERR_ABORTED; + } + } + session->was_truncated = any_truncated; + return TRANSCRIBE_OK; +} + transcribe_status run_batched_encdec_step_loop(transcribe_session * session, ggml_backend_sched_t sched, const EncDecRebuildFn & rebuild, @@ -225,6 +264,7 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * std::vector tok_buf(n, 0), pos_buf(n, 0), argmax_buf(n, 0); std::vector kvidx_buf(n, 0); std::vector finished(n, 0); + std::vector repeating(n, 0); std::vector next_tok(n, 0); for (int b = 0; b < n; ++b) { if (!valid[b]) { @@ -344,6 +384,10 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * finished[b] = 1; } else { generated[b].push_back(next_tok[b]); + if (stop_on_repetition(generated[b], "batched decode")) { + finished[b] = 1; + repeating[b] = 1; + } } } } @@ -352,13 +396,14 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * *n_steps_out = n_steps; } - // A valid row that never reached eos was cut off at the generation budget - // or the context window — report it as truncated so the family can return - // per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED. See docs/input-limits.md. + // A valid row that never reached eos was cut off at the generation budget, + // the context window, or the repetition guard — report it as truncated so + // the family can return per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED. See + // docs/input-limits.md. if (truncated_out != nullptr) { truncated_out->assign(n, 0); for (int b = 0; b < n; ++b) { - (*truncated_out)[b] = (valid[b] && !finished[b]) ? 1 : 0; + (*truncated_out)[b] = (valid[b] && (!finished[b] || repeating[b])) ? 1 : 0; } } return TRANSCRIBE_OK; diff --git a/src/transcribe-batch-util.h b/src/transcribe-batch-util.h index 9be59eed..ad15961f 100644 --- a/src/transcribe-batch-util.h +++ b/src/transcribe-batch-util.h @@ -102,6 +102,21 @@ transcribe_status decode_batch_slices(transcribe_session * session, int64_t total_mel_us, const std::function & decode_fn); +// Serial run_batch fallback: runs each utterance through `run_one` (a family's +// single-utterance run()) and snapshots it into session->batch_results. The +// per-run state transcribe_run resets (result slot, timings, truncation flag) +// is reset before every utterance, so one truncated utterance cannot mark the +// rest. A truncated or aborted utterance keeps its partial result, as it does +// from transcribe_run; a null pcm or n_samples <= 0 is recorded as +// INVALID_ARG. session->was_truncated ends true if any utterance truncated. +// Returns TRANSCRIBE_ERR_ABORTED once an utterance aborts, else OK. +using RunOneFn = std::function; +transcribe_status run_batch_serial(transcribe_session * session, + const float * const * pcm, + const int * n_samples, + int n, + const RunOneFn & run_one); + // --------------------------------------------------------------------------- // Batched encoder-decoder greedy step loop (cohere / canary / moonshine) // @@ -133,8 +148,9 @@ struct EncDecStepIO { using EncDecRebuildFn = std::function; // Run the shared greedy enc-dec step loop. Feeds `prompt_ids[0..prompt_len)` as -// uniform lockstep tokens, then generates until each row emits eos_id, the batch -// reaches `max_new` produced tokens, or the position fills `max_n_kv`. Manages +// uniform lockstep tokens, then generates until each row emits eos_id or starts +// repeating (transcribe-repetition-guard.h), the batch reaches `max_new` +// produced tokens, or the position fills `max_n_kv`. Manages // the self-attention key mask and dynamic window growth (via `rebuild`), and // appends generated tokens to generated[b] (invalid rows are skipped, finished // rows keep stepping into their own KV slab). Polls session->poll_abort() each @@ -142,8 +158,9 @@ using EncDecRebuildFn = std::function; // *n_steps_out (if non-null) receives the number of compute steps run. // // truncated_out (if non-null) is sized to n_batch and set per row: 1 when that -// (valid) row hit the generation budget (max_new) or the context window -// (max_n_kv) BEFORE emitting eos_id (transcript truncated), else 0. Lets a +// (valid) row hit the generation budget (max_new), the context window +// (max_n_kv) or the repetition guard BEFORE emitting eos_id (transcript +// truncated), else 0. Lets a // family report per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED from run_batch. transcribe_status run_batched_encdec_step_loop(transcribe_session * session, ggml_backend_sched_t sched, diff --git a/src/transcribe-repetition-guard.h b/src/transcribe-repetition-guard.h new file mode 100644 index 00000000..498213c1 --- /dev/null +++ b/src/transcribe-repetition-guard.h @@ -0,0 +1,82 @@ +// Stop rule for greedy autoregressive decoders stuck repeating themselves. +// +// Greedy argmax has no way out of a self-reinforcing state: once the likeliest +// continuation of a phrase is the phrase itself, it repeats until the decode +// budget runs out. The guard works on token ids only, so it fits decoders that +// read back a device-side argmax, and leaves output that never loops unchanged. + +#pragma once + +#include "transcribe-env.h" +#include "transcribe-log.h" + +#include +#include +#include + +namespace transcribe { + +// A block repeating at the tail is a loop once it has k_repeat_min_copies +// copies covering at least k_repeat_min_tokens. Short blocks need many more +// copies (a 1-token block needs 32), so emphatic human repetition survives. +constexpr int k_repeat_max_block = 64; +constexpr int k_repeat_min_copies = 4; +constexpr int k_repeat_min_tokens = 32; + +// Length of the block repeating at the end of ids[0, n), or 0. +inline int repeating_tail_block(const int32_t * ids, int n) { + for (int block = 1; block <= k_repeat_max_block; ++block) { + const int copies = std::max(k_repeat_min_copies, (k_repeat_min_tokens + block - 1) / block); + const int span = block * copies; + if (span > n) { + continue; + } + bool periodic = true; + for (int i = n - 1; i >= n - span + block; --i) { + if (ids[i] != ids[i - block]) { + periodic = false; + break; + } + } + if (periodic) { + return block; + } + } + return 0; +} + +// Length of ids[0, n) with every copy of the tail block but the first dropped. +inline int trim_repeating_tail(const int32_t * ids, int n, int block) { + while (block > 0 && n >= 2 * block && std::equal(ids + n - block, ids + n, ids + n - 2 * block)) { + n -= block; + } + return n; +} + +// TRANSCRIBE_NO_REPETITION_GUARD=1 turns the guard off (reference parity). +inline bool repetition_guard_enabled() { + static const bool enabled = !env::flag("TRANSCRIBE_NO_REPETITION_GUARD"); + return enabled; +} + +// Call after appending a token. On a loop, trims `ids` to one copy, logs a WARN +// tagged `who`, and returns true: the caller stops decoding and flags the run +// truncated, since whatever the audio said after the loop was never decoded. +inline bool stop_on_repetition(std::vector & ids, const char * who) { + if (!repetition_guard_enabled()) { + return false; + } + const int n = static_cast(ids.size()); + const int block = repeating_tail_block(ids.data(), n); + if (block == 0) { + return false; + } + ids.resize(static_cast(trim_repeating_tail(ids.data(), n, block))); + log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + "%s: output began repeating a %d-token block; decode stopped with the repeats dropped (%d tokens " + "kept). The transcript may be incomplete.", + who, block, static_cast(ids.size())); + return true; +} + +} // namespace transcribe diff --git a/src/transcribe.cpp b/src/transcribe.cpp index dfeb5fb4..f5500f46 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -21,6 +21,7 @@ #include "transcribe-abi.h" #include "transcribe-arch.h" #include "transcribe-backend.h" +#include "transcribe-batch-util.h" #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-model.h" @@ -2384,36 +2385,11 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * // Generic serial fallback: run each utterance in turn and snapshot it. // Correct for every family; only the per-dispatch device throughput of - // a real run_batch() is forgone. + // a real run_batch() is forgone. run_one_inner re-validates the shared + // params (idempotent) before the family run(). session->batch_results.reserve(static_cast(n)); - transcribe_status batch_status = TRANSCRIBE_OK; - for (int i = 0; i < n; ++i) { - if (session->poll_abort()) { - batch_status = TRANSCRIBE_ERR_ABORTED; - break; - } - // run_one_inner clears the scratch slot and writes this utterance's - // result; it re-validates the shared params (idempotent) and - // validates this utterance's pcm[i] / n_samples[i]. - const transcribe_status st = run_one_inner(session, pcm[i], n_samples[i], params); - if (st == TRANSCRIBE_OK) { - // capture_result (not a local field-copy) so every result field — - // including per-utterance timings and raw_text — reaches the - // batch snapshot without a second list to keep in sync. - session->batch_results.push_back(session->capture_result(st)); - } else { - // Malformed-input early returns preserve the previous scratch - // slot, so do NOT snapshot it — record an explicit empty - // failure for this utterance instead. - transcribe_session::ResultSet rs; - rs.status = st; - session->batch_results.push_back(std::move(rs)); - if (st == TRANSCRIBE_ERR_ABORTED) { - batch_status = TRANSCRIBE_ERR_ABORTED; - break; - } - } - } + const transcribe_status batch_status = transcribe::run_batch_serial( + session, pcm, n_samples, n, [&](const float * p, int ns) { return run_one_inner(session, p, ns, params); }); // On abort the loop can break early, leaving fewer than n entries; // synthesize any missing slots so the result-set view always exposes n diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 1e2318a0..db3d64be 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -214,6 +214,21 @@ transcribe_apply_warnings(transcribe_decode_budget_unit) add_test(NAME transcribe_decode_budget_unit COMMAND transcribe_decode_budget_unit) +# ----------------------------------------------------------------------------- +# Greedy-decode repetition guard (pure host, no model) + +add_executable(transcribe_repetition_guard_unit + repetition_guard_unit.cpp) + +target_link_libraries(transcribe_repetition_guard_unit PRIVATE transcribe ggml) + +target_include_directories(transcribe_repetition_guard_unit PRIVATE + ${CMAKE_SOURCE_DIR}/src) + +transcribe_apply_warnings(transcribe_repetition_guard_unit) + +add_test(NAME transcribe_repetition_guard_unit COMMAND transcribe_repetition_guard_unit) + # ----------------------------------------------------------------------------- # MOSS diarized-transcript parser unit test (pure host, no model) # ----------------------------------------------------------------------------- diff --git a/tests/repetition_guard_unit.cpp b/tests/repetition_guard_unit.cpp new file mode 100644 index 00000000..bb7b4d02 --- /dev/null +++ b/tests/repetition_guard_unit.cpp @@ -0,0 +1,144 @@ +// Repetition guard stop rule (pure host, no model). Too eager and it cuts real +// speech mid-utterance; too lax and a looping decode runs to the budget. See +// src/transcribe-repetition-guard.h. + +#include "transcribe-repetition-guard.h" + +#include +#include +#include + +namespace { + +int g_failures = 0; + +void expect(const char * what, bool ok) { + if (!ok) { + std::fprintf(stderr, "FAIL %s\n", what); + ++g_failures; + } +} + +std::vector seq(int first, int count) { + std::vector out; + for (int i = 0; i < count; ++i) { + out.push_back(first + i); + } + return out; +} + +std::vector repeat(const std::vector & block, int copies) { + std::vector out; + for (int c = 0; c < copies; ++c) { + out.insert(out.end(), block.begin(), block.end()); + } + return out; +} + +std::vector cat(std::vector a, const std::vector & b) { + a.insert(a.end(), b.begin(), b.end()); + return a; +} + +struct Decoded { + std::vector ids; + int stopped_at = -1; // 1-based token count when the guard fired, or -1 +}; + +// Feed `stream` one token at a time the way a decode loop does, calling the +// guard after every append. +Decoded decode(const std::vector & stream) { + Decoded d; + for (size_t i = 0; i < stream.size(); ++i) { + d.ids.push_back(stream[i]); + if (transcribe::stop_on_repetition(d.ids, "repetition_guard_unit")) { + d.stopped_at = static_cast(i + 1); + return d; + } + } + return d; +} + +int tail_block(const std::vector & ids) { + return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size())); +} + +} // namespace + +int main(void) { + const std::vector prefix = seq(1000, 10); + + // Detection thresholds. + expect("empty is not a loop", tail_block({}) == 0); + expect("distinct tokens are not a loop", tail_block(seq(1, 200)) == 0); + expect("1-token block x31 is not a loop", tail_block(repeat({7}, 31)) == 0); + expect("1-token block x32 is a loop", tail_block(repeat({7}, 32)) == 1); + expect("2-token block x16 is a loop", tail_block(repeat({7, 8}, 16)) == 2); + expect("4-token block x7 is not a loop", tail_block(repeat(seq(1, 4), 7)) == 0); + expect("4-token block x8 is a loop", tail_block(repeat(seq(1, 4), 8)) == 4); + expect("8-token block x3 is not a loop", tail_block(repeat(seq(1, 8), 3)) == 0); + expect("8-token block x4 is a loop", tail_block(repeat(seq(1, 8), 4)) == 8); + expect("20-token block x4 is a loop", tail_block(repeat(seq(1, 20), 4)) == 20); + expect("64-token block x4 is a loop", tail_block(repeat(seq(1, 64), 4)) == 64); + expect("65-token block is past the limit", tail_block(repeat(seq(1, 65), 6)) == 0); + expect("smallest period wins", tail_block(repeat({7, 8}, 32)) == 2); + expect("loop after a prefix", tail_block(cat(prefix, repeat(seq(1, 10), 4))) == 10); + expect("a broken final copy is not a loop", + tail_block(cat(repeat(seq(1, 10), 5), {99})) == 0); + + // Trimming keeps the prefix and one copy. + { + const std::vector ids = cat(prefix, repeat(seq(1, 10), 6)); + const int n = transcribe::trim_repeating_tail(ids.data(), static_cast(ids.size()), 10); + expect("trim keeps prefix + one copy", n == 20); + expect("trim of block 0 is a no-op", transcribe::trim_repeating_tail(ids.data(), 70, 0) == 70); + } + + if (!transcribe::repetition_guard_enabled()) { + std::fprintf(stdout, "repetition_guard_unit: guard disabled by env, skipping decode-loop cases\n"); + return g_failures > 0 ? 1 : 0; + } + + // Emphatic repetition mid-utterance must not stop the decode. + { + const std::vector stream = cat(cat(prefix, repeat(seq(1, 5), 3)), seq(2000, 40)); + const Decoded d = decode(stream); + expect("5-token phrase x3 mid-utterance keeps decoding", d.stopped_at == -1 && d.ids == stream); + } + { + const std::vector stream = cat(repeat(seq(1, 4), 3), seq(2000, 40)); + const Decoded d = decode(stream); + expect("4-token phrase x3 keeps decoding", d.stopped_at == -1 && d.ids == stream); + } + { + const std::vector stream = cat(repeat({42, 43}, 6), seq(2000, 20)); + const Decoded d = decode(stream); + expect("'no, no, no, no, no, no' keeps decoding", d.stopped_at == -1 && d.ids == stream); + } + + // A runaway loop stops as soon as it qualifies, keeping prefix + one copy. + { + const std::vector loop = seq(1, 10); + const Decoded d = decode(cat(prefix, repeat(loop, 40))); + expect("10-token loop stops after 4 copies", d.stopped_at == 10 + 4 * 10); + expect("10-token loop keeps prefix + one copy", d.ids == cat(prefix, loop)); + } + { + const std::vector loop = seq(1, 30); + const Decoded d = decode(cat(prefix, repeat(loop, 10))); + expect("30-token sentence loop stops", d.stopped_at == 10 + 4 * 30); + expect("30-token sentence loop keeps prefix + one copy", d.ids == cat(prefix, loop)); + } + { + const Decoded d = decode(repeat({5}, 100)); + expect("1-token loop stops at 32", d.stopped_at == 32); + expect("1-token loop keeps one token", d.ids == std::vector{5}); + } + + if (g_failures > 0) { + std::fprintf(stderr, "repetition_guard_unit: %d failures\n", g_failures); + return 1; + } + std::fprintf(stdout, "repetition_guard_unit: ok\n"); + return 0; +} diff --git a/tests/run_dispatch_unit.cpp b/tests/run_dispatch_unit.cpp index bacf4dd6..1ba8fb3e 100644 --- a/tests/run_dispatch_unit.cpp +++ b/tests/run_dispatch_unit.cpp @@ -1,6 +1,7 @@ // run_dispatch_unit.cpp - dispatcher-level transcribe_run behavior tests. #include "transcribe-arch.h" +#include "transcribe-batch-util.h" #include "transcribe-model.h" #include "transcribe-session.h" #include "transcribe.h" @@ -508,8 +509,105 @@ void test_release_scratch_after_run_and_batch() { g_run_throw = false; } +// --------------------------------------------------------------------------- +// Serial batch fallback truncation: one truncated utterance must not mark the +// rest (the flag is per-run state), and its partial transcript must survive. +// fake_family_run derives its status from the session flag, as every +// autoregressive family's run() does. +// --------------------------------------------------------------------------- + +namespace { + +transcribe_status fake_family_run(transcribe_session * session, + const float * pcm, + int n_samples, + const transcribe_run_params * params) { + (void) n_samples; + (void) params; + const bool truncate = pcm[0] > 0.5f; + session->clear_result(); + session->full_text = truncate ? "partial" : "complete"; + session->has_result = true; + if (truncate) { + session->was_truncated = true; + } + return session->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; +} + +transcribe_status fake_family_run_batch(transcribe_session * session, + const float * const * pcm, + const int * n_samples, + int n, + const transcribe_run_params * params) { + return transcribe::run_batch_serial( + session, pcm, n_samples, n, [&](const float * p, int ns) { return fake_family_run(session, p, ns, params); }); +} + +void check_truncated_then_clean(const transcribe::Arch & arch) { + transcribe_model model; + model.arch = &arch; + + transcribe_session session; + session.model = &model; + + transcribe_run_params params; + transcribe_run_params_init(¶ms); + + const float truncating = 1.0f, clean = 0.0f; + const float * pcm[3] = { &truncating, &clean, &clean }; + const int ns[3] = { 1, 1, 1 }; + CHECK(transcribe_run_batch(&session, pcm, ns, 3, ¶ms) == TRANSCRIBE_OK); + CHECK(transcribe_batch_n_results(&session) == 3); + CHECK(transcribe_batch_status(&session, 0) == TRANSCRIBE_ERR_OUTPUT_TRUNCATED); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 0), "partial") == 0); + CHECK(transcribe_batch_status(&session, 1) == TRANSCRIBE_OK); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 1), "complete") == 0); + CHECK(transcribe_batch_status(&session, 2) == TRANSCRIBE_OK); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 2), "complete") == 0); + CHECK(transcribe_was_truncated(&session)); +} + +void test_batch_serial_truncation_is_per_utterance() { + // Family run_batch hook falling back to its serial path. + const transcribe::Arch family_arch = { + "fake-family-serial", + nullptr, + nullptr, + fake_family_run, + fake_family_run_batch, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + }; + check_truncated_then_clean(family_arch); + + // No run_batch hook: the dispatcher's generic serial fallback. + const transcribe::Arch dispatcher_arch = { + "fake-dispatcher-serial", + nullptr, + nullptr, + fake_family_run, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + nullptr, + }; + check_truncated_then_clean(dispatcher_arch); +} + +} // namespace + int main() { test_no_run_hook_clears_and_not_implemented(); + test_batch_serial_truncation_is_per_utterance(); test_release_scratch_after_run_and_batch(); test_run_validate_failure_preserves_snapshot(); test_run_validate_success_clears_and_runs(); From 9345bcb5acf805eabab6e2f4fa6e13c1922f70ea Mon Sep 17 00:00:00 2001 From: CJ Pais Date: Mon, 21 Sep 2026 20:51:52 +0800 Subject: [PATCH 2/4] looser degenerate repetition + add new err --- .../python/src/transcribe_cpp/__init__.py | 8 +- .../python/src/transcribe_cpp/_generated.py | 3 +- bindings/python/src/transcribe_cpp/errors.py | 13 +++ bindings/python/tests/test_errors.py | 9 ++ bindings/rust/sys/src/transcribe_sys.rs | 5 +- bindings/rust/transcribe-cpp/src/error.rs | 23 ++++- bindings/rust/transcribe-cpp/src/session.rs | 24 +++-- .../swift/Sources/TranscribeCpp/ABIHash.swift | 2 +- .../swift/Sources/TranscribeCpp/Session.swift | 6 ++ .../TranscribeCpp/TranscribeError.swift | 7 +- bindings/typescript/src/_generated.ts | 3 +- bindings/typescript/src/errors.ts | 10 +- bindings/typescript/src/index.ts | 15 ++- docs/environment-variables.md | 2 +- docs/input-limits.md | 45 +++++---- examples/cli/main.cpp | 79 +++++++++------ include/transcribe.abihash | 2 +- include/transcribe.h | 43 ++++++--- src/arch/canary/model.cpp | 18 ++-- src/arch/canary_qwen/model.cpp | 9 +- src/arch/cohere/model.cpp | 11 ++- src/arch/funasr_nano/model.cpp | 14 +-- src/arch/granite/model.cpp | 14 +-- src/arch/moonshine/model.cpp | 13 +-- src/arch/moonshine_streaming/model.cpp | 14 ++- src/arch/moss/model.cpp | 9 +- src/arch/qwen3_asr/model.cpp | 9 +- src/arch/voxtral/model.cpp | 14 +-- src/causal_lm/causal_lm.cpp | 28 ++++-- src/causal_lm/causal_lm.h | 4 + src/transcribe-batch-util.cpp | 36 ++++--- src/transcribe-batch-util.h | 19 ++-- src/transcribe-repetition-guard.h | 67 +++++++++++-- src/transcribe-session.h | 18 ++++ src/transcribe.cpp | 29 +++--- .../moonshine_streaming_batch_truncation.cpp | 20 ++-- tests/repetition_guard_unit.cpp | 95 ++++++++++++++----- tests/run_dispatch_unit.cpp | 47 +++++---- 38 files changed, 549 insertions(+), 238 deletions(-) diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index 51687230..2d3afc71 100644 --- a/bindings/python/src/transcribe_cpp/__init__.py +++ b/bindings/python/src/transcribe_cpp/__init__.py @@ -39,6 +39,7 @@ ModelLoadError, NotImplementedByModel, OutOfMemory, + OutputRepetition, OutputTruncated, TranscribeError, UnsupportedRequest, @@ -108,6 +109,7 @@ "Aborted", "InputTooLong", "OutputTruncated", + "OutputRepetition", "native_version", "native_commit", "library_path", @@ -1112,9 +1114,9 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", capabilities advertise ``supports_spec_decode`` (-1 = family default, 0 = disabled, >0 = draft length; silently ignored elsewhere). - On ``Aborted`` (via :meth:`cancel`) and ``OutputTruncated`` the - partial transcript is preserved and attached to the exception as - ``partial_result``.""" + On ``Aborted`` (via :meth:`cancel`) and ``OutputTruncated`` (including + its ``OutputRepetition`` subclass) the partial transcript is preserved + and attached to the exception as ``partial_result``.""" self._cancel.clear() array, n_samples = _pcm_to_carray(pcm) params = _build_run_params(task, language, target_language, timestamps, diff --git a/bindings/python/src/transcribe_cpp/_generated.py b/bindings/python/src/transcribe_cpp/_generated.py index ac9763c2..9c7d4715 100644 --- a/bindings/python/src/transcribe_cpp/_generated.py +++ b/bindings/python/src/transcribe_cpp/_generated.py @@ -13,7 +13,7 @@ # Stable digest of the ABI surface below (structs, enums, macros, layout, # prototypes). A native provider package echoes this back so the API # package can reject an ABI-mismatched provider before dlopen. -PUBLIC_HEADER_HASH = "7df72bf9e667b8c2" +PUBLIC_HEADER_HASH = "9866413f80138057" # === enum constants === TRANSCRIBE_OK = 0 @@ -35,6 +35,7 @@ TRANSCRIBE_ERR_UNSUPPORTED_ITN = 16 TRANSCRIBE_ERR_INPUT_TOO_LONG = 17 TRANSCRIBE_ERR_OUTPUT_TRUNCATED = 18 +TRANSCRIBE_ERR_OUTPUT_REPETITION = 19 TRANSCRIBE_ABI_MODEL_LOAD_PARAMS = 0 TRANSCRIBE_ABI_SESSION_PARAMS = 1 TRANSCRIBE_ABI_RUN_PARAMS = 2 diff --git a/bindings/python/src/transcribe_cpp/errors.py b/bindings/python/src/transcribe_cpp/errors.py index 18609c8d..c92e3f16 100644 --- a/bindings/python/src/transcribe_cpp/errors.py +++ b/bindings/python/src/transcribe_cpp/errors.py @@ -22,6 +22,7 @@ TRANSCRIBE_ERR_INVALID_ARG as ERR_INVALID_ARG, TRANSCRIBE_ERR_NOT_IMPLEMENTED as ERR_NOT_IMPLEMENTED, TRANSCRIBE_ERR_OOM as ERR_OOM, + TRANSCRIBE_ERR_OUTPUT_REPETITION as ERR_OUTPUT_REPETITION, TRANSCRIBE_ERR_OUTPUT_TRUNCATED as ERR_OUTPUT_TRUNCATED, TRANSCRIBE_ERR_SAMPLE_RATE as ERR_SAMPLE_RATE, TRANSCRIBE_ERR_UNSUPPORTED_ARCH as ERR_UNSUPPORTED_ARCH, @@ -113,6 +114,17 @@ class OutputTruncated(TranscribeError): partial_result: "Optional[Result]" = None +class OutputRepetition(OutputTruncated): + """The decode was stopped because the output began repeating itself — + the transcript is incomplete by contract. + + A subclass of :class:`OutputTruncated`, so a handler for incomplete + transcripts catches both. ``partial_result`` holds the partial transcript + with the repeats dropped (one copy kept), or None when the status surfaced + outside a result-bearing call. + """ + + _STATUS_TO_EXC = { ERR_INVALID_ARG: InvalidArgument, ERR_NOT_IMPLEMENTED: NotImplementedByModel, @@ -132,6 +144,7 @@ class OutputTruncated(TranscribeError): ERR_UNSUPPORTED_ITN: UnsupportedRequest, ERR_INPUT_TOO_LONG: InputTooLong, ERR_OUTPUT_TRUNCATED: OutputTruncated, + ERR_OUTPUT_REPETITION: OutputRepetition, } diff --git a/bindings/python/tests/test_errors.py b/bindings/python/tests/test_errors.py index 47066cd3..3aaee5d2 100644 --- a/bindings/python/tests/test_errors.py +++ b/bindings/python/tests/test_errors.py @@ -41,6 +41,7 @@ def test_every_status_maps_to_documented_subclass(): errors.ERR_UNSUPPORTED_ITN: t.UnsupportedRequest, errors.ERR_INPUT_TOO_LONG: t.InputTooLong, errors.ERR_OUTPUT_TRUNCATED: t.OutputTruncated, + errors.ERR_OUTPUT_REPETITION: t.OutputRepetition, } # The mapping table covers every non-OK status the header defines, and # nothing else (a new C status must be mapped deliberately, not by @@ -64,6 +65,14 @@ def test_unknown_status_degrades_to_base_class(): assert exc.status == 999 +def test_output_repetition_is_an_output_truncated(): + # A handler that keeps the partial of an incomplete transcript catches both. + exc = errors.exception_for_status(errors.ERR_OUTPUT_REPETITION, "looped", "run") + assert isinstance(exc, t.OutputRepetition) + assert isinstance(exc, t.OutputTruncated) + assert exc.partial_result is None + + def test_exception_for_status_builds_without_raising(): exc = errors.exception_for_status(errors.ERR_ABORTED, "aborted", "run") assert isinstance(exc, t.Aborted) diff --git a/bindings/rust/sys/src/transcribe_sys.rs b/bindings/rust/sys/src/transcribe_sys.rs index 363cd6e3..fc648364 100644 --- a/bindings/rust/sys/src/transcribe_sys.rs +++ b/bindings/rust/sys/src/transcribe_sys.rs @@ -1,11 +1,11 @@ // @generated by `cargo xtask bindgen` from include/transcribe/extensions.h // DO NOT EDIT BY HAND. Regenerate: `cargo xtask bindgen`. -// Pinned to include/transcribe.abihash = 7df72bf9e667b8c2 +// Pinned to include/transcribe.abihash = 9866413f80138057 /// The public-ABI digest these bindings were generated against /// (sha256/16 over the normalized FFI surface). The load-time version /// gate and the CI drift check both anchor on this value. -pub const PUBLIC_HEADER_HASH: &str = "7df72bf9e667b8c2"; +pub const PUBLIC_HEADER_HASH: &str = "9866413f80138057"; /* automatically generated by rust-bindgen 0.72.1 */ @@ -35,6 +35,7 @@ impl transcribe_status { pub const TRANSCRIBE_ERR_UNSUPPORTED_ITN: transcribe_status = transcribe_status(16); pub const TRANSCRIBE_ERR_INPUT_TOO_LONG: transcribe_status = transcribe_status(17); pub const TRANSCRIBE_ERR_OUTPUT_TRUNCATED: transcribe_status = transcribe_status(18); + pub const TRANSCRIBE_ERR_OUTPUT_REPETITION: transcribe_status = transcribe_status(19); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] diff --git a/bindings/rust/transcribe-cpp/src/error.rs b/bindings/rust/transcribe-cpp/src/error.rs index a1aee480..f874dbb1 100644 --- a/bindings/rust/transcribe-cpp/src/error.rs +++ b/bindings/rust/transcribe-cpp/src/error.rs @@ -66,6 +66,15 @@ pub enum Error { /// The (incomplete) transcript produced before truncation. partial: Option>, }, + /// `TRANSCRIBE_ERR_OUTPUT_REPETITION` — the decode was stopped because the + /// output began repeating itself; the transcript is incomplete by contract. + /// The partial transcript, with the repeats dropped, is always preserved. + #[error("output began repeating before end-of-stream: {message}")] + OutputRepetition { + message: String, + /// The (incomplete) transcript produced before the loop, one copy kept. + partial: Option>, + }, /// The loaded library's base version disagrees with the headers this crate /// was generated against (the pre-1.0 version lock). Raised on first use. #[error("native library version mismatch: {0}")] @@ -103,18 +112,20 @@ impl Error { Error::InputTooLong(_) => S::TRANSCRIBE_ERR_INPUT_TOO_LONG, Error::Aborted { .. } => S::TRANSCRIBE_ERR_ABORTED, Error::OutputTruncated { .. } => S::TRANSCRIBE_ERR_OUTPUT_TRUNCATED, + Error::OutputRepetition { .. } => S::TRANSCRIBE_ERR_OUTPUT_REPETITION, _ => S::TRANSCRIBE_OK, }; s.0 as i32 } /// The partial transcript carried by [`Error::Aborted`] / - /// [`Error::OutputTruncated`], if any. `None` for every other variant. + /// [`Error::OutputTruncated`] / [`Error::OutputRepetition`], if any. `None` + /// for every other variant. pub fn partial(&self) -> Option<&Transcript> { match self { - Error::Aborted { partial, .. } | Error::OutputTruncated { partial, .. } => { - partial.as_deref() - } + Error::Aborted { partial, .. } + | Error::OutputTruncated { partial, .. } + | Error::OutputRepetition { partial, .. } => partial.as_deref(), _ => None, } } @@ -171,6 +182,10 @@ pub(crate) fn error_for_status(status: sys::transcribe_status, context: &str) -> message: msg, partial: None, }, + S::TRANSCRIBE_ERR_OUTPUT_REPETITION => Error::OutputRepetition { + message: msg, + partial: None, + }, _ => Error::Other(msg), } } diff --git a/bindings/rust/transcribe-cpp/src/session.rs b/bindings/rust/transcribe-cpp/src/session.rs index 220f2b63..f15f66e3 100644 --- a/bindings/rust/transcribe-cpp/src/session.rs +++ b/bindings/rust/transcribe-cpp/src/session.rs @@ -137,16 +137,18 @@ impl Session { unsafe { sys::transcribe_was_aborted(self.ptr) } } - /// Whether the most recent decode stopped at the generation budget before - /// end-of-stream (the transcript is incomplete). + /// Whether the most recent decode stopped before end-of-stream, at the + /// generation budget or because the output began repeating (the transcript + /// is incomplete). pub fn was_truncated(&self) -> bool { unsafe { sys::transcribe_was_truncated(self.ptr) } } /// Transcribe one buffer of 16 kHz mono float32 PCM in `[-1, 1]`. /// - /// On an aborted or truncated decode the partial transcript is preserved - /// on the returned [`Error::Aborted`] / [`Error::OutputTruncated`]. + /// On an aborted, truncated, or repetition-stopped decode the partial + /// transcript is preserved on the returned [`Error::Aborted`] / + /// [`Error::OutputTruncated`] / [`Error::OutputRepetition`]. pub fn run(&mut self, pcm: &[f32], options: &RunOptions) -> Result { let (params, _lang, _target, _family) = build_run_params(options)?; let n = clamp_len(pcm.len())?; @@ -171,9 +173,10 @@ impl Session { match status { s if s == sys::transcribe_status::TRANSCRIBE_OK => Ok(self.materialize_run()), s if s == sys::transcribe_status::TRANSCRIBE_ERR_ABORTED - || s == sys::transcribe_status::TRANSCRIBE_ERR_OUTPUT_TRUNCATED => + || s == sys::transcribe_status::TRANSCRIBE_ERR_OUTPUT_TRUNCATED + || s == sys::transcribe_status::TRANSCRIBE_ERR_OUTPUT_REPETITION => { - // Partial transcript is preserved by the C API for both. + // Partial transcript is preserved by the C API for all three. let partial = Box::new(self.materialize_run()); Err(attach_partial(error_for_status(s, "run"), partial)) } @@ -229,6 +232,7 @@ impl Session { results.push(Ok(self.materialize_batch(i))); } else if st == sys::transcribe_status::TRANSCRIBE_ERR_ABORTED || st == sys::transcribe_status::TRANSCRIBE_ERR_OUTPUT_TRUNCATED + || st == sys::transcribe_status::TRANSCRIBE_ERR_OUTPUT_REPETITION { let partial = Box::new(self.materialize_batch(i)); let err = attach_partial(error_for_status(st, "run_batch utterance"), partial); @@ -460,8 +464,8 @@ fn clamp_len(len: usize) -> Result { i32::try_from(len).map_err(|_| Error::InvalidArgument(format!("length {len} exceeds i32::MAX"))) } -/// Replace the (partial-less) Aborted/OutputTruncated error with one carrying -/// the materialized partial transcript. +/// Replace the (partial-less) Aborted/OutputTruncated/OutputRepetition error +/// with one carrying the materialized partial transcript. fn attach_partial(err: Error, partial: Box) -> Error { match err { Error::Aborted { message, .. } => Error::Aborted { @@ -472,6 +476,10 @@ fn attach_partial(err: Error, partial: Box) -> Error { message, partial: Some(partial), }, + Error::OutputRepetition { message, .. } => Error::OutputRepetition { + message, + partial: Some(partial), + }, other => other, } } diff --git a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift index 5ea6cbc1..535db416 100644 --- a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift +++ b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift @@ -13,7 +13,7 @@ import CTranscribe extension Transcribe { /// sha256/16 of the normalized public FFI surface, pinned to the value in /// include/transcribe.abihash at the time this binding was last reviewed. - public static let pinnedHeaderHash = "7df72bf9e667b8c2" + public static let pinnedHeaderHash = "9866413f80138057" /// The public-ABI digest this binding was reviewed against (16 hex chars). public static func headerHash() -> String { pinnedHeaderHash } diff --git a/bindings/swift/Sources/TranscribeCpp/Session.swift b/bindings/swift/Sources/TranscribeCpp/Session.swift index ce4850c0..a5a1e599 100644 --- a/bindings/swift/Sources/TranscribeCpp/Session.swift +++ b/bindings/swift/Sources/TranscribeCpp/Session.swift @@ -105,6 +105,9 @@ public final class Session { case TRANSCRIBE_ERR_OUTPUT_TRUNCATED: return .failure(TranscribeError.outputTruncated( message: TranscribeError.message(s, context), partial: batchTranscript(i))) + case TRANSCRIBE_ERR_OUTPUT_REPETITION: + return .failure(TranscribeError.outputRepetition( + message: TranscribeError.message(s, context), partial: batchTranscript(i))) default: return .failure(TranscribeError.make(s, context: context)) } @@ -176,6 +179,9 @@ public final class Session { case TRANSCRIBE_ERR_OUTPUT_TRUNCATED: throw TranscribeError.outputTruncated( message: TranscribeError.message(status, context), partial: readTranscript()) + case TRANSCRIBE_ERR_OUTPUT_REPETITION: + throw TranscribeError.outputRepetition( + message: TranscribeError.message(status, context), partial: readTranscript()) default: throw TranscribeError.make(status, context: context) } diff --git a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift index d42c82f3..830ee583 100644 --- a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift +++ b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift @@ -5,7 +5,7 @@ import CTranscribe /// failures stay distinct (requirements §3): "no such provider" is not /// "provider can't satisfy this request". /// -/// `.aborted` / `.outputTruncated` carry the preserved partial `Transcript` +/// `.aborted` / `.outputTruncated` / `.outputRepetition` carry the preserved partial `Transcript` /// (the C side keeps partial output readable after those statuses); it is `nil` /// when the error is built outside a run (e.g. by `check`). /// @@ -25,6 +25,9 @@ public enum TranscribeError: Error { case inputTooLong(String) case aborted(message: String, partial: Transcript?) case outputTruncated(message: String, partial: Transcript?) + /// The decode was stopped because the output began repeating itself; the + /// repeats are dropped from `partial`, which is incomplete. + case outputRepetition(message: String, partial: Transcript?) case versionMismatch(String) case busy(String) case other(status: Int32, message: String) @@ -64,6 +67,8 @@ public enum TranscribeError: Error { return .aborted(message: message, partial: nil) case TRANSCRIBE_ERR_OUTPUT_TRUNCATED: return .outputTruncated(message: message, partial: nil) + case TRANSCRIBE_ERR_OUTPUT_REPETITION: + return .outputRepetition(message: message, partial: nil) default: return .other(status: raw, message: message) } diff --git a/bindings/typescript/src/_generated.ts b/bindings/typescript/src/_generated.ts index fff2ca59..8864ed42 100644 --- a/bindings/typescript/src/_generated.ts +++ b/bindings/typescript/src/_generated.ts @@ -11,7 +11,7 @@ // Stable digest of the ABI surface (structs, enums, macros, layout, // prototypes), computed by the Python oracle and pinned here so a header // ABI change turns this binding's drift check red for conscious review. -export const PUBLIC_HEADER_HASH = "7df72bf9e667b8c2"; +export const PUBLIC_HEADER_HASH = "9866413f80138057"; // === enum constants === export const TRANSCRIBE_OK = 0; @@ -33,6 +33,7 @@ export const TRANSCRIBE_ERR_UNSUPPORTED_PNC = 15; export const TRANSCRIBE_ERR_UNSUPPORTED_ITN = 16; export const TRANSCRIBE_ERR_INPUT_TOO_LONG = 17; export const TRANSCRIBE_ERR_OUTPUT_TRUNCATED = 18; +export const TRANSCRIBE_ERR_OUTPUT_REPETITION = 19; export const TRANSCRIBE_ABI_MODEL_LOAD_PARAMS = 0; export const TRANSCRIBE_ABI_SESSION_PARAMS = 1; export const TRANSCRIBE_ABI_RUN_PARAMS = 2; diff --git a/bindings/typescript/src/errors.ts b/bindings/typescript/src/errors.ts index 41ec2e5d..be517cf8 100644 --- a/bindings/typescript/src/errors.ts +++ b/bindings/typescript/src/errors.ts @@ -8,7 +8,7 @@ export class TranscribeError extends Error { readonly status: number; /** Set on per-utterance failures from a batch run. */ utteranceIndex?: number; - /** Any partial transcript recovered before the error (set on Aborted / OutputTruncated). */ + /** Any partial transcript recovered before the error (set on Aborted / OutputTruncated / OutputRepetition). */ partialResult?: TranscriptionResult; constructor(message: string, status: number = g.TRANSCRIBE_OK) { @@ -43,6 +43,13 @@ export class Aborted extends TranscribeError {} /** Raised when decode hits the context/generation cap; carries the partial in `partialResult`. */ export class OutputTruncated extends TranscribeError {} +/** + * Raised when decode is stopped because the output began repeating itself; carries + * the partial (repeats dropped) in `partialResult`. Extends OutputTruncated, so a + * handler for incomplete transcripts catches both. + */ +export class OutputRepetition extends OutputTruncated {} + const STATUS_TO_EXC: Record TranscribeError> = { [g.TRANSCRIBE_ERR_INVALID_ARG]: InvalidArgument, [g.TRANSCRIBE_ERR_NOT_IMPLEMENTED]: NotImplementedByModel, @@ -62,6 +69,7 @@ const STATUS_TO_EXC: Record TranscribeErr [g.TRANSCRIBE_ERR_UNSUPPORTED_ITN]: UnsupportedRequest, [g.TRANSCRIBE_ERR_INPUT_TOO_LONG]: InputTooLong, [g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED]: OutputTruncated, + [g.TRANSCRIBE_ERR_OUTPUT_REPETITION]: OutputRepetition, }; /** Build (do not throw) the mapped exception for a status. */ diff --git a/bindings/typescript/src/index.ts b/bindings/typescript/src/index.ts index 54a651e8..27f41680 100644 --- a/bindings/typescript/src/index.ts +++ b/bindings/typescript/src/index.ts @@ -21,6 +21,7 @@ import { InvalidArgument, ModelLoadError, NotImplementedByModel, + OutputRepetition, OutputTruncated, TranscribeError, UnsupportedRequest, @@ -807,7 +808,8 @@ export class Session { if ( status === g.TRANSCRIBE_ERR_ABORTED || - status === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED + status === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED || + status === g.TRANSCRIBE_ERR_OUTPUT_REPETITION ) { const partial: TranscriptionResult = { ...materialize(n, singleAccessors(n, h)), @@ -817,7 +819,9 @@ export class Session { const exc = status === g.TRANSCRIBE_ERR_ABORTED ? new Aborted(`run aborted`, status) - : new OutputTruncated(`run output truncated`, status); + : status === g.TRANSCRIBE_ERR_OUTPUT_REPETITION + ? new OutputRepetition(`run stopped: output began repeating`, status) + : new OutputTruncated(`run output truncated`, status); exc.partialResult = partial; throw exc; } @@ -918,12 +922,15 @@ export class Session { error.utteranceIndex = i; if ( st === g.TRANSCRIBE_ERR_ABORTED || - st === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED + st === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED || + st === g.TRANSCRIBE_ERR_OUTPUT_REPETITION ) { error.partialResult = { ...materialize(n, batchAccessors(n, h, i)), aborted: st === g.TRANSCRIBE_ERR_ABORTED, - truncated: st === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED, + truncated: + st === g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED || + st === g.TRANSCRIBE_ERR_OUTPUT_REPETITION, }; } out.push({ ok: false, error }); diff --git a/docs/environment-variables.md b/docs/environment-variables.md index 1beb8e73..848a6257 100644 --- a/docs/environment-variables.md +++ b/docs/environment-variables.md @@ -28,7 +28,7 @@ of tests. | `TRANSCRIBE_FORCE_FLASH` | Force flash attention on. Wins over `TRANSCRIBE_NO_FLASH` if both are set. | | `TRANSCRIBE_CONV_DIRECT_DW` / `TRANSCRIBE_CONV_NO_DIRECT_DW` | Force the depthwise-conv dispatch to the direct `conv_2d_dw` path / the im2col path, overriding the per-family backend default. | | `TRANSCRIBE_CONV_DIRECT_PW` / `TRANSCRIBE_CONV_NO_DIRECT_PW` | Force the pointwise-conv dispatch to direct `mul_mat` / im2col, overriding the backend default. | -| `TRANSCRIBE_NO_REPETITION_GUARD` | Turn off the stop for greedy decodes that start repeating themselves (see [`input-limits.md`](input-limits.md)). For byte-exact reference parity; a looping decode then runs to its budget. Read once per process. | +| `TRANSCRIBE_NO_REPETITION_GUARD` | Turn off the stop for greedy decodes that start repeating themselves, and the trim of a repeating tail at a budget stop (see [`input-limits.md`](input-limits.md)). For byte-exact reference parity; a looping decode then runs to its budget and keeps its repeats. Read once per process. | | `TRANSCRIBE_DUMP_DIR=` | Enable the per-stage tensor dumper; writes `.f32` + `.json` per dumped tensor into ``. The basis for the numerical-comparison harness (`scripts/compare_tensors.py`). | | `TRANSCRIBE_PERF_DEBUG` | Print a per-stage timing breakdown to stderr (DEBUG log) on the families that profile (`cohere`, `granite`, `canary`, `canary_qwen`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `whisper`). For whisper, a value containing `cpu` or `all` additionally prints the CPU sub-section breakdown. | | `TRANSCRIBE_VOXTRAL_REALTIME_STREAM_TIMING` | Print a per-component streaming wall-time breakdown at stream finalize (voxtral_realtime). | diff --git a/docs/input-limits.md b/docs/input-limits.md index feced3a3..ed3c6c2a 100644 --- a/docs/input-limits.md +++ b/docs/input-limits.md @@ -14,9 +14,10 @@ were trained on and warn when you cross it. Whatever the bucket, the library never truncates silently: an over-length input is rejected up front with `TRANSCRIBE_ERR_INPUT_TOO_LONG`; a transcript that runs into the context or generation budget mid-decode returns the hard status -`TRANSCRIBE_ERR_OUTPUT_TRUNCATED` (with the partial transcript still readable -and `transcribe_was_truncated()` set); and a soft-window family logs a `WARN` -and proceeds. Every model reports its usable limit through +`TRANSCRIBE_ERR_OUTPUT_TRUNCATED`, and one stopped because the output began +looping returns `TRANSCRIBE_ERR_OUTPUT_REPETITION` (both with the partial +transcript still readable and `transcribe_was_truncated()` set); and a +soft-window family logs a `WARN` and proceeds. Every model reports its usable limit through `transcribe_capabilities::max_audio_ms` (or, per-session, `transcribe_session_get_limits()`) so you can check before you call. (Streaming is the one exception to the non-OK truncation rule — see below.) @@ -97,15 +98,25 @@ chunk size. A greedy decode can also fall into repeating one phrase until the budget runs out. The greedy families (`canary`, `canary_qwen`, `cohere`, `funasr_nano`, `granite`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `voxtral`) -stop as soon as a block of up to 64 tokens has repeated at least 4 times and -the copies cover at least 32 tokens. The repeats are dropped, leaving one copy. -Whatever the audio said after the loop was never decoded, so the run reports -the same `OUTPUT_TRUNCATED` status and `WARN` as a budget stop. `whisper` is -excluded because it recovers from loops with its own temperature fallback. -`voxtral_realtime` is excluded because it emits one token per audio frame, so -its padding tokens repeat through any silence. Set -`TRANSCRIBE_NO_REPETITION_GUARD=1` to turn the guard off, e.g. for -byte-exact reference parity. +stop as soon as a block of up to 64 tokens has repeated at least 8 times and +the copies cover at least 64 tokens (a single repeated token needs 64 copies, +a sentence-length block 8). The bar is high on purpose: stopping early loses +whatever the audio said after the loop, so a line sung or chanted a few times +must not trigger it. The repeats are dropped, leaving one copy, and the run +returns `TRANSCRIBE_ERR_OUTPUT_REPETITION` with a `WARN`. Like +`OUTPUT_TRUNCATED` it is result-bearing: the partial transcript is readable +and `transcribe_was_truncated()` is true. + +A decode that runs out of budget has already failed, so the cleanup bar there +is lower: if the partial ends in a block repeated at least 3 times over at +least 32 tokens, the repeats are dropped before the transcript is returned +(still `OUTPUT_TRUNCATED`, with a `WARN` saying how many tokens were dropped). + +`whisper` is excluded because it recovers from loops with its own temperature +fallback. `voxtral_realtime` is excluded because it emits one token per audio +frame, so its padding tokens repeat through any silence. Set +`TRANSCRIBE_NO_REPETITION_GUARD=1` to turn off both the stop and the +budget-stop cleanup, e.g. for byte-exact reference parity. ### 3. Soft window — warn and proceed @@ -178,13 +189,13 @@ with `TRANSCRIBE_ERR_INPUT_TOO_LONG` (one-shot and batch) or surfaced via | Input within limit and decode completes | `TRANSCRIBE_OK` | — | full transcript | | Over-length, hard-cap family | `TRANSCRIBE_ERR_INPUT_TOO_LONG` | `ERROR` via callback | no transcript (rejected before the decode) | | Generation ran long mid-decode | `TRANSCRIBE_ERR_OUTPUT_TRUNCATED` | `WARN` via callback | partial transcript readable; `transcribe_was_truncated() == true` | -| Greedy decode started repeating | `TRANSCRIBE_ERR_OUTPUT_TRUNCATED` | `WARN` via callback | partial transcript readable, repeats dropped; `transcribe_was_truncated() == true` | +| Greedy decode started repeating | `TRANSCRIBE_ERR_OUTPUT_REPETITION` | `WARN` via callback | partial transcript readable, repeats dropped; `transcribe_was_truncated() == true` | | Over-window, soft-window family | `TRANSCRIBE_OK` | `WARN` via callback | full transcript (accuracy may be degraded) | | Chunked / unbounded family | `TRANSCRIBE_OK` | — | full transcript | | Cache/graph allocation failed | `TRANSCRIBE_ERR_OOM` | `ERROR` via callback | no transcript (no silent context shrink) | -In `transcribe_run_batch`, `INPUT_TOO_LONG` and `OUTPUT_TRUNCATED` are -per-utterance statuses (`transcribe_batch_status(session, i)`); the whole-batch +In `transcribe_run_batch`, `INPUT_TOO_LONG`, `OUTPUT_TRUNCATED`, and +`OUTPUT_REPETITION` are per-utterance statuses (`transcribe_batch_status(session, i)`); the whole-batch call returns `TRANSCRIBE_OK`. `transcribe_was_truncated(session)` is reset at the top of every @@ -193,8 +204,8 @@ lifecycle as `transcribe_was_aborted`). ## Streaming is the exception -`TRANSCRIBE_ERR_OUTPUT_TRUNCATED` is an **offline-only** status -(`transcribe_run` / `transcribe_run_batch`). An active stream is incremental +`TRANSCRIBE_ERR_OUTPUT_TRUNCATED` and `TRANSCRIBE_ERR_OUTPUT_REPETITION` are +**offline-only** statuses (`transcribe_run` / `transcribe_run_batch`). An active stream is incremental and has its own terminal-state machine (`transcribe_stream_*`, IDLE/ACTIVE/FINISHED/FAILED), and `stream_feed` / `stream_finalize` return the status of *that step*, not a verdict on the whole transcript. So when a diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index cce1a200..9eae9ac4 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -151,6 +151,16 @@ std::string raw_text_json(const char * raw, const char * clean) { return out; } +// A decode cut short before end-of-stream (budget or repetition stop): non-OK, +// but the partial transcript is preserved (see docs/input-limits.md). +bool is_cut_short(transcribe_status st) { + return st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED || st == TRANSCRIBE_ERR_OUTPUT_REPETITION; +} + +const char * cut_short_label(transcribe_status st) { + return st == TRANSCRIBE_ERR_OUTPUT_REPETITION ? "stopped repeating" : "truncated"; +} + // ",\"speakers\":[...]" fragment: the "who spoke when" rows. Emitted only // when the run produced speaker segments. p is omitted unless finite // (NaN — "model provides no confidence" — is not representable in JSON). @@ -899,6 +909,7 @@ int main(int argc, char ** argv) { int n_ok = 0; int n_truncated = 0; // result-bearing: hit the generation cap, partial hyp emitted + int n_repeating = 0; // result-bearing: stopped when the output looped, partial hyp emitted int n_fail = 0; // no usable result (wav load / backend / unsupported / whole-batch) // Offline batched path: group up to batch_size utterances into one @@ -968,15 +979,16 @@ int main(int argc, char ** argv) { } for (size_t k = 0; k < src_index.size(); ++k) { - const std::string & wav = wav_paths[src_index[k]]; - const transcribe_status ust = transcribe_batch_status(ctx, static_cast(k)); - // OUTPUT_TRUNCATED is result-bearing: the partial transcript is - // preserved and readable via transcribe_batch_full_text (see - // transcribe.h). Emit it as the hyp so downstream tooling scores - // the partial rather than an empty string; the error field below - // still tags it so the truncation stays visible. - const bool result_present = ust == TRANSCRIBE_OK || ust == TRANSCRIBE_ERR_OUTPUT_TRUNCATED; - const char * text = ""; + const std::string & wav = wav_paths[src_index[k]]; + const transcribe_status ust = transcribe_batch_status(ctx, static_cast(k)); + // OUTPUT_TRUNCATED / OUTPUT_REPETITION are result-bearing: the + // partial transcript is preserved and readable via + // transcribe_batch_full_text (see transcribe.h). Emit it as the hyp + // so downstream tooling scores the partial rather than an empty + // string; the error field below still tags it so the stop stays + // visible. + const bool result_present = ust == TRANSCRIBE_OK || is_cut_short(ust); + const char * text = ""; if (result_present) { const char * t = transcribe_batch_full_text(ctx, static_cast(k)); if (t && *t) { @@ -987,6 +999,8 @@ int main(int argc, char ** argv) { ++n_ok; } else if (ust == TRANSCRIBE_ERR_OUTPUT_TRUNCATED) { ++n_truncated; + } else if (ust == TRANSCRIBE_ERR_OUTPUT_REPETITION) { + ++n_repeating; } else { ++n_fail; } @@ -1019,8 +1033,8 @@ int main(int argc, char ** argv) { std::printf("[%zu/%zu] %s", src_index[k] + 1, total, wav.c_str()); if (ust == TRANSCRIBE_OK) { std::printf("\n text: %s\n", text); - } else if (ust == TRANSCRIBE_ERR_OUTPUT_TRUNCATED) { - std::printf(" (truncated)\n text: %s\n", text); + } else if (is_cut_short(ust)) { + std::printf(" (%s)\n text: %s\n", cut_short_label(ust), text); } else { std::printf(" ERROR: %s\n", transcribe_status_string(ust)); } @@ -1107,12 +1121,12 @@ int main(int argc, char ** argv) { run_st = transcribe_run(ctx, pcm.data(), static_cast(pcm.size()), &rp); } - // OUTPUT_TRUNCATED is result-bearing: the partial transcript is - // preserved and readable via transcribe_full_text (see transcribe.h). - // Emit it as the hyp so downstream tooling scores the partial rather - // than an empty string; the error field below still tags it so the - // truncation stays visible. - const bool result_present = run_st == TRANSCRIBE_OK || run_st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + // OUTPUT_TRUNCATED / OUTPUT_REPETITION are result-bearing: the partial + // transcript is preserved and readable via transcribe_full_text (see + // transcribe.h). Emit it as the hyp so downstream tooling scores the + // partial rather than an empty string; the error field below still + // tags it so the stop stays visible. + const bool result_present = run_st == TRANSCRIBE_OK || is_cut_short(run_st); const char * text = ""; if (result_present) { const char * t = transcribe_full_text(ctx); @@ -1124,6 +1138,8 @@ int main(int argc, char ** argv) { ++n_ok; } else if (run_st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED) { ++n_truncated; + } else if (run_st == TRANSCRIBE_ERR_OUTPUT_REPETITION) { + ++n_repeating; } else { ++n_fail; } @@ -1160,8 +1176,8 @@ int main(int argc, char ** argv) { std::printf("[%zu/%zu] %s", i + 1, wav_paths.size(), wav.c_str()); if (run_st == TRANSCRIBE_OK) { std::printf("\n text: %s\n", text); - } else if (run_st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED) { - std::printf(" (truncated)\n text: %s\n", text); + } else if (is_cut_short(run_st)) { + std::printf(" (%s)\n text: %s\n", cut_short_label(run_st), text); } else { std::printf(" ERROR: %s\n", transcribe_status_string(run_st)); } @@ -1172,13 +1188,13 @@ int main(int argc, char ** argv) { } if (!args.batch_jsonl) { - std::fprintf(stderr, "batch: %d ok, %d truncated, %d failed out of %zu\n", n_ok, n_truncated, n_fail, - wav_paths.size()); + std::fprintf(stderr, "batch: %d ok, %d truncated, %d stopped repeating, %d failed out of %zu\n", n_ok, + n_truncated, n_repeating, n_fail, wav_paths.size()); } transcribe_session_free(ctx); transcribe_model_free(model); - // OUTPUT_TRUNCATED is result-bearing and does not fail the batch, but + // OUTPUT_TRUNCATED / OUTPUT_REPETITION are result-bearing and do not fail the batch, but // hard per-utterance failures must remain visible to automation. return n_fail > 0 || !output_ok ? EXIT_FAILURE : EXIT_SUCCESS; } @@ -1416,19 +1432,22 @@ int main(int argc, char ** argv) { } } std::printf("run: %s\n", transcribe_status_string(run_st)); - // OUTPUT_TRUNCATED and ABORTED are non-OK but preserve the partial - // transcript (see docs/input-limits.md), so show the result for them - // too — just flagged. - const bool result_present = - run_st == TRANSCRIBE_OK || run_st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED || run_st == TRANSCRIBE_ERR_ABORTED; + // OUTPUT_TRUNCATED, OUTPUT_REPETITION and ABORTED are non-OK but + // preserve the partial transcript (see docs/input-limits.md), so show + // the result for them too — just flagged. + const bool result_present = run_st == TRANSCRIBE_OK || is_cut_short(run_st) || run_st == TRANSCRIBE_ERR_ABORTED; if (result_present) { const char * text = transcribe_full_text(ctx); std::printf("text: %s\n", (text && *text) ? text : "(empty)"); output_ok = write_output_file(output, args.output_path, text) && output_ok; - // A truncated decode hit the model's context/output budget before - // end-of-stream; the text above is incomplete. - if (transcribe_was_truncated(ctx)) { + // A decode cut short before end-of-stream; the text above is + // incomplete. + if (run_st == TRANSCRIBE_ERR_OUTPUT_REPETITION) { + std::printf( + " note: decode stopped when the output began repeating " + "itself (repeats dropped); transcript is incomplete\n"); + } else if (transcribe_was_truncated(ctx)) { std::printf( " note: output truncated (hit the model's " "context/generation cap before end-of-stream); " diff --git a/include/transcribe.abihash b/include/transcribe.abihash index b0e23c5d..f126c9a7 100644 --- a/include/transcribe.abihash +++ b/include/transcribe.abihash @@ -1 +1 @@ -7df72bf9e667b8c2 +9866413f80138057 diff --git a/include/transcribe.h b/include/transcribe.h index 552c1c7a..94be629d 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -269,12 +269,11 @@ typedef enum { /* * Returned by transcribe_run when the decode stopped because it hit * the model's context / generation budget BEFORE the model emitted - * end-of-stream — i.e. the transcript is incomplete. A greedy decode - * that falls into repeating itself is also stopped early and reported - * here, with the repeats dropped from the partial. This is the - * "started, couldn't finish" counterpart to INPUT_TOO_LONG, and it is - * a hard non-OK status by design: a truncated transcript must not be - * mistaken for a complete one. + * end-of-stream — i.e. the transcript is incomplete. If the partial + * ends in a phrase repeating itself, the repeats are dropped (one + * copy kept). This is the "started, couldn't finish" counterpart to + * INPUT_TOO_LONG, and it is a hard non-OK status by design: a + * truncated transcript must not be mistaken for a complete one. * * The partial transcript IS preserved and readable through the normal * result accessors (transcribe_full_text, segments, words, tokens), @@ -295,6 +294,25 @@ typedef enum { * See docs/input-limits.md for the full contract. */ TRANSCRIBE_ERR_OUTPUT_TRUNCATED = 18, + /* + * Returned by transcribe_run when a greedy decode fell into repeating + * the same block of tokens over and over and was stopped early, + * BEFORE the model emitted end-of-stream. The repeats are dropped + * (one copy kept), but whatever the audio said after the loop was + * never decoded, so the transcript is incomplete. + * + * Result-bearing exactly like OUTPUT_TRUNCATED: the partial transcript + * is readable through the normal result accessors, and + * transcribe_was_truncated() is true. The two codes differ only in why + * the decode stopped: the budget ran out (OUTPUT_TRUNCATED) versus the + * output started looping (this code). Re-running the same audio gives + * the same result; splitting it at a different point may not loop. + * + * Per-utterance in transcribe_run_batch, like OUTPUT_TRUNCATED. Not + * used by streaming. Set TRANSCRIBE_NO_REPETITION_GUARD=1 to disable + * the check. See docs/input-limits.md. + */ + TRANSCRIBE_ERR_OUTPUT_REPETITION = 19, } transcribe_status; /* @@ -1680,8 +1698,9 @@ TRANSCRIBE_API bool transcribe_was_aborted(const struct transcribe_session * ses /* * Supplemental flag for output truncation. True if the most recent decode - * stopped at the model's context / generation cap before end-of-stream, - * leaving the transcript incomplete. The partial transcript is preserved + * stopped before end-of-stream, at the model's context / generation cap or + * because the output started repeating itself, leaving the transcript + * incomplete. The partial transcript is preserved * and readable through the normal result accessors. Reset to false at the * start of each new decode — transcribe_run, transcribe_run_batch, and * transcribe_stream_begin (the same lifecycle as transcribe_was_aborted). @@ -1690,10 +1709,10 @@ TRANSCRIBE_API bool transcribe_was_aborted(const struct transcribe_session * ses * Two paths set it, and they differ in whether a status also reports it: * * - Offline (transcribe_run / transcribe_run_batch): the flag is true - * exactly when the run returned TRANSCRIBE_ERR_OUTPUT_TRUNCATED (or, in - * a batch, when a per-utterance status is OUTPUT_TRUNCATED), so the run - * status is the authoritative signal and this accessor is a convenience - * for a caller that has lost it. + * exactly when the run returned TRANSCRIBE_ERR_OUTPUT_TRUNCATED or + * TRANSCRIBE_ERR_OUTPUT_REPETITION (or, in a batch, when a per-utterance + * status is one of those), so the run status is the authoritative signal + * and this accessor is a convenience for a caller that has lost it. * * - Streaming (transcribe_stream_*): OUTPUT_TRUNCATED is NOT used. An * active stream has its own terminal-state machine, and stream_feed / diff --git a/src/arch/canary/model.cpp b/src/arch/canary/model.cpp index e824a35e..569a7106 100644 --- a/src/arch/canary/model.cpp +++ b/src/arch/canary/model.cpp @@ -1280,8 +1280,8 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "canary run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -1347,8 +1347,8 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "canary run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -1367,6 +1367,7 @@ transcribe_status run(transcribe_session * session, "generation budget / decoder context (%d) before end-of-stream; " "the transcript may be incomplete.", static_cast(generated_ids.size()), cc->kv_cache.n_ctx); + transcribe::trim_repetition_at_budget_stop(generated_ids, "canary run"); } commit_result(); @@ -1374,7 +1375,7 @@ transcribe_status run(transcribe_session * session, // Partial transcript committed above; a truncated decode returns the hard // OUTPUT_TRUNCATED status (the result stays readable, like an aborted run). - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // =========================================================================== @@ -1802,11 +1803,12 @@ transcribe_status run_batch(transcribe_session * session, rs.status = TRANSCRIBE_OK; // Per-utterance truncation parity with the single-shot path: a valid row // that hit the generation budget / context window before eos reports - // TRANSCRIBE_ERR_OUTPUT_TRUNCATED (partial transcript retained). Only - // override an otherwise-OK status — never a worse one. + // TRANSCRIBE_ERR_OUTPUT_TRUNCATED, and one the repetition guard stopped + // reports TRANSCRIBE_ERR_OUTPUT_REPETITION (partial transcript retained + // either way). Only override an otherwise-OK status — never a worse one. if (rs.status == TRANSCRIBE_OK && b < static_cast(truncated.size()) && truncated[b]) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/canary_qwen/model.cpp b/src/arch/canary_qwen/model.cpp index fdd61d5f..2c83701b 100644 --- a/src/arch/canary_qwen/model.cpp +++ b/src/arch/canary_qwen/model.cpp @@ -1201,8 +1201,8 @@ transcribe_status run(transcribe_session * context, cur_past += 1; n_steps += 1; if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "canary_qwen run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -1217,6 +1217,7 @@ transcribe_status run(transcribe_session * context, "the generation budget before end-of-stream; the transcript may " "be incomplete.", static_cast(generated_ids.size())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "canary_qwen run"); } if (profile_decode) { @@ -1270,7 +1271,7 @@ transcribe_status run(transcribe_session * context, // A truncated decode returns OUTPUT_TRUNCATED; the partial transcript above // stays readable (like an aborted run). - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } } // namespace @@ -1659,7 +1660,7 @@ transcribe_status run_batch(transcribe_session * session, // a TRANSCRIBE_OK status, never a worse one. if (b < static_cast(truncated.size()) && truncated[b] && rs.status == TRANSCRIBE_OK) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/cohere/model.cpp b/src/arch/cohere/model.cpp index 92ec38fb..03dc5826 100644 --- a/src/arch/cohere/model.cpp +++ b/src/arch/cohere/model.cpp @@ -1219,7 +1219,7 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "cohere run")) { - cc->was_truncated = true; + cc->mark_repetition_stop(); break; } } @@ -1294,7 +1294,7 @@ transcribe_status run(transcribe_session * session, if (next_token != eos_id) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "cohere run")) { - cc->was_truncated = true; + cc->mark_repetition_stop(); break; } } @@ -1313,6 +1313,9 @@ transcribe_status run(transcribe_session * session, "incomplete.", static_cast(generated_ids.size())); } + if (cc->was_truncated && !cc->stopped_on_repetition) { + transcribe::trim_repetition_at_budget_stop(generated_ids, "cohere run"); + } // Build the result. max_timestamp_kind == NONE means text but no // alignment data: full_text plus one segment (text == full_text, @@ -1323,7 +1326,7 @@ transcribe_status run(transcribe_session * session, // Output truncation is a hard status: the partial transcript is committed // and stays readable (like an aborted run), but we surface the truncation // rather than reporting a clean OK. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // =========================================================================== @@ -1732,7 +1735,7 @@ transcribe_status run_batch(transcribe_session * session, // otherwise-OK status, never a worse one. if (rs.status == TRANSCRIBE_OK && b < static_cast(truncated.size()) && truncated[b]) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/funasr_nano/model.cpp b/src/arch/funasr_nano/model.cpp index fd6628bc..29df291e 100644 --- a/src/arch/funasr_nano/model.cpp +++ b/src/arch/funasr_nano/model.cpp @@ -900,8 +900,8 @@ transcribe_status run(transcribe_session * session, cur_past += 1; n_steps += 1; if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "funasr_nano run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -919,6 +919,7 @@ transcribe_status run(transcribe_session * session, "the generation budget before end-of-stream; the transcript may be " "incomplete.", static_cast(generated_ids.size())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "funasr_nano run"); } if (!generated_ids.empty() && generated_ids.back() == eos_id) { @@ -943,7 +944,7 @@ transcribe_status run(transcribe_session * session, // The partial transcript is fully populated above; a truncated decode // returns the hard OUTPUT_TRUNCATED status (the result stays readable, // like an aborted run). See docs/input-limits.md. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } } // namespace @@ -1361,11 +1362,12 @@ transcribe_status run_batch(transcribe_session * session, rs.segments.push_back(std::move(seg)); // Per-utterance truncation parity with single-shot run(): a row cut at // the generation budget / KV window before eos reports - // TRANSCRIBE_ERR_OUTPUT_TRUNCATED (partial transcript retained). Only - // override a TRANSCRIBE_OK status, never a worse one. + // TRANSCRIBE_ERR_OUTPUT_TRUNCATED, and one the repetition guard stopped + // reports TRANSCRIBE_ERR_OUTPUT_REPETITION (partial transcript retained + // either way). Only override a TRANSCRIBE_OK status, never a worse one. if (b < static_cast(truncated.size()) && truncated[b] && rs.status == TRANSCRIBE_OK) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/granite/model.cpp b/src/arch/granite/model.cpp index 9ebd549d..46202cb0 100644 --- a/src/arch/granite/model.cpp +++ b/src/arch/granite/model.cpp @@ -1242,8 +1242,8 @@ transcribe_status run(transcribe_session * ctx_base, } gen_ids.push_back(next_id); if (transcribe::stop_on_repetition(gen_ids, "granite run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } @@ -1289,6 +1289,7 @@ transcribe_status run(transcribe_session * ctx_base, "generation budget before end-of-stream; the transcript may be " "incomplete.", static_cast(gen_ids.size())); + transcribe::trim_repetition_at_budget_stop(gen_ids, "granite run"); } // Detokenize. @@ -1304,7 +1305,7 @@ transcribe_status run(transcribe_session * ctx_base, // before EOS) is a hard status, not a silent success: surface it so the // caller can distinguish a complete transcript from one cut short. The // partial transcript is still attached above for inspection. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // Offline batched decode (transcribe_run_batch). Serial mel + Conformer @@ -1803,11 +1804,12 @@ transcribe_status run_batch(transcribe_session * session, finalize_granite_result(cm, params, transcript, audio_ms, rs); // Per-utterance truncation parity with single-shot run(): a row cut at // the generation budget / KV window before eos reports - // TRANSCRIBE_ERR_OUTPUT_TRUNCATED (partial transcript retained). Only - // override a TRANSCRIBE_OK status, never a worse one. + // TRANSCRIBE_ERR_OUTPUT_TRUNCATED, and one the repetition guard stopped + // reports TRANSCRIBE_ERR_OUTPUT_REPETITION (partial transcript retained + // either way). Only override a TRANSCRIBE_OK status, never a worse one. if (b < static_cast(truncated.size()) && truncated[b] && rs.status == TRANSCRIBE_OK) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/moonshine/model.cpp b/src/arch/moonshine/model.cpp index 8d6ffc31..29b1656f 100644 --- a/src/arch/moonshine/model.cpp +++ b/src/arch/moonshine/model.cpp @@ -711,8 +711,8 @@ transcribe_status run(transcribe_session * session, if (next_token != eos) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "moonshine run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -741,8 +741,8 @@ transcribe_status run(transcribe_session * session, if (next_token != eos) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "moonshine run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -760,6 +760,7 @@ transcribe_status run(transcribe_session * session, "incomplete. This model is intended for short utterances. See " "transcribe_capabilities.max_audio_ms.", static_cast(generated_ids.size()), max_pos); + transcribe::trim_repetition_at_budget_stop(generated_ids, "moonshine run"); } cc->t_decode_us = ggml_time_us() - t_decode_start; @@ -790,7 +791,7 @@ transcribe_status run(transcribe_session * session, } // Truncation is a hard status; the partial transcript stays readable. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // Offline batched decode (transcribe_run_batch). Mirrors src/arch/cohere + @@ -1101,7 +1102,7 @@ transcribe_status run_batch(transcribe_session * session, // Per-utterance truncation parity with the single-shot path. Only // override an otherwise-OK status — never a worse one. if (rs.status == TRANSCRIBE_OK && b < static_cast(truncated.size()) && truncated[b]) { - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = 0; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/moonshine_streaming/model.cpp b/src/arch/moonshine_streaming/model.cpp index 48d1edfe..4a18af06 100644 --- a/src/arch/moonshine_streaming/model.cpp +++ b/src/arch/moonshine_streaming/model.cpp @@ -997,8 +997,8 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, if (next_token != eos) { generated_ids.push_back(next_token); if (transcribe::stop_on_repetition(generated_ids, "moonshine_streaming run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -1015,6 +1015,7 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, "See transcribe_capabilities.max_audio_ms.", static_cast(generated_ids.size()), hit_duration_budget ? "generation budget" : "position cap", gen_cap); + transcribe::trim_repetition_at_budget_stop(generated_ids, "moonshine_streaming run"); } cc->t_decode_us += ggml_time_us() - t_decode_start; @@ -1200,7 +1201,7 @@ transcribe_status run(transcribe_session * session, // Remap truncation to a hard status only at this offline entry: // decode_from_kv_cache returns OK (it's shared with the streaming finalize // path, which must NOT surface OUTPUT_TRUNCATED). Partial text stays readable. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // Streaming hooks. @@ -2119,12 +2120,13 @@ transcribe_status run_batch(transcribe_session * session, // Batched truncation: a valid row that never reached eos exhausted the // output budget (n_ctx_cap = position cap clamped to cache capacity). - // Mirror the serial path (WARN + flag). + // Mirror the serial path (WARN + flag + repeating-tail trim). { int n_truncated = 0; for (int b = 0; b < n; ++b) { if (valid[b] && !finished[b]) { ++n_truncated; + transcribe::trim_repetition_at_budget_stop(generated[b], "moonshine_streaming run_batch"); } } if (n_truncated > 0) { @@ -2160,7 +2162,9 @@ transcribe_status run_batch(transcribe_session * session, rs.result_kind = TRANSCRIBE_TIMESTAMPS_NONE; rs.has_result = true; // Per-utterance truncation parity (offline run_batch, not streaming). - rs.status = (!finished[b] || repeating[b]) ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + rs.status = repeating[b] ? TRANSCRIBE_ERR_OUTPUT_REPETITION : + !finished[b] ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : + TRANSCRIBE_OK; rs.t_mel_us = 0; rs.t_encode_us = enc_us / valid_count; rs.t_decode_us = dec_us / valid_count; diff --git a/src/arch/moss/model.cpp b/src/arch/moss/model.cpp index 605ae14b..a944f6f8 100644 --- a/src/arch/moss/model.cpp +++ b/src/arch/moss/model.cpp @@ -971,8 +971,8 @@ transcribe_status run(transcribe_session * session, cc->kv_cache.n = cur_past + 1; cc->kv_cache.head = cur_past + 1; if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "moss run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } if (cc->poll_abort()) { @@ -984,6 +984,7 @@ transcribe_status run(transcribe_session * session, cc->was_truncated = true; log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "moss run: output truncated at %d tokens", static_cast(generated_ids.size())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "moss run"); } if (!generated_ids.empty() && generated_ids.back() == eos_id) { generated_ids.pop_back(); @@ -1029,7 +1030,7 @@ transcribe_status run(transcribe_session * session, install_transcript(*cc, params, raw_text, audio_ms); cc->has_result = true; - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // --------------------------------------------------------------------------- @@ -1324,7 +1325,7 @@ transcribe_status run_batch(transcribe_session * session, transcribe_session::ResultSet rs = finalize_utterance(cm, params, generated[b], n_samples[b]); if (b < static_cast(truncated.size()) && truncated[b]) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } cc->batch_results.push_back(std::move(rs)); } diff --git a/src/arch/qwen3_asr/model.cpp b/src/arch/qwen3_asr/model.cpp index fb1a02b3..db970ed9 100644 --- a/src/arch/qwen3_asr/model.cpp +++ b/src/arch/qwen3_asr/model.cpp @@ -971,8 +971,8 @@ transcribe_status run(transcribe_session * session, cc->kv_cache.head = cur_past + 1; t_step_get_us += ggml_time_us() - t_comp1; if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "qwen3_asr run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -989,6 +989,7 @@ transcribe_status run(transcribe_session * session, "generation budget before end-of-stream; the transcript may be " "incomplete.", static_cast(generated_ids.size())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "qwen3_asr run"); } // Map granular counters to the debug-print shape. With graph reuse all @@ -1088,7 +1089,7 @@ transcribe_status run(transcribe_session * session, // A truncated decode returns OUTPUT_TRUNCATED; the partial transcript above // stays readable (like an aborted run). - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // =========================================================================== @@ -1685,7 +1686,7 @@ transcribe_status run_batch(transcribe_session * session, // Per-utterance truncation parity with the single-shot path. if (b < static_cast(truncated.size()) && truncated[b]) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/arch/voxtral/model.cpp b/src/arch/voxtral/model.cpp index 9739df9f..b536a79c 100644 --- a/src/arch/voxtral/model.cpp +++ b/src/arch/voxtral/model.cpp @@ -981,8 +981,8 @@ transcribe_status run(transcribe_session * session, cc->kv_cache.n = cur_past + 1; cc->kv_cache.head = cur_past + 1; if (next_tok != eos_id && transcribe::stop_on_repetition(generated_ids, "voxtral run")) { - cc->was_truncated = true; - repeating = true; + cc->mark_repetition_stop(); + repeating = true; break; } } @@ -998,6 +998,7 @@ transcribe_status run(transcribe_session * session, "generation budget before end-of-stream; the transcript may be " "incomplete.", static_cast(generated_ids.size())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "voxtral run"); } if (!generated_ids.empty() && generated_ids.back() == eos_id) { @@ -1028,7 +1029,7 @@ transcribe_status run(transcribe_session * session, // Output truncation is a hard status: the partial transcript stays readable // (like an aborted run) but the caller is told, not given a clean OK. - return cc->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return cc->truncation_status(); } // --------------------------------------------------------------------------- @@ -1543,11 +1544,12 @@ transcribe_status run_batch(transcribe_session * session, rs.segments.push_back(std::move(seg)); // Per-utterance truncation parity with single-shot run(): a row cut at // the generation budget / KV window before eos reports - // TRANSCRIBE_ERR_OUTPUT_TRUNCATED (partial transcript retained). Only - // override a TRANSCRIBE_OK status, never a worse one. + // TRANSCRIBE_ERR_OUTPUT_TRUNCATED, and one the repetition guard stopped + // reports TRANSCRIBE_ERR_OUTPUT_REPETITION (partial transcript retained + // either way). Only override a TRANSCRIBE_OK status, never a worse one. if (b < static_cast(truncated.size()) && truncated[b] && rs.status == TRANSCRIBE_OK) { cc->was_truncated = true; - rs.status = TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + rs.status = transcribe::decode_stop_status(truncated[b]); } rs.t_mel_us = mel_us / valid_count; rs.t_encode_us = enc_us / valid_count; diff --git a/src/causal_lm/causal_lm.cpp b/src/causal_lm/causal_lm.cpp index 001d9e58..93971f91 100644 --- a/src/causal_lm/causal_lm.cpp +++ b/src/causal_lm/causal_lm.cpp @@ -900,6 +900,7 @@ transcribe_status run_batched_step_loop(transcribe_session * sess std::vector n_past = state.n_past; const std::vector & valid = state.valid; std::vector finished(n, 1); + std::vector repeating(n, 0); for (int b = 0; b < n; ++b) { if (valid[b]) { finished[b] = (next_tok[b] == eos_id); @@ -964,8 +965,12 @@ transcribe_status run_batched_step_loop(transcribe_session * sess if (n_past[b] < max_n_kv) { mask_buf[base + n_past[b]] = mz; } - if (tok == eos_id || stop_on_repetition(generated[b], "batched decode") || - static_cast(generated[b].size()) >= max_new || n_past[b] + 1 > max_n_kv) { + if (tok == eos_id) { + finished[b] = 1; + } else if (stop_on_repetition(generated[b], "batched decode")) { + finished[b] = 1; + repeating[b] = 1; + } else if (static_cast(generated[b].size()) >= max_new || n_past[b] + 1 > max_n_kv) { finished[b] = 1; } else { all_done = false; @@ -978,16 +983,25 @@ transcribe_status run_batched_step_loop(transcribe_session * sess stats->step_us = ggml_time_us() - t_step0; } - // A valid row was truncated if it stopped for a reason OTHER than eos + // A valid row was cut off if it stopped for a reason OTHER than eos // (generation budget, KV window, repetition guard). `finished` is set on // every stop reason, so it can't discriminate; the signal is the last sampled token: // `next_tok[b] != eos_id` means the row was cut off mid-transcript (it is // frozen once the row finishes). See docs/input-limits.md. - if (truncated_out != nullptr) { - truncated_out->assign(n, 0); - for (int b = 0; b < n; ++b) { - (*truncated_out)[b] = (valid[b] && next_tok[b] != eos_id) ? 1 : 0; + std::vector stop(n, k_stop_eos); + for (int b = 0; b < n; ++b) { + if (!valid[b] || next_tok[b] == eos_id) { + continue; } + if (repeating[b]) { + stop[b] = k_stop_repetition; + } else { + stop[b] = k_stop_budget; + trim_repetition_at_budget_stop(generated[b], "batched decode"); + } + } + if (truncated_out != nullptr) { + *truncated_out = std::move(stop); } return TRANSCRIBE_OK; } diff --git a/src/causal_lm/causal_lm.h b/src/causal_lm/causal_lm.h index b5f477c8..8ab253e5 100644 --- a/src/causal_lm/causal_lm.h +++ b/src/causal_lm/causal_lm.h @@ -340,6 +340,10 @@ struct StepLoopStats { // session->poll_abort() once per step. The step graph must already be built // and allocated on `sched`. Returns TRANSCRIBE_ERR_ABORTED on abort, // TRANSCRIBE_ERR_GGUF on a compute failure, else TRANSCRIBE_OK. +// +// truncated_out (if non-null) receives each row's transcribe::DecodeStop, as +// in run_batched_encdec_step_loop; a budget-stopped row has its repeating tail +// trimmed. transcribe_status run_batched_step_loop(transcribe_session * session, ggml_backend_sched_t sched, const StepBatchedIO & io, diff --git a/src/transcribe-batch-util.cpp b/src/transcribe-batch-util.cpp index 3901e61b..4ce61792 100644 --- a/src/transcribe-batch-util.cpp +++ b/src/transcribe-batch-util.cpp @@ -210,14 +210,16 @@ transcribe_status run_batch_serial(transcribe_session * session, return TRANSCRIBE_ERR_ABORTED; } session->clear_result(); - session->t_mel_us = 0; - session->t_encode_us = 0; - session->t_decode_us = 0; - session->was_truncated = false; + session->t_mel_us = 0; + session->t_encode_us = 0; + session->t_decode_us = 0; + session->was_truncated = false; + session->stopped_on_repetition = false; const transcribe_status st = (pcm[i] == nullptr || n_samples[i] <= 0) ? TRANSCRIBE_ERR_INVALID_ARG : run_one(pcm[i], n_samples[i]); - any_truncated = any_truncated || st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + any_truncated = + any_truncated || st == TRANSCRIBE_ERR_OUTPUT_TRUNCATED || st == TRANSCRIBE_ERR_OUTPUT_REPETITION; // The slot was cleared above, so has_result means this utterance wrote // it: keep partials (truncated, aborted), never a stale snapshot. if (st == TRANSCRIBE_OK || session->has_result) { @@ -396,16 +398,24 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * *n_steps_out = n_steps; } - // A valid row that never reached eos was cut off at the generation budget, - // the context window, or the repetition guard — report it as truncated so - // the family can return per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED. See - // docs/input-limits.md. - if (truncated_out != nullptr) { - truncated_out->assign(n, 0); - for (int b = 0; b < n; ++b) { - (*truncated_out)[b] = (valid[b] && (!finished[b] || repeating[b])) ? 1 : 0; + // A valid row that never reached eos was cut off at the generation budget + // or context window, or by the repetition guard. Report why, so the family + // can return the per-utterance status. See docs/input-limits.md. + std::vector stop(n, k_stop_eos); + for (int b = 0; b < n; ++b) { + if (!valid[b]) { + continue; + } + if (repeating[b]) { + stop[b] = k_stop_repetition; + } else if (!finished[b]) { + stop[b] = k_stop_budget; + trim_repetition_at_budget_stop(generated[b], "batched decode"); } } + if (truncated_out != nullptr) { + *truncated_out = std::move(stop); + } return TRANSCRIBE_OK; } diff --git a/src/transcribe-batch-util.h b/src/transcribe-batch-util.h index ad15961f..3fd22dac 100644 --- a/src/transcribe-batch-util.h +++ b/src/transcribe-batch-util.h @@ -106,9 +106,10 @@ transcribe_status decode_batch_slices(transcribe_session * session, // single-utterance run()) and snapshots it into session->batch_results. The // per-run state transcribe_run resets (result slot, timings, truncation flag) // is reset before every utterance, so one truncated utterance cannot mark the -// rest. A truncated or aborted utterance keeps its partial result, as it does -// from transcribe_run; a null pcm or n_samples <= 0 is recorded as -// INVALID_ARG. session->was_truncated ends true if any utterance truncated. +// rest. A truncated, repetition-stopped, or aborted utterance keeps its partial +// result, as it does from transcribe_run; a null pcm or n_samples <= 0 is +// recorded as INVALID_ARG. session->was_truncated ends true if any utterance +// truncated or stopped on repetition. // Returns TRANSCRIBE_ERR_ABORTED once an utterance aborts, else OK. using RunOneFn = std::function; transcribe_status run_batch_serial(transcribe_session * session, @@ -157,11 +158,13 @@ using EncDecRebuildFn = std::function; // step. Returns TRANSCRIBE_ERR_ABORTED / TRANSCRIBE_ERR_GGUF / TRANSCRIBE_OK; // *n_steps_out (if non-null) receives the number of compute steps run. // -// truncated_out (if non-null) is sized to n_batch and set per row: 1 when that -// (valid) row hit the generation budget (max_new), the context window -// (max_n_kv) or the repetition guard BEFORE emitting eos_id (transcript -// truncated), else 0. Lets a -// family report per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED from run_batch. +// truncated_out (if non-null) is sized to n_batch and set per row to why that +// row stopped (transcribe::DecodeStop): k_stop_budget when a valid row hit the +// generation budget (max_new) or the context window (max_n_kv) before eos_id, +// k_stop_repetition when the repetition guard stopped it, else k_stop_eos. +// Non-zero means the transcript was cut off; decode_stop_status maps it to the +// per-utterance status. A budget-stopped row has its repeating tail trimmed +// (trim_repetition_at_budget_stop). transcribe_status run_batched_encdec_step_loop(transcribe_session * session, ggml_backend_sched_t sched, const EncDecRebuildFn & rebuild, diff --git a/src/transcribe-repetition-guard.h b/src/transcribe-repetition-guard.h index 498213c1..c87da09a 100644 --- a/src/transcribe-repetition-guard.h +++ b/src/transcribe-repetition-guard.h @@ -4,11 +4,16 @@ // continuation of a phrase is the phrase itself, it repeats until the decode // budget runs out. The guard works on token ids only, so it fits decoders that // read back a device-side argmax, and leaves output that never loops unchanged. +// +// Two bars. Stopping early cuts off whatever the audio said after the loop, so +// it needs strong evidence: many copies. A decode that already hit its budget +// has failed anyway, so dropping a repeating tail needs less. #pragma once #include "transcribe-env.h" #include "transcribe-log.h" +#include "transcribe.h" #include #include @@ -16,17 +21,22 @@ namespace transcribe { -// A block repeating at the tail is a loop once it has k_repeat_min_copies -// copies covering at least k_repeat_min_tokens. Short blocks need many more -// copies (a 1-token block needs 32), so emphatic human repetition survives. -constexpr int k_repeat_max_block = 64; -constexpr int k_repeat_min_copies = 4; -constexpr int k_repeat_min_tokens = 32; +// A block repeating at the tail is a loop once it has min_copies copies +// covering at least min_tokens. Short blocks need many more copies (a 1-token +// block needs 64 to stop), so emphatic and sung repetition survives. +constexpr int k_repeat_max_block = 64; +constexpr int k_repeat_min_copies = 8; +constexpr int k_repeat_min_tokens = 64; +constexpr int k_budget_trim_min_copies = 3; +constexpr int k_budget_trim_min_tokens = 32; // Length of the block repeating at the end of ids[0, n), or 0. -inline int repeating_tail_block(const int32_t * ids, int n) { +inline int repeating_tail_block(const int32_t * ids, + int n, + int min_copies = k_repeat_min_copies, + int min_tokens = k_repeat_min_tokens) { for (int block = 1; block <= k_repeat_max_block; ++block) { - const int copies = std::max(k_repeat_min_copies, (k_repeat_min_tokens + block - 1) / block); + const int copies = std::max(min_copies, (min_tokens + block - 1) / block); const int span = block * copies; if (span > n) { continue; @@ -60,8 +70,9 @@ inline bool repetition_guard_enabled() { } // Call after appending a token. On a loop, trims `ids` to one copy, logs a WARN -// tagged `who`, and returns true: the caller stops decoding and flags the run -// truncated, since whatever the audio said after the loop was never decoded. +// tagged `who`, and returns true: the caller stops decoding and reports +// TRANSCRIBE_ERR_OUTPUT_REPETITION, since whatever the audio said after the +// loop was never decoded. inline bool stop_on_repetition(std::vector & ids, const char * who) { if (!repetition_guard_enabled()) { return false; @@ -79,4 +90,40 @@ inline bool stop_on_repetition(std::vector & ids, const char * who) { return true; } +// Call once when a decode stopped at its budget or context window before eos. +// Drops the repeats of a block repeating at the tail, at the lower budget-stop +// bar, and logs what it dropped. +inline void trim_repetition_at_budget_stop(std::vector & ids, const char * who) { + if (!repetition_guard_enabled()) { + return; + } + const int n = static_cast(ids.size()); + const int block = repeating_tail_block(ids.data(), n, k_budget_trim_min_copies, k_budget_trim_min_tokens); + if (block == 0) { + return; + } + ids.resize(static_cast(trim_repeating_tail(ids.data(), n, block))); + log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "%s: dropped %d tokens of a repeating %d-token block at the budget stop", who, + n - static_cast(ids.size()), block); +} + +// Why a batched decode row stopped, as reported through a shared step loop's +// truncated_out. Non-zero means the row never reached eos. +enum DecodeStop : char { + k_stop_eos = 0, + k_stop_budget = 1, // generation budget or context window + k_stop_repetition = 2, // stop_on_repetition +}; + +inline transcribe_status decode_stop_status(char stop) { + switch (stop) { + case k_stop_eos: + return TRANSCRIBE_OK; + case k_stop_repetition: + return TRANSCRIBE_ERR_OUTPUT_REPETITION; + default: + return TRANSCRIBE_ERR_OUTPUT_TRUNCATED; + } +} + } // namespace transcribe diff --git a/src/transcribe-session.h b/src/transcribe-session.h index 3aa67b69..73cc69c7 100644 --- a/src/transcribe-session.h +++ b/src/transcribe-session.h @@ -244,6 +244,24 @@ struct transcribe_session { // couldn't start). See docs/input-limits.md. bool was_truncated = false; + // Set with was_truncated when the repetition guard stopped the decode + // (transcribe-repetition-guard.h) rather than the budget, so the run + // reports TRANSCRIBE_ERR_OUTPUT_REPETITION. Cleared with was_truncated. + bool stopped_on_repetition = false; + + void mark_repetition_stop() { + was_truncated = true; + stopped_on_repetition = true; + } + + // Status of a run() whose decode finished: OK, or the stop that cut it short. + transcribe_status truncation_status() const { + if (stopped_on_repetition) { + return TRANSCRIBE_ERR_OUTPUT_REPETITION; + } + return was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + } + // Streaming state. Lifecycle (stream_state) is separated from the // result snapshot so clear_result() can wipe per-call data without // churning the IDLE/ACTIVE/FINISHED/FAILED machine, which the diff --git a/src/transcribe.cpp b/src/transcribe.cpp index f5500f46..2acf9a60 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -156,6 +156,8 @@ extern "C" const char * transcribe_status_string(int status) { return "input audio too long for model context"; case TRANSCRIBE_ERR_OUTPUT_TRUNCATED: return "output truncated: decode hit the context/generation cap before end-of-stream"; + case TRANSCRIBE_ERR_OUTPUT_REPETITION: + return "output repetition: decode stopped when the output began repeating itself"; default: return "unknown status"; } @@ -1827,6 +1829,7 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session session->t_decode_us = 0; session->was_aborted = false; session->was_truncated = false; + session->stopped_on_repetition = false; session->stream_state = TRANSCRIBE_STREAM_ACTIVE; session->stream_commit_policy = commit_policy; session->stream_stable_prefix_agreement_n = stable_prefix_agreement_n; @@ -2212,16 +2215,17 @@ static transcribe_status run_one_inner(struct transcribe_session * sess *committed = true; } session->clear_result(); - session->t_mel_us = 0; - session->t_encode_us = 0; - session->t_decode_us = 0; - session->was_aborted = false; - session->was_truncated = false; + session->t_mel_us = 0; + session->t_encode_us = 0; + session->t_decode_us = 0; + session->was_aborted = false; + session->was_truncated = false; + session->stopped_on_repetition = false; // Force stream_state to IDLE: clear_result deliberately preserves // lifecycle state, but a well-formed transcribe_run subsumes any // prior FINISHED/FAILED stream — after a one-shot run the context // is no longer meaningfully in a streaming lifecycle. - session->stream_state = TRANSCRIBE_STREAM_IDLE; + session->stream_state = TRANSCRIBE_STREAM_IDLE; if (session->model == nullptr || session->model->arch == nullptr || session->model->arch->run == nullptr) { return TRANSCRIBE_ERR_NOT_IMPLEMENTED; @@ -2353,12 +2357,13 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * // Past this point we commit to producing a fresh batch result. session->clear_result(); - session->t_mel_us = 0; - session->t_encode_us = 0; - session->t_decode_us = 0; - session->was_aborted = false; - session->was_truncated = false; - session->stream_state = TRANSCRIBE_STREAM_IDLE; + session->t_mel_us = 0; + session->t_encode_us = 0; + session->t_decode_us = 0; + session->was_aborted = false; + session->was_truncated = false; + session->stopped_on_repetition = false; + session->stream_state = TRANSCRIBE_STREAM_IDLE; session->batch_results.clear(); // Release the compute scratch once the batch has run, whichever path it diff --git a/tests/moonshine_streaming_batch_truncation.cpp b/tests/moonshine_streaming_batch_truncation.cpp index a2ff6078..c569f632 100644 --- a/tests/moonshine_streaming_batch_truncation.cpp +++ b/tests/moonshine_streaming_batch_truncation.cpp @@ -3,14 +3,18 @@ // // Moonshine's cap is on output (max_length decode tokens), not input, so a // long clip runs the decoder into the cap before end-of-stream. In a batch, -// that must surface as a per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED on the -// affected row (with its partial text retained), while a short row that -// finishes normally stays TRANSCRIBE_OK and the whole-batch call still returns -// OK. transcribe_was_truncated() is also set. See docs/input-limits.md. +// that must surface as a per-utterance cut-short status on the affected row +// (with its partial text retained), while a short row that finishes normally +// stays TRANSCRIBE_OK and the whole-batch call still returns OK. +// transcribe_was_truncated() is also set. See docs/input-limits.md. +// +// Past its window the tiny model can fall into a loop before it reaches the +// cap, and then the repetition guard stops it first. Which stop fires depends +// on the model's output, so row 1 accepts either cut-short status. // // Batch makeup: // row 0 = jfk.wav (~11 s) -> completes under the cap -> OK -// row 1 = love-loss.wav (~197 s) -> exceeds the cap -> OUTPUT_TRUNCATED +// row 1 = love-loss.wav (~197 s) -> exceeds the cap -> OUTPUT_TRUNCATED / OUTPUT_REPETITION // // Gating: // - TRANSCRIBE_BUILD_REAL_MODEL_TESTS (CMake, default OFF) builds it. @@ -111,9 +115,11 @@ int main() { CHECK(transcribe_run_batch(s, pcms, lens, 2, nullptr) == TRANSCRIBE_OK); CHECK_EQ_INT(transcribe_batch_n_results(s), 2); - // Row 0 (short) completes; row 1 (long) hits the output cap. + // Row 0 (short) completes; row 1 (long) is cut short at the output cap or + // by the repetition guard. CHECK(transcribe_batch_status(s, 0) == TRANSCRIBE_OK); - CHECK(transcribe_batch_status(s, 1) == TRANSCRIBE_ERR_OUTPUT_TRUNCATED); + const transcribe_status st1 = transcribe_batch_status(s, 1); + CHECK(st1 == TRANSCRIBE_ERR_OUTPUT_TRUNCATED || st1 == TRANSCRIBE_ERR_OUTPUT_REPETITION); // Both rows keep their (partial, for row 1) transcript. for (int i = 0; i < 2; ++i) { diff --git a/tests/repetition_guard_unit.cpp b/tests/repetition_guard_unit.cpp index bb7b4d02..25e0583d 100644 --- a/tests/repetition_guard_unit.cpp +++ b/tests/repetition_guard_unit.cpp @@ -63,37 +63,64 @@ int tail_block(const std::vector & ids) { return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size())); } +int budget_tail_block(const std::vector & ids) { + return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size()), + transcribe::k_budget_trim_min_copies, transcribe::k_budget_trim_min_tokens); +} + +std::vector budget_trimmed(std::vector ids) { + transcribe::trim_repetition_at_budget_stop(ids, "repetition_guard_unit"); + return ids; +} + } // namespace int main(void) { const std::vector prefix = seq(1000, 10); - // Detection thresholds. + // Stop thresholds: 8 copies covering at least 64 tokens. expect("empty is not a loop", tail_block({}) == 0); expect("distinct tokens are not a loop", tail_block(seq(1, 200)) == 0); - expect("1-token block x31 is not a loop", tail_block(repeat({7}, 31)) == 0); - expect("1-token block x32 is a loop", tail_block(repeat({7}, 32)) == 1); - expect("2-token block x16 is a loop", tail_block(repeat({7, 8}, 16)) == 2); - expect("4-token block x7 is not a loop", tail_block(repeat(seq(1, 4), 7)) == 0); - expect("4-token block x8 is a loop", tail_block(repeat(seq(1, 4), 8)) == 4); - expect("8-token block x3 is not a loop", tail_block(repeat(seq(1, 8), 3)) == 0); - expect("8-token block x4 is a loop", tail_block(repeat(seq(1, 8), 4)) == 8); - expect("20-token block x4 is a loop", tail_block(repeat(seq(1, 20), 4)) == 20); - expect("64-token block x4 is a loop", tail_block(repeat(seq(1, 64), 4)) == 64); - expect("65-token block is past the limit", tail_block(repeat(seq(1, 65), 6)) == 0); - expect("smallest period wins", tail_block(repeat({7, 8}, 32)) == 2); - expect("loop after a prefix", tail_block(cat(prefix, repeat(seq(1, 10), 4))) == 10); - expect("a broken final copy is not a loop", - tail_block(cat(repeat(seq(1, 10), 5), {99})) == 0); + expect("1-token block x63 is not a loop", tail_block(repeat({ 7 }, 63)) == 0); + expect("1-token block x64 is a loop", tail_block(repeat({ 7 }, 64)) == 1); + expect("2-token block x31 is not a loop", tail_block(repeat({ 7, 8 }, 31)) == 0); + expect("2-token block x32 is a loop", tail_block(repeat({ 7, 8 }, 32)) == 2); + expect("4-token block x15 is not a loop", tail_block(repeat(seq(1, 4), 15)) == 0); + expect("4-token block x16 is a loop", tail_block(repeat(seq(1, 4), 16)) == 4); + expect("8-token block x7 is not a loop", tail_block(repeat(seq(1, 8), 7)) == 0); + expect("8-token block x8 is a loop", tail_block(repeat(seq(1, 8), 8)) == 8); + expect("20-token block x4 is not a loop", tail_block(repeat(seq(1, 20), 4)) == 0); + expect("20-token block x7 is not a loop", tail_block(repeat(seq(1, 20), 7)) == 0); + expect("20-token block x8 is a loop", tail_block(repeat(seq(1, 20), 8)) == 20); + expect("64-token block x8 is a loop", tail_block(repeat(seq(1, 64), 8)) == 64); + expect("65-token block is past the limit", tail_block(repeat(seq(1, 65), 10)) == 0); + expect("smallest period wins", tail_block(repeat({ 7, 8 }, 64)) == 2); + expect("loop after a prefix", tail_block(cat(prefix, repeat(seq(1, 10), 8))) == 10); + expect("a broken final copy is not a loop", tail_block(cat(repeat(seq(1, 10), 9), { 99 })) == 0); + + // Budget-stop thresholds: 3 copies covering at least 32 tokens. + expect("budget: 20-token block x2 is not a loop", budget_tail_block(repeat(seq(1, 20), 2)) == 0); + expect("budget: 20-token block x3 is a loop", budget_tail_block(repeat(seq(1, 20), 3)) == 20); + expect("budget: 8-token block x3 is not a loop", budget_tail_block(repeat(seq(1, 8), 3)) == 0); + expect("budget: 8-token block x4 is a loop", budget_tail_block(repeat(seq(1, 8), 4)) == 8); + expect("budget: 1-token block x31 is not a loop", budget_tail_block(repeat({ 7 }, 31)) == 0); + expect("budget: 1-token block x32 is a loop", budget_tail_block(repeat({ 7 }, 32)) == 1); // Trimming keeps the prefix and one copy. { const std::vector ids = cat(prefix, repeat(seq(1, 10), 6)); - const int n = transcribe::trim_repeating_tail(ids.data(), static_cast(ids.size()), 10); + const int n = transcribe::trim_repeating_tail(ids.data(), static_cast(ids.size()), 10); expect("trim keeps prefix + one copy", n == 20); expect("trim of block 0 is a no-op", transcribe::trim_repeating_tail(ids.data(), 70, 0) == 70); } + // Stop reasons reported by the batched step loops. + expect("eos row is OK", transcribe::decode_stop_status(transcribe::k_stop_eos) == TRANSCRIBE_OK); + expect("budget row is OUTPUT_TRUNCATED", + transcribe::decode_stop_status(transcribe::k_stop_budget) == TRANSCRIBE_ERR_OUTPUT_TRUNCATED); + expect("repetition row is OUTPUT_REPETITION", + transcribe::decode_stop_status(transcribe::k_stop_repetition) == TRANSCRIBE_ERR_OUTPUT_REPETITION); + if (!transcribe::repetition_guard_enabled()) { std::fprintf(stdout, "repetition_guard_unit: guard disabled by env, skipping decode-loop cases\n"); return g_failures > 0 ? 1 : 0; @@ -111,28 +138,48 @@ int main(void) { expect("4-token phrase x3 keeps decoding", d.stopped_at == -1 && d.ids == stream); } { - const std::vector stream = cat(repeat({42, 43}, 6), seq(2000, 20)); + const std::vector stream = cat(repeat({ 42, 43 }, 6), seq(2000, 20)); const Decoded d = decode(stream); expect("'no, no, no, no, no, no' keeps decoding", d.stopped_at == -1 && d.ids == stream); } + { + // A sung line or chant repeated a handful of times, then more speech. + const std::vector stream = cat(cat(prefix, repeat(seq(1, 12), 7)), seq(2000, 40)); + const Decoded d = decode(stream); + expect("12-token line x7 keeps decoding", d.stopped_at == -1 && d.ids == stream); + } + // A runaway loop stops as soon as it qualifies, keeping prefix + one copy. { - const std::vector loop = seq(1, 10); - const Decoded d = decode(cat(prefix, repeat(loop, 40))); - expect("10-token loop stops after 4 copies", d.stopped_at == 10 + 4 * 10); + const std::vector loop = seq(1, 10); + const Decoded d = decode(cat(prefix, repeat(loop, 40))); + expect("10-token loop stops after 8 copies", d.stopped_at == 10 + 8 * 10); expect("10-token loop keeps prefix + one copy", d.ids == cat(prefix, loop)); } { const std::vector loop = seq(1, 30); const Decoded d = decode(cat(prefix, repeat(loop, 10))); - expect("30-token sentence loop stops", d.stopped_at == 10 + 4 * 30); + expect("30-token sentence loop stops", d.stopped_at == 10 + 8 * 30); expect("30-token sentence loop keeps prefix + one copy", d.ids == cat(prefix, loop)); } { - const Decoded d = decode(repeat({5}, 100)); - expect("1-token loop stops at 32", d.stopped_at == 32); - expect("1-token loop keeps one token", d.ids == std::vector{5}); + const Decoded d = decode(repeat({ 5 }, 100)); + expect("1-token loop stops at 64", d.stopped_at == 64); + expect("1-token loop keeps one token", d.ids == std::vector{ 5 }); + } + + // At a budget stop, a shorter repeating tail is dropped; anything below + // the budget bar is left alone. + { + const std::vector loop = seq(1, 20); + expect("budget stop drops a 3-copy tail", budget_trimmed(cat(prefix, repeat(loop, 3))) == cat(prefix, loop)); + const std::vector twice = cat(prefix, repeat(loop, 2)); + expect("budget stop keeps a 2-copy tail", budget_trimmed(twice) == twice); + const std::vector no_no = cat(prefix, repeat({ 42, 43 }, 6)); + expect("budget stop keeps 'no, no, no, no, no, no'", budget_trimmed(no_no) == no_no); + const std::vector clean = cat(prefix, seq(2000, 40)); + expect("budget stop leaves non-repeating output alone", budget_trimmed(clean) == clean); } if (g_failures > 0) { diff --git a/tests/run_dispatch_unit.cpp b/tests/run_dispatch_unit.cpp index 1ba8fb3e..e3bf2150 100644 --- a/tests/run_dispatch_unit.cpp +++ b/tests/run_dispatch_unit.cpp @@ -510,10 +510,10 @@ void test_release_scratch_after_run_and_batch() { } // --------------------------------------------------------------------------- -// Serial batch fallback truncation: one truncated utterance must not mark the -// rest (the flag is per-run state), and its partial transcript must survive. -// fake_family_run derives its status from the session flag, as every -// autoregressive family's run() does. +// Serial batch fallback truncation: one truncated or repetition-stopped +// utterance must not mark the rest (the flags are per-run state), and its +// partial transcript must survive. fake_family_run derives its status from the +// session flags, as every autoregressive family's run() does. // --------------------------------------------------------------------------- namespace { @@ -524,14 +524,17 @@ transcribe_status fake_family_run(transcribe_session * session, const transcribe_run_params * params) { (void) n_samples; (void) params; - const bool truncate = pcm[0] > 0.5f; + const bool repeat = pcm[0] > 1.5f; + const bool truncate = pcm[0] > 0.5f && !repeat; session->clear_result(); - session->full_text = truncate ? "partial" : "complete"; + session->full_text = repeat ? "looped" : truncate ? "partial" : "complete"; session->has_result = true; - if (truncate) { + if (repeat) { + session->mark_repetition_stop(); + } else if (truncate) { session->was_truncated = true; } - return session->was_truncated ? TRANSCRIBE_ERR_OUTPUT_TRUNCATED : TRANSCRIBE_OK; + return session->truncation_status(); } transcribe_status fake_family_run_batch(transcribe_session * session, @@ -553,18 +556,28 @@ void check_truncated_then_clean(const transcribe::Arch & arch) { transcribe_run_params params; transcribe_run_params_init(¶ms); - const float truncating = 1.0f, clean = 0.0f; - const float * pcm[3] = { &truncating, &clean, &clean }; - const int ns[3] = { 1, 1, 1 }; - CHECK(transcribe_run_batch(&session, pcm, ns, 3, ¶ms) == TRANSCRIBE_OK); - CHECK(transcribe_batch_n_results(&session) == 3); - CHECK(transcribe_batch_status(&session, 0) == TRANSCRIBE_ERR_OUTPUT_TRUNCATED); - CHECK(std::strcmp(transcribe_batch_full_text(&session, 0), "partial") == 0); + const float repeating = 2.0f, truncating = 1.0f, clean = 0.0f; + const float * pcm[4] = { &repeating, &clean, &truncating, &clean }; + const int ns[4] = { 1, 1, 1, 1 }; + CHECK(transcribe_run_batch(&session, pcm, ns, 4, ¶ms) == TRANSCRIBE_OK); + CHECK(transcribe_batch_n_results(&session) == 4); + CHECK(transcribe_batch_status(&session, 0) == TRANSCRIBE_ERR_OUTPUT_REPETITION); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 0), "looped") == 0); CHECK(transcribe_batch_status(&session, 1) == TRANSCRIBE_OK); CHECK(std::strcmp(transcribe_batch_full_text(&session, 1), "complete") == 0); - CHECK(transcribe_batch_status(&session, 2) == TRANSCRIBE_OK); - CHECK(std::strcmp(transcribe_batch_full_text(&session, 2), "complete") == 0); + CHECK(transcribe_batch_status(&session, 2) == TRANSCRIBE_ERR_OUTPUT_TRUNCATED); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 2), "partial") == 0); + CHECK(transcribe_batch_status(&session, 3) == TRANSCRIBE_OK); + CHECK(std::strcmp(transcribe_batch_full_text(&session, 3), "complete") == 0); + CHECK(transcribe_was_truncated(&session)); + + // Single-shot: the repetition stop is its own status, keeps its partial, + // and does not leak into the next run. + CHECK(transcribe_run(&session, &repeating, 1, ¶ms) == TRANSCRIBE_ERR_OUTPUT_REPETITION); + CHECK(std::strcmp(transcribe_full_text(&session), "looped") == 0); CHECK(transcribe_was_truncated(&session)); + CHECK(transcribe_run(&session, &clean, 1, ¶ms) == TRANSCRIBE_OK); + CHECK(!transcribe_was_truncated(&session)); } void test_batch_serial_truncation_is_per_utterance() { From 5f48eeb74f2e8311c070674e1d1df118f3c20f22 Mon Sep 17 00:00:00 2001 From: CJ Pais Date: Sat, 26 Sep 2026 09:28:56 +0800 Subject: [PATCH 3/4] review --- .../python/src/transcribe_cpp/__init__.py | 3 +- .../TranscribeCpp/TranscribeError.swift | 24 +++++++ .../TranscribeCppTests/NoModelTests.swift | 10 +++ docs/input-limits.md | 12 ++-- include/transcribe.h | 4 +- src/arch/moonshine_streaming/model.cpp | 37 ++++++++--- src/causal_lm/causal_lm.cpp | 6 +- src/causal_lm/causal_lm.h | 6 +- src/transcribe-batch-util.cpp | 6 +- src/transcribe-batch-util.h | 6 +- src/transcribe-repetition-guard.h | 66 ++++++++++++------- tests/api_smoke.c | 1 + tests/repetition_guard_unit.cpp | 53 +++++++++++++-- 13 files changed, 177 insertions(+), 57 deletions(-) diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index 032ae4c7..3187ee24 100644 --- a/bindings/python/src/transcribe_cpp/__init__.py +++ b/bindings/python/src/transcribe_cpp/__init__.py @@ -1130,7 +1130,8 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", "transcribe_run") except (Aborted, OutputTruncated) as exc: # The C API preserves the partial transcript on the session for - # exactly these two statuses; surface it rather than discard it. + # these statuses (OutputRepetition included, as an OutputTruncated + # subclass); surface it rather than discard it. exc.partial_result = self._materialize() raise return self._materialize() diff --git a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift index 830ee583..ce2c8b03 100644 --- a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift +++ b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift @@ -74,6 +74,30 @@ public enum TranscribeError: Error { } } + /// The partial transcript carried by `.aborted` / `.outputTruncated` / + /// `.outputRepetition`, if any; `nil` for every other case. + public var partial: Transcript? { + switch self { + case .aborted(_, let partial), .outputTruncated(_, let partial), .outputRepetition(_, let partial): + return partial + default: + return nil + } + } + + /// True when the decode stopped before end-of-stream, at the generation + /// budget (`.outputTruncated`) or because the output began repeating + /// (`.outputRepetition`), as `transcribe_was_truncated` reports it. The + /// transcript in `partial` is incomplete. + public var isTruncated: Bool { + switch self { + case .outputTruncated, .outputRepetition: + return true + default: + return false + } + } + /// Throw the mapped error unless `status` is `TRANSCRIBE_OK`. static func check(_ status: transcribe_status, context: String = "") throws { guard status != TRANSCRIBE_OK else { return } diff --git a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift index 9de46724..61fc38e1 100644 --- a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift +++ b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift @@ -38,6 +38,16 @@ final class NoModelTests: XCTestCase { XCTAssertFalse(Transcribe.statusString(3).isEmpty) // ERR_FILE_NOT_FOUND } + func testTruncationStatusesShareOneCheck() { + // A handler for incomplete transcripts covers both cut-short statuses. + XCTAssertTrue(TranscribeError.outputTruncated(message: "", partial: nil).isTruncated) + XCTAssertTrue(TranscribeError.outputRepetition(message: "", partial: nil).isTruncated) + XCTAssertFalse(TranscribeError.aborted(message: "", partial: nil).isTruncated) + XCTAssertFalse(TranscribeError.inputTooLong("").isTruncated) + XCTAssertNil(TranscribeError.outputRepetition(message: "", partial: nil).partial) + XCTAssertNil(TranscribeError.inputTooLong("").partial) + } + func testAtLeastOneDevice() { XCTAssertGreaterThanOrEqual(Transcribe.devices().count, 1) } diff --git a/docs/input-limits.md b/docs/input-limits.md index ed3c6c2a..597a6677 100644 --- a/docs/input-limits.md +++ b/docs/input-limits.md @@ -98,9 +98,11 @@ chunk size. A greedy decode can also fall into repeating one phrase until the budget runs out. The greedy families (`canary`, `canary_qwen`, `cohere`, `funasr_nano`, `granite`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `voxtral`) -stop as soon as a block of up to 64 tokens has repeated at least 8 times and -the copies cover at least 64 tokens (a single repeated token needs 64 copies, -a sentence-length block 8). The bar is high on purpose: stopping early loses +stop as soon as a block of up to 128 tokens has repeated verbatim enough +times: 8 copies covering at least 64 tokens, but never more than 192 tokens' +worth, and never fewer than 4 copies. A single repeated token needs 64 +copies, a sentence-length block (up to 27 tokens) 8, and a paragraph-length +block of 48 tokens or more 4. The bar is high on purpose: stopping early loses whatever the audio said after the loop, so a line sung or chanted a few times must not trigger it. The repeats are dropped, leaving one copy, and the run returns `TRANSCRIBE_ERR_OUTPUT_REPETITION` with a `WARN`. Like @@ -108,8 +110,8 @@ returns `TRANSCRIBE_ERR_OUTPUT_REPETITION` with a `WARN`. Like and `transcribe_was_truncated()` is true. A decode that runs out of budget has already failed, so the cleanup bar there -is lower: if the partial ends in a block repeated at least 3 times over at -least 32 tokens, the repeats are dropped before the transcript is returned +is lower: if the partial ends in a block of up to 256 tokens repeated at +least 3 times over at least 32 tokens, the repeats are dropped before the transcript is returned (still `OUTPUT_TRUNCATED`, with a `WARN` saying how many tokens were dropped). `whisper` is excluded because it recovers from loops with its own temperature diff --git a/include/transcribe.h b/include/transcribe.h index aaaeaeaf..ae05c95a 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -1721,7 +1721,9 @@ TRANSCRIBE_API bool transcribe_was_aborted(const struct transcribe_session * ses * reached its absolute position cap (forcing the stream to FAILED would * discard the committed text the caller has been consuming). There, this * flag is the ONLY signal of truncation: a streaming caller must check - * it after finalize. + * it after finalize. A family that re-decodes the stream from the start + * on each feed (moonshine_streaming) sets it from its latest decode, so + * after finalize it describes the final transcript. * * Distinct from the "couldn't start" rejection: input that cannot fit at * all is rejected before the decode with TRANSCRIBE_ERR_INPUT_TOO_LONG; diff --git a/src/arch/moonshine_streaming/model.cpp b/src/arch/moonshine_streaming/model.cpp index 222d0588..63e4febc 100644 --- a/src/arch/moonshine_streaming/model.cpp +++ b/src/arch/moonshine_streaming/model.cpp @@ -824,7 +824,8 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, MoonshineStreamingModel * cm, int T_enc, const transcribe_run_params * params, - bool emit_dumps) { + bool emit_dumps, + bool interim) { (void) params; if (cc->poll_abort()) { @@ -844,6 +845,15 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, const auto & hp = cm->hparams; const int64_t t_decode_start = ggml_time_us(); + // The truncation flags describe the transcript this decode produces: a + // stream re-decodes from BOS on each feed, and after finalize the flags + // must reflect the final transcript, not an earlier partial. An interim + // (per-feed) decode logs its stop at DEBUG; stream_finalize warns once + // for the transcript it keeps. + cc->was_truncated = false; + cc->stopped_on_repetition = false; + const transcribe_log_level stop_log_level = interim ? TRANSCRIBE_LOG_LEVEL_DEBUG : TRANSCRIBE_LOG_LEVEL_WARN; + auto try_dump = [emit_dumps](const char * name, ggml_tensor * t, const char * stage) { if (!emit_dumps || t == nullptr) { return; @@ -996,7 +1006,7 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, } if (next_token != eos) { generated_ids.push_back(next_token); - if (transcribe::stop_on_repetition(generated_ids, "moonshine_streaming run")) { + if (transcribe::stop_on_repetition(generated_ids, "moonshine_streaming run", stop_log_level)) { cc->mark_repetition_stop(); repeating = true; break; @@ -1009,13 +1019,13 @@ transcribe_status decode_from_kv_cache(MoonshineStreamingSession * cc, if (!repeating && next_token != eos) { cc->was_truncated = true; const bool hit_duration_budget = (gen_cap < max_pos) || (max_pos <= 0); - transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + transcribe::log_msg(stop_log_level, "moonshine run: output truncated at %d tokens — decode reached the " "%s (%d) before end-of-stream; the transcript may be incomplete. " "See transcribe_capabilities.max_audio_ms.", static_cast(generated_ids.size()), hit_duration_budget ? "generation budget" : "position cap", gen_cap); - transcribe::trim_repetition_at_budget_stop(generated_ids, "moonshine_streaming run"); + transcribe::trim_repetition_at_budget_stop(generated_ids, "moonshine_streaming run", stop_log_level); } cc->t_decode_us += ggml_time_us() - t_decode_start; @@ -1134,7 +1144,7 @@ transcribe_status decode_from_committed_enc(MoonshineStreamingSession * cc, } // 4. AR decoder loop. - return decode_from_kv_cache(cc, cm, T_enc, params, emit_dumps); + return decode_from_kv_cache(cc, cm, T_enc, params, emit_dumps, /*interim=*/false); } // Internal one-shot inference helper. Encoder over the full PCM, then @@ -1260,7 +1270,8 @@ void reset_result_text_only(MoonshineStreamingSession * cc) { transcribe_status decode_partial(MoonshineStreamingSession * cc, MoonshineStreamingModel * cm, const transcribe_run_params * params, - bool emit_dumps) { + bool emit_dumps, + bool interim) { const int T_enc = cc->stream_T_emitted; if (T_enc <= 0) { return TRANSCRIBE_ERR_INVALID_ARG; @@ -1275,7 +1286,7 @@ transcribe_status decode_partial(MoonshineStreamingSession * cc, st != TRANSCRIBE_OK) { return st; } - if (auto st = decode_from_kv_cache(cc, cm, T_enc, params, emit_dumps); st != TRANSCRIBE_OK) { + if (auto st = decode_from_kv_cache(cc, cm, T_enc, params, emit_dumps, interim); st != TRANSCRIBE_OK) { return st; } cc->stream_last_decoded_T = T_enc; @@ -1624,7 +1635,7 @@ transcribe_status stream_feed(transcribe_session * session, const std::string prev_full_text = cc->full_text; if (auto st = decode_partial(cc, cm, &cc->stream_run_params, - /*emit_dumps=*/false); + /*emit_dumps=*/false, /*interim=*/true); st != TRANSCRIBE_OK) { return st; } @@ -1742,11 +1753,19 @@ transcribe_status stream_finalize(transcribe_session * session, transcribe_strea const std::string prev_full_text = cc->full_text; if (T_enc > cc->stream_last_decoded_T || !cc->has_result) { if (auto st = decode_partial(cc, cm, &cc->stream_run_params, - /*emit_dumps=*/true); + /*emit_dumps=*/true, /*interim=*/false); st != TRANSCRIBE_OK) { write_update(st); return st; } + } else if (cc->was_truncated) { + // The last feed's interim decode is the final transcript, and it only + // logged its stop at DEBUG. + transcribe::log_msg( + TRANSCRIBE_LOG_LEVEL_WARN, + "moonshine_streaming stream: the final transcript %s before end-of-stream; it may be " + "incomplete.", + cc->stopped_on_repetition ? "stopped when the output began repeating" : "reached the generation budget"); } // Commit the entire result at finalize: tokens, words, and diff --git a/src/causal_lm/causal_lm.cpp b/src/causal_lm/causal_lm.cpp index fd11f3d8..284aa780 100644 --- a/src/causal_lm/causal_lm.cpp +++ b/src/causal_lm/causal_lm.cpp @@ -892,7 +892,7 @@ transcribe_status run_batched_step_loop(transcribe_session * sess const StepBatchedState & state, std::vector> & generated, StepLoopStats * stats, - std::vector * truncated_out) { + std::vector * stop_out) { const int n = n_batch; // Per-row working state. @@ -1000,8 +1000,8 @@ transcribe_status run_batched_step_loop(transcribe_session * sess trim_repetition_at_budget_stop(generated[b], "batched decode"); } } - if (truncated_out != nullptr) { - *truncated_out = std::move(stop); + if (stop_out != nullptr) { + *stop_out = std::move(stop); } return TRANSCRIBE_OK; } diff --git a/src/causal_lm/causal_lm.h b/src/causal_lm/causal_lm.h index f53b54f8..b2129d86 100644 --- a/src/causal_lm/causal_lm.h +++ b/src/causal_lm/causal_lm.h @@ -341,7 +341,7 @@ struct StepLoopStats { // and allocated on `sched`. Returns TRANSCRIBE_ERR_ABORTED on abort, // TRANSCRIBE_ERR_GGUF on a compute failure, else TRANSCRIBE_OK. // -// truncated_out (if non-null) receives each row's transcribe::DecodeStop, as +// stop_out (if non-null) receives each row's transcribe::DecodeStop, as // in run_batched_encdec_step_loop; a budget-stopped row has its repeating tail // trimmed. transcribe_status run_batched_step_loop(transcribe_session * session, @@ -353,7 +353,7 @@ transcribe_status run_batched_step_loop(transcribe_session * sess int max_new, const StepBatchedState & state, std::vector> & generated, - StepLoopStats * stats = nullptr, - std::vector * truncated_out = nullptr); + StepLoopStats * stats = nullptr, + std::vector * stop_out = nullptr); } // namespace transcribe::causal_lm diff --git a/src/transcribe-batch-util.cpp b/src/transcribe-batch-util.cpp index f1bba455..5fa509b6 100644 --- a/src/transcribe-batch-util.cpp +++ b/src/transcribe-batch-util.cpp @@ -251,7 +251,7 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * const std::vector & valid, std::vector> & generated, int * n_steps_out, - std::vector * truncated_out) { + std::vector * stop_out) { const int n = n_batch; const ggml_fp16_t f16_zero = ggml_fp32_to_fp16(0.0f); const ggml_fp16_t f16_ninf = ggml_fp32_to_fp16(-std::numeric_limits::infinity()); @@ -419,8 +419,8 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * trim_repetition_at_budget_stop(generated[b], "batched decode"); } } - if (truncated_out != nullptr) { - *truncated_out = std::move(stop); + if (stop_out != nullptr) { + *stop_out = std::move(stop); } return TRANSCRIBE_OK; } diff --git a/src/transcribe-batch-util.h b/src/transcribe-batch-util.h index d4d04c93..cb2586d6 100644 --- a/src/transcribe-batch-util.h +++ b/src/transcribe-batch-util.h @@ -159,7 +159,7 @@ using EncDecRebuildFn = std::function & valid, std::vector> & generated, - int * n_steps_out = nullptr, - std::vector * truncated_out = nullptr); + int * n_steps_out = nullptr, + std::vector * stop_out = nullptr); } // namespace transcribe diff --git a/src/transcribe-repetition-guard.h b/src/transcribe-repetition-guard.h index c87da09a..d6d37d6e 100644 --- a/src/transcribe-repetition-guard.h +++ b/src/transcribe-repetition-guard.h @@ -21,22 +21,37 @@ namespace transcribe { -// A block repeating at the tail is a loop once it has min_copies copies -// covering at least min_tokens. Short blocks need many more copies (a 1-token -// block needs 64 to stop), so emphatic and sung repetition survives. -constexpr int k_repeat_max_block = 64; -constexpr int k_repeat_min_copies = 8; -constexpr int k_repeat_min_tokens = 64; -constexpr int k_budget_trim_min_copies = 3; -constexpr int k_budget_trim_min_tokens = 32; +// When a block repeating at the tail counts as a loop. The evidence is how +// many tokens repeat verbatim, so a block needs `copies` copies spanning at +// least `min_tokens`, but once that span passes `max_tokens` it only has to +// cover `max_tokens`, with at least `min_copies` copies. Short blocks need many +// more copies (a 1-token block needs 64 to stop), so emphatic and sung +// repetition survives; a paragraph-length loop stops before the budget. +struct RepeatBar { + int max_block; // longest block looked for + int copies; // copies a block needs... + int min_tokens; // ...spanning at least this many tokens... + int max_tokens; // ...or at most this many... + int min_copies; // ...in no fewer than this many copies +}; + +// Stop mid-decode: 8 copies of a block up to 27 tokens, tapering to 4 copies +// of a 48-128-token block (192-512 tokens), which fits the decode budgets. +constexpr RepeatBar k_stop_bar = { 128, 8, 64, 192, 4 }; +// Trim at a budget stop: 3 copies spanning at least 32 tokens (max_tokens +// never binds). The trim runs once per decode, so it looks for longer blocks. +constexpr RepeatBar k_budget_trim_bar = { 256, 3, 32, 256 * 3, 3 }; + +// Copies of a `block`-token block that make a loop under `bar`. +constexpr int repeat_copies_needed(const RepeatBar & bar, int block) { + const int span = std::min(std::max(bar.copies * block, bar.min_tokens), bar.max_tokens); + return std::max(bar.min_copies, (span + block - 1) / block); +} // Length of the block repeating at the end of ids[0, n), or 0. -inline int repeating_tail_block(const int32_t * ids, - int n, - int min_copies = k_repeat_min_copies, - int min_tokens = k_repeat_min_tokens) { - for (int block = 1; block <= k_repeat_max_block; ++block) { - const int copies = std::max(min_copies, (min_tokens + block - 1) / block); +inline int repeating_tail_block(const int32_t * ids, int n, const RepeatBar & bar = k_stop_bar) { + for (int block = 1; block <= bar.max_block; ++block) { + const int copies = repeat_copies_needed(bar, block); const int span = block * copies; if (span > n) { continue; @@ -69,11 +84,14 @@ inline bool repetition_guard_enabled() { return enabled; } -// Call after appending a token. On a loop, trims `ids` to one copy, logs a WARN -// tagged `who`, and returns true: the caller stops decoding and reports +// Call after appending a token. On a loop, trims `ids` to one copy, logs at +// `level` (a WARN unless the decode is an interim one) tagged `who`, and +// returns true: the caller stops decoding and reports // TRANSCRIBE_ERR_OUTPUT_REPETITION, since whatever the audio said after the // loop was never decoded. -inline bool stop_on_repetition(std::vector & ids, const char * who) { +inline bool stop_on_repetition(std::vector & ids, + const char * who, + transcribe_log_level level = TRANSCRIBE_LOG_LEVEL_WARN) { if (!repetition_guard_enabled()) { return false; } @@ -83,7 +101,7 @@ inline bool stop_on_repetition(std::vector & ids, const char * who) { return false; } ids.resize(static_cast(trim_repeating_tail(ids.data(), n, block))); - log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + log_msg(level, "%s: output began repeating a %d-token block; decode stopped with the repeats dropped (%d tokens " "kept). The transcript may be incomplete.", who, block, static_cast(ids.size())); @@ -92,23 +110,25 @@ inline bool stop_on_repetition(std::vector & ids, const char * who) { // Call once when a decode stopped at its budget or context window before eos. // Drops the repeats of a block repeating at the tail, at the lower budget-stop -// bar, and logs what it dropped. -inline void trim_repetition_at_budget_stop(std::vector & ids, const char * who) { +// bar, and logs what it dropped at `level`. +inline void trim_repetition_at_budget_stop(std::vector & ids, + const char * who, + transcribe_log_level level = TRANSCRIBE_LOG_LEVEL_WARN) { if (!repetition_guard_enabled()) { return; } const int n = static_cast(ids.size()); - const int block = repeating_tail_block(ids.data(), n, k_budget_trim_min_copies, k_budget_trim_min_tokens); + const int block = repeating_tail_block(ids.data(), n, k_budget_trim_bar); if (block == 0) { return; } ids.resize(static_cast(trim_repeating_tail(ids.data(), n, block))); - log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "%s: dropped %d tokens of a repeating %d-token block at the budget stop", who, + log_msg(level, "%s: dropped %d tokens of a repeating %d-token block at the budget stop", who, n - static_cast(ids.size()), block); } // Why a batched decode row stopped, as reported through a shared step loop's -// truncated_out. Non-zero means the row never reached eos. +// stop_out. Non-zero means the row never reached eos. enum DecodeStop : char { k_stop_eos = 0, k_stop_budget = 1, // generation budget or context window diff --git a/tests/api_smoke.c b/tests/api_smoke.c index ce776b72..0fb680a0 100644 --- a/tests/api_smoke.c +++ b/tests/api_smoke.c @@ -76,6 +76,7 @@ static void test_status_string(void) { TRANSCRIBE_ERR_UNSUPPORTED_ITN, TRANSCRIBE_ERR_INPUT_TOO_LONG, TRANSCRIBE_ERR_OUTPUT_TRUNCATED, + TRANSCRIBE_ERR_OUTPUT_REPETITION, }; for (size_t i = 0; i < sizeof(all) / sizeof(all[0]); ++i) { const char * s = transcribe_status_string(all[i]); diff --git a/tests/repetition_guard_unit.cpp b/tests/repetition_guard_unit.cpp index 25e0583d..a4c7f255 100644 --- a/tests/repetition_guard_unit.cpp +++ b/tests/repetition_guard_unit.cpp @@ -64,8 +64,7 @@ int tail_block(const std::vector & ids) { } int budget_tail_block(const std::vector & ids) { - return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size()), - transcribe::k_budget_trim_min_copies, transcribe::k_budget_trim_min_tokens); + return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size()), transcribe::k_budget_trim_bar); } std::vector budget_trimmed(std::vector ids) { @@ -78,7 +77,8 @@ std::vector budget_trimmed(std::vector ids) { int main(void) { const std::vector prefix = seq(1000, 10); - // Stop thresholds: 8 copies covering at least 64 tokens. + // Stop thresholds: 8 copies covering at least 64 tokens, tapering to 4 + // copies once the copies cover 192 tokens. expect("empty is not a loop", tail_block({}) == 0); expect("distinct tokens are not a loop", tail_block(seq(1, 200)) == 0); expect("1-token block x63 is not a loop", tail_block(repeat({ 7 }, 63)) == 0); @@ -92,8 +92,29 @@ int main(void) { expect("20-token block x4 is not a loop", tail_block(repeat(seq(1, 20), 4)) == 0); expect("20-token block x7 is not a loop", tail_block(repeat(seq(1, 20), 7)) == 0); expect("20-token block x8 is a loop", tail_block(repeat(seq(1, 20), 8)) == 20); - expect("64-token block x8 is a loop", tail_block(repeat(seq(1, 64), 8)) == 64); - expect("65-token block is past the limit", tail_block(repeat(seq(1, 65), 10)) == 0); + expect("30-token block x6 is not a loop", tail_block(repeat(seq(1, 30), 6)) == 0); + expect("30-token block x7 is a loop", tail_block(repeat(seq(1, 30), 7)) == 30); + expect("47-token block x4 is not a loop", tail_block(repeat(seq(1, 47), 4)) == 0); + expect("47-token block x5 is a loop", tail_block(repeat(seq(1, 47), 5)) == 47); + expect("48-token block x3 is not a loop", tail_block(repeat(seq(1, 48), 3)) == 0); + expect("48-token block x4 is a loop", tail_block(repeat(seq(1, 48), 4)) == 48); + expect("64-token block x4 is a loop", tail_block(repeat(seq(1, 64), 4)) == 64); + expect("128-token block x3 is not a loop", tail_block(repeat(seq(1, 128), 3)) == 0); + expect("128-token block x4 is a loop", tail_block(repeat(seq(1, 128), 4)) == 128); + expect("129-token block is past the limit", tail_block(repeat(seq(1, 129), 10)) == 0); + { + // Longer blocks never need more copies, and past the taper the copies + // always span at least 192 tokens. + bool monotonic = true; + bool spans_192 = true; + for (int block = 2; block <= transcribe::k_stop_bar.max_block; ++block) { + const int copies = transcribe::repeat_copies_needed(transcribe::k_stop_bar, block); + monotonic = monotonic && copies <= transcribe::repeat_copies_needed(transcribe::k_stop_bar, block - 1); + spans_192 = spans_192 && (block < 24 || copies * block >= 192); + } + expect("copies needed never grow with the block", monotonic); + expect("tapered copies span at least 192 tokens", spans_192); + } expect("smallest period wins", tail_block(repeat({ 7, 8 }, 64)) == 2); expect("loop after a prefix", tail_block(cat(prefix, repeat(seq(1, 10), 8))) == 10); expect("a broken final copy is not a loop", tail_block(cat(repeat(seq(1, 10), 9), { 99 })) == 0); @@ -105,6 +126,10 @@ int main(void) { expect("budget: 8-token block x4 is a loop", budget_tail_block(repeat(seq(1, 8), 4)) == 8); expect("budget: 1-token block x31 is not a loop", budget_tail_block(repeat({ 7 }, 31)) == 0); expect("budget: 1-token block x32 is a loop", budget_tail_block(repeat({ 7 }, 32)) == 1); + expect("budget: 100-token block x2 is not a loop", budget_tail_block(repeat(seq(1, 100), 2)) == 0); + expect("budget: 100-token block x3 is a loop", budget_tail_block(repeat(seq(1, 100), 3)) == 100); + expect("budget: 256-token block x3 is a loop", budget_tail_block(repeat(seq(1, 256), 3)) == 256); + expect("budget: 257-token block is past the limit", budget_tail_block(repeat(seq(1, 257), 3)) == 0); // Trimming keeps the prefix and one copy. { @@ -149,6 +174,12 @@ int main(void) { const Decoded d = decode(stream); expect("12-token line x7 keeps decoding", d.stopped_at == -1 && d.ids == stream); } + { + // A long passage repeated a few times, then more speech. + const std::vector stream = cat(cat(prefix, repeat(seq(1, 40), 4)), seq(2000, 40)); + const Decoded d = decode(stream); + expect("40-token passage x4 keeps decoding", d.stopped_at == -1 && d.ids == stream); + } // A runaway loop stops as soon as it qualifies, keeping prefix + one copy. { @@ -160,9 +191,16 @@ int main(void) { { const std::vector loop = seq(1, 30); const Decoded d = decode(cat(prefix, repeat(loop, 10))); - expect("30-token sentence loop stops", d.stopped_at == 10 + 8 * 30); + expect("30-token sentence loop stops after 7 copies", d.stopped_at == 10 + 7 * 30); expect("30-token sentence loop keeps prefix + one copy", d.ids == cat(prefix, loop)); } + { + // A paragraph loop, past the old 64-token block limit. + const std::vector loop = seq(1, 100); + const Decoded d = decode(cat(prefix, repeat(loop, 10))); + expect("100-token paragraph loop stops after 4 copies", d.stopped_at == 10 + 4 * 100); + expect("100-token paragraph loop keeps prefix + one copy", d.ids == cat(prefix, loop)); + } { const Decoded d = decode(repeat({ 5 }, 100)); expect("1-token loop stops at 64", d.stopped_at == 64); @@ -176,6 +214,9 @@ int main(void) { expect("budget stop drops a 3-copy tail", budget_trimmed(cat(prefix, repeat(loop, 3))) == cat(prefix, loop)); const std::vector twice = cat(prefix, repeat(loop, 2)); expect("budget stop keeps a 2-copy tail", budget_trimmed(twice) == twice); + const std::vector paragraph = seq(3000, 150); + expect("budget stop drops a 3-copy 150-token tail", + budget_trimmed(cat(prefix, repeat(paragraph, 3))) == cat(prefix, paragraph)); const std::vector no_no = cat(prefix, repeat({ 42, 43 }, 6)); expect("budget stop keeps 'no, no, no, no, no, no'", budget_trimmed(no_no) == no_no); const std::vector clean = cat(prefix, seq(2000, 40)); From 2f78b64d350bd436a1ea0e268df4faae62364550 Mon Sep 17 00:00:00 2001 From: CJ Pais Date: Sat, 26 Sep 2026 09:30:31 +0800 Subject: [PATCH 4/4] slim docs --- docs/input-limits.md | 23 ++--------------------- 1 file changed, 2 insertions(+), 21 deletions(-) diff --git a/docs/input-limits.md b/docs/input-limits.md index 597a6677..499737c2 100644 --- a/docs/input-limits.md +++ b/docs/input-limits.md @@ -98,27 +98,8 @@ chunk size. A greedy decode can also fall into repeating one phrase until the budget runs out. The greedy families (`canary`, `canary_qwen`, `cohere`, `funasr_nano`, `granite`, `moonshine`, `moonshine_streaming`, `moss`, `qwen3_asr`, `voxtral`) -stop as soon as a block of up to 128 tokens has repeated verbatim enough -times: 8 copies covering at least 64 tokens, but never more than 192 tokens' -worth, and never fewer than 4 copies. A single repeated token needs 64 -copies, a sentence-length block (up to 27 tokens) 8, and a paragraph-length -block of 48 tokens or more 4. The bar is high on purpose: stopping early loses -whatever the audio said after the loop, so a line sung or chanted a few times -must not trigger it. The repeats are dropped, leaving one copy, and the run -returns `TRANSCRIBE_ERR_OUTPUT_REPETITION` with a `WARN`. Like -`OUTPUT_TRUNCATED` it is result-bearing: the partial transcript is readable -and `transcribe_was_truncated()` is true. - -A decode that runs out of budget has already failed, so the cleanup bar there -is lower: if the partial ends in a block of up to 256 tokens repeated at -least 3 times over at least 32 tokens, the repeats are dropped before the transcript is returned -(still `OUTPUT_TRUNCATED`, with a `WARN` saying how many tokens were dropped). - -`whisper` is excluded because it recovers from loops with its own temperature -fallback. `voxtral_realtime` is excluded because it emits one token per audio -frame, so its padding tokens repeat through any silence. Set -`TRANSCRIBE_NO_REPETITION_GUARD=1` to turn off both the stop and the -budget-stop cleanup, e.g. for byte-exact reference parity. +have some protection against this, so that you don't infinitely decode +on sequences which are identical and are obviously looping ### 3. Soft window — warn and proceed