diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index 068cc6b8..3187ee24 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, @@ -1128,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/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..ce2c8b03 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,11 +67,37 @@ 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) } } + /// 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/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 5d6927fa..848a6257 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, 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 e53123f8..499737c2 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.) @@ -94,6 +95,12 @@ 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`) +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 | Families | Window | Behavior | @@ -165,12 +172,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_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 @@ -179,8 +187,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 d15a9bde..ae05c95a 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -269,10 +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. 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), @@ -293,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; /* @@ -1678,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). @@ -1688,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 / @@ -1700,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/canary/model.cpp b/src/arch/canary/model.cpp index 3581b7b4..4676df1c 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 @@ -1280,6 +1282,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->mark_repetition_stop(); + repeating = true; + break; + } } } } else { @@ -1352,6 +1359,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->mark_repetition_stop(); + repeating = true; + break; + } } } } @@ -1360,13 +1372,15 @@ 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 " "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 +1388,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(); } // =========================================================================== @@ -1489,21 +1503,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, @@ -1815,11 +1816,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 46a2c08d..e317560b 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 @@ -1156,6 +1157,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)); @@ -1200,17 +1202,24 @@ 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->mark_repetition_stop(); + 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 " "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) { @@ -1264,7 +1273,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 @@ -1355,21 +1364,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 @@ -1666,7 +1662,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 7d593af5..f86cc016 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->mark_repetition_stop(); + break; + } } } } else { @@ -1298,6 +1303,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->mark_repetition_stop(); + break; + } } } } @@ -1314,6 +1323,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, @@ -1324,7 +1336,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(); } // =========================================================================== @@ -1446,21 +1458,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, @@ -1746,7 +1745,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 301955d0..21970c98 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 @@ -874,6 +875,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; @@ -905,20 +907,27 @@ 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->mark_repetition_stop(); + 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 " "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 +952,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 @@ -1050,21 +1059,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 @@ -1374,11 +1370,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 f51d195c..3d621fbf 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 @@ -1239,11 +1240,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->mark_repetition_stop(); + repeating = true; + break; + } const int32_t pos = T_prompt + step_i; // RoPE position const int64_t kv_idx = pos; // KV write row @@ -1279,14 +1286,15 @@ 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 " "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. @@ -1302,7 +1310,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 @@ -1455,21 +1463,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 @@ -1814,11 +1809,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 22e4bebc..11b7eb36 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->mark_repetition_stop(); + 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->mark_repetition_stop(); + 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 " @@ -748,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; @@ -778,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 + @@ -856,21 +869,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 +1056,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 +1069,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); } @@ -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 68626064..63e4febc 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" @@ -823,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()) { @@ -843,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; @@ -977,6 +988,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,20 +1006,26 @@ 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", stop_log_level)) { + cc->mark_repetition_stop(); + 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, + 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", stop_log_level); } cc->t_decode_us += ggml_time_us() - t_decode_start; @@ -1126,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 @@ -1193,7 +1211,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. @@ -1252,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; @@ -1267,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; @@ -1616,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; } @@ -1734,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 @@ -1801,21 +1828,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 +2025,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 +2127,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; + } } } } @@ -2119,12 +2139,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 +2181,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] ? 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 fb67560e..2cfb8a4f 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 @@ -922,6 +923,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)); @@ -970,15 +972,21 @@ 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->mark_repetition_stop(); + 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())); + transcribe::trim_repetition_at_budget_stop(generated_ids, "moss run"); } if (!generated_ids.empty() && generated_ids.back() == eos_id) { generated_ids.pop_back(); @@ -1024,7 +1032,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(); } // --------------------------------------------------------------------------- @@ -1054,30 +1062,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, @@ -1341,7 +1327,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 53e02b61..e8e62204 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 @@ -933,6 +934,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(); @@ -970,19 +972,26 @@ 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->mark_repetition_stop(); + 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 " "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 @@ -1082,7 +1091,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(); } // =========================================================================== @@ -1472,21 +1481,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 @@ -1692,7 +1688,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 f817d8d5..f5f973b0 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" @@ -947,6 +948,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; @@ -980,18 +982,25 @@ 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->mark_repetition_stop(); + 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 " "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) { @@ -1022,7 +1031,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(); } // --------------------------------------------------------------------------- @@ -1039,21 +1048,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, @@ -1550,11 +1546,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 e3736e57..284aa780 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 @@ -891,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. @@ -899,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); @@ -963,7 +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 || 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; @@ -976,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 - // (generation budget or KV window). `finished` is set on every stop - // reason, so it can't discriminate; the signal is the last sampled token: + // 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 (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 7d0b384f..b2129d86 100644 --- a/src/causal_lm/causal_lm.h +++ b/src/causal_lm/causal_lm.h @@ -333,12 +333,17 @@ 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 // and allocated on `sched`. Returns TRANSCRIBE_ERR_ABORTED on abort, // TRANSCRIBE_ERR_GGUF on a compute failure, else TRANSCRIBE_OK. +// +// 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, ggml_backend_sched_t sched, const StepBatchedIO & io, @@ -348,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 7ab16f14..5fa509b6 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,46 @@ 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; + 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 || 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) { + 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, @@ -210,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()); @@ -225,6 +266,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]) { @@ -350,6 +392,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; + } } } } @@ -359,13 +405,22 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * } // 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. - 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; + // 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 (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 961a64ae..cb2586d6 100644 --- a/src/transcribe-batch-util.h +++ b/src/transcribe-batch-util.h @@ -102,6 +102,22 @@ 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, 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, + const float * const * pcm, + const int * n_samples, + int n, + const RunOneFn & run_one); + // --------------------------------------------------------------------------- // Batched encoder-decoder greedy step loop (cohere / canary / moonshine) // @@ -134,18 +150,22 @@ 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 // 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) or the context window -// (max_n_kv) BEFORE emitting eos_id (transcript truncated), else 0. Lets a -// family report per-utterance TRANSCRIBE_ERR_OUTPUT_TRUNCATED from run_batch. +// stop_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, @@ -158,7 +178,7 @@ transcribe_status run_batched_encdec_step_loop(transcribe_session * int n_batch, const std::vector & 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 new file mode 100644 index 00000000..d6d37d6e --- /dev/null +++ b/src/transcribe-repetition-guard.h @@ -0,0 +1,149 @@ +// 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. +// +// 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 +#include + +namespace transcribe { + +// 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, 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; + } + 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 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, + transcribe_log_level level = TRANSCRIBE_LOG_LEVEL_WARN) { + 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(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())); + 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 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_bar); + if (block == 0) { + return; + } + ids.resize(static_cast(trim_repeating_tail(ids.data(), n, block))); + 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 +// 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 + 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 dfeb5fb4..2acf9a60 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" @@ -155,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"; } @@ -1826,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; @@ -2211,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; @@ -2352,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 @@ -2384,36 +2390,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/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/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 new file mode 100644 index 00000000..a4c7f255 --- /dev/null +++ b/tests/repetition_guard_unit.cpp @@ -0,0 +1,232 @@ +// 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())); +} + +int budget_tail_block(const std::vector & ids) { + return transcribe::repeating_tail_block(ids.data(), static_cast(ids.size()), transcribe::k_budget_trim_bar); +} + +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); + + // 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); + 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("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); + + // 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); + 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. + { + 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); + } + + // 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; + } + + // 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 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 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. + { + 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 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); + 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 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)); + expect("budget stop leaves non-repeating output alone", budget_trimmed(clean) == clean); + } + + 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..e3bf2150 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,118 @@ void test_release_scratch_after_run_and_batch() { g_run_throw = false; } +// --------------------------------------------------------------------------- +// 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 { + +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 repeat = pcm[0] > 1.5f; + const bool truncate = pcm[0] > 0.5f && !repeat; + session->clear_result(); + session->full_text = repeat ? "looped" : truncate ? "partial" : "complete"; + session->has_result = true; + if (repeat) { + session->mark_repetition_stop(); + } else if (truncate) { + session->was_truncated = true; + } + return session->truncation_status(); +} + +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 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_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() { + // 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();