From 986757bc508524c16f7fd00a3700012120a9d8d2 Mon Sep 17 00:00:00 2001 From: Yuku Kotani Date: Tue, 25 Aug 2026 17:29:59 +0900 Subject: [PATCH] feat: add recognition context for Qwen3-ASR --- bindings/python/README.md | 7 ++ .../python/src/transcribe_cpp/__init__.py | 24 +++-- .../python/src/transcribe_cpp/_generated.py | 7 +- bindings/python/tests/test_errors.py | 17 +++- bindings/rust/sys/src/transcribe_sys.rs | 10 +- bindings/rust/transcribe-cpp/README.md | 13 +++ bindings/rust/transcribe-cpp/src/session.rs | 14 ++- bindings/rust/transcribe-cpp/src/types.rs | 3 + .../rust/transcribe-cpp/tests/no_model.rs | 5 + bindings/swift/README.md | 16 +++ .../swift/Sources/TranscribeCpp/ABIHash.swift | 2 +- .../swift/Sources/TranscribeCpp/Options.swift | 20 ++-- .../TranscribeCppTests/NoModelTests.swift | 13 ++- bindings/typescript/README.md | 9 ++ bindings/typescript/src/_generated.ts | 7 +- bindings/typescript/src/index.ts | 3 + bindings/typescript/src/types.ts | 7 +- bindings/typescript/test/batch.test.mjs | 5 +- bindings/typescript/test/streaming.test.mjs | 5 +- bindings/typescript/test/transcribe.test.mjs | 5 +- docs/input-limits.md | 7 ++ docs/models/qwen3-asr.md | 13 ++- examples/cli/main.cpp | 14 +++ include/transcribe.abihash | 2 +- include/transcribe.h | 24 ++++- src/arch/qwen3_asr/capabilities.cpp | 6 +- src/arch/qwen3_asr/model.cpp | 58 ++++++++--- src/arch/qwen3_asr/qwen3_asr.h | 11 +++ src/transcribe-session.h | 3 +- src/transcribe.cpp | 59 ++++++++--- tests/CMakeLists.txt | 23 +++++ tests/api_smoke.c | 9 ++ tests/cli_context_model_smoke.cmake | 24 +++++ tests/objc_context_header_smoke.m | 14 +++ tests/qwen3_asr_bpe_parity.cpp | 28 ++++++ tests/qwen3_asr_e2e_smoke.cpp | 47 +++++++++ tests/qwen3_asr_smoke.cpp | 20 ++++ tests/run_dispatch_unit.cpp | 99 ++++++++++++++++++- tests/stream_dispatch_unit.cpp | 35 ++++++- 39 files changed, 620 insertions(+), 68 deletions(-) create mode 100644 tests/cli_context_model_smoke.cmake create mode 100644 tests/objc_context_header_smoke.m diff --git a/bindings/python/README.md b/bindings/python/README.md index 094bb9cb..ad757550 100644 --- a/bindings/python/README.md +++ b/bindings/python/README.md @@ -43,6 +43,13 @@ and the one-shot `transcribe()` helper. result = session.run(pcm, pnc="off", itn="on") ``` +Models advertising `model.supports("context")` accept best-effort recognition +background text for names and terminology. It is not an instruction prompt. + +```python +result = session.run(pcm, context="Vocabulary: GGUF, ggml, Qwen3-ASR") +``` + Streaming models expose incremental transcription with committed/tentative text views — see `examples/stream_wav.py`: diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index e52adb8c..ed2474ac 100644 --- a/bindings/python/src/transcribe_cpp/__init__.py +++ b/bindings/python/src/transcribe_cpp/__init__.py @@ -60,7 +60,7 @@ CommitPolicy = Literal["auto", "on_finalize", "stable_prefix"] Feature = Literal[ "initial_prompt", "temperature_fallback", "long_form", - "cancellation", "pnc", "itn", "diarization", + "cancellation", "pnc", "itn", "diarization", "context", ] __all__ = [ @@ -248,6 +248,7 @@ "pnc": _generated.TRANSCRIBE_FEATURE_PNC, "itn": _generated.TRANSCRIBE_FEATURE_ITN, "diarization": _generated.TRANSCRIBE_FEATURE_DIARIZATION, + "context": _generated.TRANSCRIBE_FEATURE_CONTEXT, } @@ -653,7 +654,7 @@ def _stream_update_from(u) -> StreamUpdate: def _build_run_params(task, language, target_language, timestamps, keep_special_tags, spec_k_drafts, diarize="default", - pnc="default", itn="default"): + pnc="default", itn="default", context=None): if not isinstance(spec_k_drafts, int) or spec_k_drafts < -1: raise InvalidArgument( f"spec_k_drafts must be -1 (family default), 0 (disabled), or a " @@ -668,6 +669,7 @@ def _build_run_params(task, language, target_language, timestamps, params.diarize = _enum(_DIARIZE, diarize, "diarize") params.language = language.encode("utf-8") if language else None params.target_language = target_language.encode("utf-8") if target_language else None + params.context = context.encode("utf-8") if context else None params.keep_special_tags = keep_special_tags params.spec_k_drafts = spec_k_drafts return params @@ -1095,6 +1097,7 @@ def _resolve_family(self, family: "FamilyExtension", slot: str): def run(self, pcm: PCMLike, *, task: Task = "transcribe", language: str | None = None, target_language: str | None = None, + context: str | None = None, timestamps: Timestamps = "auto", pnc: Pnc = "default", itn: Itn = "default", @@ -1104,8 +1107,10 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", family: FamilyExtension | None = None) -> Result: """Transcribe 16 kHz mono float32 PCM and return a materialized Result. - ``pnc`` controls punctuation/capitalization and ``itn`` controls - inverse text normalization on models advertising those features. + ``context`` supplies best-effort recognition background text on models + advertising the ``context`` feature. ``pnc`` controls punctuation and + capitalization; ``itn`` controls inverse text normalization on models + advertising those features. ``family`` is an optional family-specific extension (e.g. WhisperRunOptions) carrying per-run knobs for models that accept it. ``spec_k_drafts`` tunes speculative decoding on models whose @@ -1118,7 +1123,7 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", self._cancel.clear() array, n_samples = _pcm_to_carray(pcm) params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, spec_k_drafts, diarize, pnc, itn) + keep_special_tags, spec_k_drafts, diarize, pnc, itn, context) ext = self._resolve_family(family, "run") if family is not None else None if ext is not None: params.family = ctypes.cast( @@ -1136,6 +1141,7 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", language: str | None = None, target_language: str | None = None, + context: str | None = None, timestamps: Timestamps = "auto", pnc: Pnc = "default", itn: Itn = "default", @@ -1176,7 +1182,7 @@ def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", counts[k] = n params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, spec_k_drafts, diarize, pnc, itn) + keep_special_tags, spec_k_drafts, diarize, pnc, itn, context) ext = self._resolve_family(family, "run") if family is not None else None if ext is not None: params.family = ctypes.cast( @@ -1227,6 +1233,7 @@ def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", def stream(self, *, task: Task = "transcribe", language: str | None = None, target_language: str | None = None, timestamps: Timestamps = "none", + context: str | None = None, pnc: Pnc = "default", itn: Itn = "default", diarize: Diarize = "default", keep_special_tags: bool = False, commit_policy: CommitPolicy = "auto", @@ -1243,7 +1250,7 @@ def stream(self, *, task: Task = "transcribe", language: str | None = None, # spec_k_drafts is an offline-decode knob; streaming always uses the # family default (-1). run_params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, -1, diarize, pnc, itn) + keep_special_tags, -1, diarize, pnc, itn, context) sp = _StreamParams() _lib.transcribe_stream_params_init(_byref(sp)) sp.commit_policy = _enum(_COMMIT_POLICIES, commit_policy, "commit_policy") @@ -1489,6 +1496,7 @@ def transcribe( task: Task = "transcribe", language: str | None = None, target_language: str | None = None, + context: str | None = None, timestamps: Timestamps = "auto", pnc: Pnc = "default", itn: Itn = "default", @@ -1507,7 +1515,7 @@ def transcribe( ``family`` / ``spec_k_drafts`` pass through to :meth:`Session.run`. """ session_opts = dict(n_threads=n_threads, kv_type=kv_type, n_ctx=n_ctx) - run_opts = dict(task=task, language=language, target_language=target_language, + run_opts = dict(task=task, language=language, target_language=target_language, context=context, timestamps=timestamps, pnc=pnc, itn=itn, diarize=diarize, keep_special_tags=keep_special_tags, spec_k_drafts=spec_k_drafts, family=family) diff --git a/bindings/python/src/transcribe_cpp/_generated.py b/bindings/python/src/transcribe_cpp/_generated.py index ac9763c2..e4e8e9de 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 = "c5cab57f2043f3c9" # === enum constants === TRANSCRIBE_OK = 0 @@ -95,6 +95,7 @@ TRANSCRIBE_FEATURE_PNC = 4 TRANSCRIBE_FEATURE_ITN = 5 TRANSCRIBE_FEATURE_DIARIZATION = 6 +TRANSCRIBE_FEATURE_CONTEXT = 7 TRANSCRIBE_STREAM_IDLE = 0 TRANSCRIBE_STREAM_ACTIVE = 1 TRANSCRIBE_STREAM_FINISHED = 2 @@ -167,7 +168,7 @@ class transcribe_whisper_chunk_trace(_c.Structure): transcribe_device_info._fields_ = [("struct_size", _c.c_uint64), ("name", _c.c_char_p), ("description", _c.c_char_p), ("kind", _c.c_char_p), ("device_id", _c.c_char_p), ("memory_total", _c.c_uint64), ("memory_free", _c.c_uint64), ("device_type", _c.c_int)] transcribe_model_load_params._fields_ = [("struct_size", _c.c_uint64), ("backend", _c.c_int), ("device", _c.c_void_p)] transcribe_session_params._fields_ = [("struct_size", _c.c_uint64), ("n_threads", _c.c_int), ("kv_type", _c.c_int), ("n_ctx", _c.c_int32)] -transcribe_run_params._fields_ = [("struct_size", _c.c_uint64), ("task", _c.c_int), ("timestamps", _c.c_int), ("pnc", _c.c_int), ("itn", _c.c_int), ("diarize", _c.c_int), ("language", _c.c_char_p), ("target_language", _c.c_char_p), ("keep_special_tags", _c.c_bool), ("family", _c.POINTER(transcribe_ext)), ("spec_k_drafts", _c.c_int32)] +transcribe_run_params._fields_ = [("struct_size", _c.c_uint64), ("task", _c.c_int), ("timestamps", _c.c_int), ("pnc", _c.c_int), ("itn", _c.c_int), ("diarize", _c.c_int), ("language", _c.c_char_p), ("target_language", _c.c_char_p), ("keep_special_tags", _c.c_bool), ("family", _c.POINTER(transcribe_ext)), ("spec_k_drafts", _c.c_int32), ("context", _c.c_char_p)] transcribe_capabilities._fields_ = [("struct_size", _c.c_uint64), ("native_sample_rate", _c.c_int32), ("n_languages", _c.c_int), ("languages", _c.POINTER(_c.c_char_p)), ("max_timestamp_kind", _c.c_int), ("supports_language_detect", _c.c_bool), ("supports_translate", _c.c_bool), ("supports_streaming", _c.c_bool), ("supports_spec_decode", _c.c_bool), ("max_audio_ms", _c.c_int64), ("n_translate_target_languages", _c.c_int), ("translate_target_languages", _c.POINTER(_c.c_char_p))] transcribe_session_limits._fields_ = [("struct_size", _c.c_uint64), ("effective_n_ctx", _c.c_int32), ("effective_max_audio_ms", _c.c_int64), ("max_kv_bytes", _c.c_int64)] transcribe_stream_params._fields_ = [("struct_size", _c.c_uint64), ("family", _c.POINTER(transcribe_ext)), ("commit_policy", _c.c_int), ("stable_prefix_agreement_n", _c.c_uint32)] @@ -212,7 +213,7 @@ class transcribe_whisper_chunk_trace(_c.Structure): 'transcribe_device_info': {'size': 64, 'align': 8, 'offsets': {'struct_size': 0, 'name': 8, 'description': 16, 'kind': 24, 'device_id': 32, 'memory_total': 40, 'memory_free': 48, 'device_type': 56}}, 'transcribe_model_load_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'backend': 8, 'device': 16}}, 'transcribe_session_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'n_threads': 8, 'kv_type': 12, 'n_ctx': 16}}, - 'transcribe_run_params': {'size': 72, 'align': 8, 'offsets': {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64}}, + 'transcribe_run_params': {'size': 80, 'align': 8, 'offsets': {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64, 'context': 72}}, 'transcribe_capabilities': {'size': 56, 'align': 8, 'offsets': {'struct_size': 0, 'native_sample_rate': 8, 'n_languages': 12, 'languages': 16, 'max_timestamp_kind': 24, 'supports_language_detect': 28, 'supports_translate': 29, 'supports_streaming': 30, 'supports_spec_decode': 31, 'max_audio_ms': 32, 'n_translate_target_languages': 40, 'translate_target_languages': 48}}, 'transcribe_session_limits': {'size': 32, 'align': 8, 'offsets': {'struct_size': 0, 'effective_n_ctx': 8, 'effective_max_audio_ms': 16, 'max_kv_bytes': 24}}, 'transcribe_stream_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'family': 8, 'commit_policy': 16, 'stable_prefix_agreement_n': 20}}, diff --git a/bindings/python/tests/test_errors.py b/bindings/python/tests/test_errors.py index 47066cd3..2f1ad3fa 100644 --- a/bindings/python/tests/test_errors.py +++ b/bindings/python/tests/test_errors.py @@ -137,10 +137,25 @@ def test_pnc_and_itn_modes_map_to_native_run_params(): ) +def test_context_maps_to_native_run_params_without_normalization(): + from transcribe_cpp import _build_run_params + + params = _build_run_params( + "transcribe", None, None, "none", False, -1, + context=" Glossary: GGUF\n\u65e5\u672c\u8a9e ", + ) + assert params.context == " Glossary: GGUF\n\u65e5\u672c\u8a9e ".encode() + + empty = _build_run_params( + "transcribe", None, None, "none", False, -1, context="", + ) + assert empty.context is None + + def test_public_run_surfaces_cover_every_generic_option(): common = { "task", "language", "target_language", "timestamps", "pnc", "itn", - "diarize", "keep_special_tags", "family", + "diarize", "keep_special_tags", "family", "context", } expected = { t.Session.run: common | {"spec_k_drafts"}, diff --git a/bindings/rust/sys/src/transcribe_sys.rs b/bindings/rust/sys/src/transcribe_sys.rs index 363cd6e3..1a33de79 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 = c5cab57f2043f3c9 /// 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 = "c5cab57f2043f3c9"; /* automatically generated by rust-bindgen 0.72.1 */ @@ -346,10 +346,11 @@ pub struct transcribe_run_params { pub keep_special_tags: bool, pub family: *const transcribe_ext, pub spec_k_drafts: i32, + pub context: *const ::std::os::raw::c_char, } #[allow(clippy::unnecessary_operation, clippy::identity_op)] const _: () = { - ["Size of transcribe_run_params"][::std::mem::size_of::() - 72usize]; + ["Size of transcribe_run_params"][::std::mem::size_of::() - 80usize]; ["Alignment of transcribe_run_params"] [::std::mem::align_of::() - 8usize]; ["Offset of field: transcribe_run_params::struct_size"] @@ -374,6 +375,8 @@ const _: () = { [::std::mem::offset_of!(transcribe_run_params, family) - 56usize]; ["Offset of field: transcribe_run_params::spec_k_drafts"] [::std::mem::offset_of!(transcribe_run_params, spec_k_drafts) - 64usize]; + ["Offset of field: transcribe_run_params::context"] + [::std::mem::offset_of!(transcribe_run_params, context) - 72usize]; }; unsafe extern "C" { pub fn transcribe_run_params_init(params: *mut transcribe_run_params); @@ -441,6 +444,7 @@ impl transcribe_feature { pub const TRANSCRIBE_FEATURE_PNC: transcribe_feature = transcribe_feature(4); pub const TRANSCRIBE_FEATURE_ITN: transcribe_feature = transcribe_feature(5); pub const TRANSCRIBE_FEATURE_DIARIZATION: transcribe_feature = transcribe_feature(6); + pub const TRANSCRIBE_FEATURE_CONTEXT: transcribe_feature = transcribe_feature(7); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] diff --git a/bindings/rust/transcribe-cpp/README.md b/bindings/rust/transcribe-cpp/README.md index 95f9f969..f2747b5a 100644 --- a/bindings/rust/transcribe-cpp/README.md +++ b/bindings/rust/transcribe-cpp/README.md @@ -50,6 +50,19 @@ let result = session.run(&pcm, &options)?; # Ok::<(), transcribe_cpp::Error>(()) ``` +Models advertising `Feature::Context` accept best-effort recognition background +text for names and terminology. It is not an instruction prompt. + +```rust +use transcribe_cpp::RunOptions; +let options = RunOptions { + context: Some("Vocabulary: GGUF, ggml, Qwen3-ASR".into()), + ..Default::default() +}; +let result = session.run(&pcm, &options)?; +# Ok::<(), transcribe_cpp::Error>(()) +``` + Streaming exposes both UI-stable text and a fully materialized structured snapshot: diff --git a/bindings/rust/transcribe-cpp/src/session.rs b/bindings/rust/transcribe-cpp/src/session.rs index 220f2b63..22b9a6da 100644 --- a/bindings/rust/transcribe-cpp/src/session.rs +++ b/bindings/rust/transcribe-cpp/src/session.rs @@ -37,6 +37,8 @@ pub struct RunOptions { pub language: Option, /// Target language for translation, or `None`. pub target_language: Option, + /// Best-effort recognition background text (names, jargon, terminology). + pub context: Option, /// Keep special vocab tags (e.g. `<|...|>`) in the returned text. pub keep_special_tags: bool, /// Speculative-decode draft length. `-1` = family default, `0` = disabled. @@ -55,6 +57,7 @@ impl Default for RunOptions { diarize: Diarize::Default, language: None, target_language: None, + context: None, keep_special_tags: false, spec_k_drafts: -1, family: None, @@ -148,7 +151,7 @@ impl Session { /// On an aborted or truncated decode the partial transcript is preserved /// on the returned [`Error::Aborted`] / [`Error::OutputTruncated`]. pub fn run(&mut self, pcm: &[f32], options: &RunOptions) -> Result { - let (params, _lang, _target, _family) = build_run_params(options)?; + let (params, _lang, _target, _context, _family) = build_run_params(options)?; let n = clamp_len(pcm.len())?; // The compute path is serialized per model; hold the lock for the native @@ -191,7 +194,7 @@ impl Session { pcms: &[&[f32]], options: &RunOptions, ) -> Result>> { - let (params, _lang, _target, _family) = build_run_params(options)?; + let (params, _lang, _target, _context, _family) = build_run_params(options)?; let ptrs: Vec<*const f32> = pcms.iter().map(|p| p.as_ptr()).collect(); let lens: Vec = pcms .iter() @@ -267,7 +270,7 @@ impl Session { /// Dropping the returned `Stream` abandons it and returns the session to /// idle. pub fn stream(&mut self, run: &RunOptions, stream: &StreamOptions) -> Result> { - let (run_params, _lang, _target, _family) = build_run_params(run)?; + let (run_params, _lang, _target, _context, _family) = build_run_params(run)?; let (stream_params, _stream_family) = build_stream_params(stream); { // Claim the model's compute lease for the whole stream lifetime: a @@ -422,6 +425,7 @@ type RunParamsBundle = ( sys::transcribe_run_params, Option, Option, + Option, Option, ); @@ -442,8 +446,10 @@ fn build_run_params(o: &RunOptions) -> Result { let lang = o.language.as_deref().map(CString::new).transpose()?; let target = o.target_language.as_deref().map(CString::new).transpose()?; + let context = o.context.as_deref().map(CString::new).transpose()?; params.language = lang.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()); params.target_language = target.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()); + params.context = context.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()); let family = o .family @@ -452,7 +458,7 @@ fn build_run_params(o: &RunOptions) -> Result { .transpose()?; params.family = family.as_ref().map_or(std::ptr::null(), |f| f.ext_ptr()); - Ok((params, lang, target, family)) + Ok((params, lang, target, context, family)) } /// PCM/utterance lengths cross the ABI as `int`; reject anything that overflows. diff --git a/bindings/rust/transcribe-cpp/src/types.rs b/bindings/rust/transcribe-cpp/src/types.rs index 02089439..86bb4bfc 100644 --- a/bindings/rust/transcribe-cpp/src/types.rs +++ b/bindings/rust/transcribe-cpp/src/types.rs @@ -211,6 +211,8 @@ pub enum Feature { Itn, /// Produces structured speaker attribution. Diarization, + /// Accepts best-effort recognition background text. + Context, } impl Feature { @@ -224,6 +226,7 @@ impl Feature { Feature::Pnc => F::TRANSCRIBE_FEATURE_PNC, Feature::Itn => F::TRANSCRIBE_FEATURE_ITN, Feature::Diarization => F::TRANSCRIBE_FEATURE_DIARIZATION, + Feature::Context => F::TRANSCRIBE_FEATURE_CONTEXT, } } } diff --git a/bindings/rust/transcribe-cpp/tests/no_model.rs b/bindings/rust/transcribe-cpp/tests/no_model.rs index 89a750be..3f78c8c3 100644 --- a/bindings/rust/transcribe-cpp/tests/no_model.rs +++ b/bindings/rust/transcribe-cpp/tests/no_model.rs @@ -44,10 +44,15 @@ fn generic_text_control_options_round_trip() { let options = RunOptions { pnc: Pnc::Off, itn: Itn::On, + context: Some(" Glossary: GGUF\n\u{65e5}\u{672c}\u{8a9e} ".into()), ..Default::default() }; assert_eq!(options.pnc, Pnc::Off); assert_eq!(options.itn, Itn::On); + assert_eq!( + options.context.as_deref(), + Some(" Glossary: GGUF\n\u{65e5}\u{672c}\u{8a9e} ") + ); } #[test] diff --git a/bindings/swift/README.md b/bindings/swift/README.md index fd34d7b6..64d0b32a 100644 --- a/bindings/swift/README.md +++ b/bindings/swift/README.md @@ -72,6 +72,14 @@ let options = RunOptions(pnc: .off, itn: .on) let transcript = try session.run(pcm, options: options) ``` +Models advertising `model.supports(.context)` accept best-effort recognition +background text for names and terminology. It is not an instruction prompt. + +```swift +let options = RunOptions(context: "Vocabulary: GGUF, ggml, Qwen3-ASR") +let transcript = try session.run(pcm, options: options) +``` + Streaming models expose committed/tentative text for UI display: ```swift @@ -135,3 +143,11 @@ task cancellation when no custom token is installed. The xcframework also exposes the raw C module as `CTranscribe`. Objective-C and C++ callers use the bundled C headers directly, for example `#import `. + +```objc +#import + +transcribe_run_params params; +transcribe_run_params_init(¶ms); +params.context = "Vocabulary: GGUF, ggml, Qwen3-ASR"; +``` diff --git a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift index 5ea6cbc1..e97d708b 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 = "c5cab57f2043f3c9" /// 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/Options.swift b/bindings/swift/Sources/TranscribeCpp/Options.swift index 10ff81f3..679b8bcd 100644 --- a/bindings/swift/Sources/TranscribeCpp/Options.swift +++ b/bindings/swift/Sources/TranscribeCpp/Options.swift @@ -81,7 +81,7 @@ public enum Diarize: Sendable { } public enum Feature: Sendable { - case initialPrompt, temperatureFallback, longForm, cancellation, pnc, itn, diarization + case initialPrompt, temperatureFallback, longForm, cancellation, pnc, itn, diarization, context var cValue: transcribe_feature { switch self { case .initialPrompt: return TRANSCRIBE_FEATURE_INITIAL_PROMPT @@ -91,6 +91,7 @@ public enum Feature: Sendable { case .pnc: return TRANSCRIBE_FEATURE_PNC case .itn: return TRANSCRIBE_FEATURE_ITN case .diarization: return TRANSCRIBE_FEATURE_DIARIZATION + case .context: return TRANSCRIBE_FEATURE_CONTEXT } } } @@ -130,6 +131,8 @@ public struct RunOptions: Sendable { public var language: String? /// Target language for translation; `nil` otherwise. public var targetLanguage: String? + /// Best-effort recognition background text (names, jargon, terminology). + public var context: String? /// Keep special `<|...|>` tags in the returned text. public var keepSpecialTags: Bool /// Speculative-decode draft length: -1 = family default, 0 = disabled. @@ -148,6 +151,7 @@ public struct RunOptions: Sendable { diarize: Diarize = .default, language: String? = nil, targetLanguage: String? = nil, + context: String? = nil, keepSpecialTags: Bool = false, specKDrafts: Int32 = -1, family: RunExtension? = nil @@ -159,14 +163,15 @@ public struct RunOptions: Sendable { self.diarize = diarize self.language = language self.targetLanguage = targetLanguage + self.context = context self.keepSpecialTags = keepSpecialTags self.specKDrafts = specKDrafts self.family = family } /// Materialize a `transcribe_run_params` and run `body` with a pointer to - /// it. The `language` / `target_language` C strings are kept alive for the - /// duration of `body` (the C side copies them before returning). + /// it. The `language` / `target_language` / `context` C strings are kept + /// alive for the duration of `body` (the C side copies them before returning). func withCParams(_ body: (UnsafePointer) throws -> R) rethrows -> R { var params = transcribe_run_params() transcribe_run_params_init(¶ms) @@ -181,9 +186,12 @@ public struct RunOptions: Sendable { params.language = lang return try withOptionalCString(targetLanguage) { tgt in params.target_language = tgt - return try withRunExtension(family) { ext in - params.family = ext - return try withUnsafePointer(to: ¶ms) { try body($0) } + return try withOptionalCString(context) { context in + params.context = context + return try withRunExtension(family) { ext in + params.family = ext + return try withUnsafePointer(to: ¶ms) { try body($0) } + } } } } diff --git a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift index 9de46724..01c13b19 100644 --- a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift +++ b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift @@ -90,7 +90,12 @@ final class NoModelTests: XCTestCase { // not shadow Swift's concurrency `Task`). Lock the public name + `task:` // option here so an accidental rename is caught without a model. func testTranscriptionTaskOptionRoundTrips() { - let translate = RunOptions(task: .translate, pnc: .off, itn: .on, diarize: .on) + let translate = RunOptions( + task: .translate, + pnc: .off, + itn: .on, + diarize: .on, + context: " Glossary: GGUF\n\u{65E5}\u{672C}\u{8A9E} ") guard case .translate = translate.task else { return XCTFail("task option did not round-trip to .translate") } @@ -99,6 +104,12 @@ final class NoModelTests: XCTestCase { guard case .off = translate.pnc else { return XCTFail("Pnc.off") } guard case .on = translate.itn else { return XCTFail("Itn.on") } guard case .on = translate.diarize else { return XCTFail("Diarize.on") } + XCTAssertEqual(translate.context, " Glossary: GGUF\n\u{65E5}\u{672C}\u{8A9E} ") + translate.withCParams { params in + XCTAssertEqual( + String(cString: params.pointee.context), + " Glossary: GGUF\n\u{65E5}\u{672C}\u{8A9E} ") + } } func testJunkFileIsModelLoadError() throws { diff --git a/bindings/typescript/README.md b/bindings/typescript/README.md index 416a45ee..c0173dc3 100644 --- a/bindings/typescript/README.md +++ b/bindings/typescript/README.md @@ -46,6 +46,15 @@ and streams. const result = await model.transcribe(pcm, { pnc: "off", itn: "on" }); ``` +Models advertising `model.supports("context")` accept best-effort recognition +background text for names and terminology. It is not an instruction prompt. + +```ts +const result = await model.transcribe(pcm, { + context: "Vocabulary: GGUF, ggml, Qwen3-ASR", +}); +``` + ### Streaming ```ts diff --git a/bindings/typescript/src/_generated.ts b/bindings/typescript/src/_generated.ts index fff2ca59..826e6b9f 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 = "c5cab57f2043f3c9"; // === enum constants === export const TRANSCRIBE_OK = 0; @@ -93,6 +93,7 @@ export const TRANSCRIBE_FEATURE_CANCELLATION = 3; export const TRANSCRIBE_FEATURE_PNC = 4; export const TRANSCRIBE_FEATURE_ITN = 5; export const TRANSCRIBE_FEATURE_DIARIZATION = 6; +export const TRANSCRIBE_FEATURE_CONTEXT = 7; export const TRANSCRIBE_STREAM_IDLE = 0; export const TRANSCRIBE_STREAM_ACTIVE = 1; export const TRANSCRIBE_STREAM_FINISHED = 2; @@ -121,7 +122,7 @@ export const STRUCT_LAYOUT: Record = { 'transcribe_device_info': { size: 64, align: 8, offsets: {'struct_size': 0, 'name': 8, 'description': 16, 'kind': 24, 'device_id': 32, 'memory_total': 40, 'memory_free': 48, 'device_type': 56} }, 'transcribe_model_load_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'backend': 8, 'device': 16} }, 'transcribe_session_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'n_threads': 8, 'kv_type': 12, 'n_ctx': 16} }, - 'transcribe_run_params': { size: 72, align: 8, offsets: {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64} }, + 'transcribe_run_params': { size: 80, align: 8, offsets: {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64, 'context': 72} }, 'transcribe_capabilities': { size: 56, align: 8, offsets: {'struct_size': 0, 'native_sample_rate': 8, 'n_languages': 12, 'languages': 16, 'max_timestamp_kind': 24, 'supports_language_detect': 28, 'supports_translate': 29, 'supports_streaming': 30, 'supports_spec_decode': 31, 'max_audio_ms': 32, 'n_translate_target_languages': 40, 'translate_target_languages': 48} }, 'transcribe_session_limits': { size: 32, align: 8, offsets: {'struct_size': 0, 'effective_n_ctx': 8, 'effective_max_audio_ms': 16, 'max_kv_bytes': 24} }, 'transcribe_stream_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'family': 8, 'commit_policy': 16, 'stable_prefix_agreement_n': 20} }, @@ -166,7 +167,7 @@ export function defineTypes(koffi: any): Record { T['transcribe_device_info'] = koffi.struct({ struct_size: 'uint64_t', name: 'char *', description: 'char *', kind: 'char *', device_id: 'char *', memory_total: 'uint64_t', memory_free: 'uint64_t', device_type: 'int' }); T['transcribe_model_load_params'] = koffi.struct({ struct_size: 'uint64_t', backend: 'int', device: 'void *' }); T['transcribe_session_params'] = koffi.struct({ struct_size: 'uint64_t', n_threads: 'int', kv_type: 'int', n_ctx: 'int32_t' }); - T['transcribe_run_params'] = koffi.struct({ struct_size: 'uint64_t', task: 'int', timestamps: 'int', pnc: 'int', itn: 'int', diarize: 'int', language: 'char *', target_language: 'char *', keep_special_tags: 'bool', family: 'void *', spec_k_drafts: 'int32_t' }); + T['transcribe_run_params'] = koffi.struct({ struct_size: 'uint64_t', task: 'int', timestamps: 'int', pnc: 'int', itn: 'int', diarize: 'int', language: 'char *', target_language: 'char *', keep_special_tags: 'bool', family: 'void *', spec_k_drafts: 'int32_t', context: 'char *' }); T['transcribe_capabilities'] = koffi.struct({ struct_size: 'uint64_t', native_sample_rate: 'int32_t', n_languages: 'int', languages: 'void *', max_timestamp_kind: 'int', supports_language_detect: 'bool', supports_translate: 'bool', supports_streaming: 'bool', supports_spec_decode: 'bool', max_audio_ms: 'int64_t', n_translate_target_languages: 'int', translate_target_languages: 'void *' }); T['transcribe_session_limits'] = koffi.struct({ struct_size: 'uint64_t', effective_n_ctx: 'int32_t', effective_max_audio_ms: 'int64_t', max_kv_bytes: 'int64_t' }); T['transcribe_stream_params'] = koffi.struct({ struct_size: 'uint64_t', family: 'void *', commit_policy: 'int', stable_prefix_agreement_n: 'uint32_t' }); diff --git a/bindings/typescript/src/index.ts b/bindings/typescript/src/index.ts index 54a651e8..91e62bc6 100644 --- a/bindings/typescript/src/index.ts +++ b/bindings/typescript/src/index.ts @@ -115,6 +115,7 @@ const FEATURES: Record = { pnc: g.TRANSCRIBE_FEATURE_PNC, itn: g.TRANSCRIBE_FEATURE_ITN, diarization: g.TRANSCRIBE_FEATURE_DIARIZATION, + context: g.TRANSCRIBE_FEATURE_CONTEXT, }; // ---- helpers --------------------------------------------------------------- @@ -846,6 +847,7 @@ export class Session { if (opts.language !== undefined) p.language = opts.language; if (opts.targetLanguage !== undefined) p.target_language = opts.targetLanguage; + if (opts.context !== undefined) p.context = opts.context; if (opts.keepSpecialTags !== undefined) p.keep_special_tags = opts.keepSpecialTags; if (opts.specKDrafts !== undefined) p.spec_k_drafts = opts.specKDrafts; @@ -942,6 +944,7 @@ export class Session { task: opts.task, language: opts.language, targetLanguage: opts.targetLanguage, + context: opts.context, timestamps: opts.timestamps, pnc: opts.pnc, itn: opts.itn, diff --git a/bindings/typescript/src/types.ts b/bindings/typescript/src/types.ts index 0a439518..70768a42 100644 --- a/bindings/typescript/src/types.ts +++ b/bindings/typescript/src/types.ts @@ -16,7 +16,8 @@ export type Feature = | "cancellation" | "pnc" | "itn" - | "diarization"; + | "diarization" + | "context"; /** Mono float32 PCM at the model's native sample rate (16 kHz for v1). */ export type PcmLike = Float32Array | number[] | ArrayBuffer | Buffer; @@ -153,6 +154,8 @@ export interface TranscribeOptions { task?: Task; language?: string; targetLanguage?: string; + /** Best-effort recognition background text (names, jargon, terminology). */ + context?: string; /** Default "auto" (richest the model supports, per-family). */ timestamps?: TimestampKind; /** Punctuation and capitalization control; default preserves the family default. */ @@ -207,6 +210,8 @@ export interface StreamOptions { task?: Task; language?: string; targetLanguage?: string; + /** Best-effort recognition background text. */ + context?: string; timestamps?: TimestampKind; pnc?: Pnc; itn?: Itn; diff --git a/bindings/typescript/test/batch.test.mjs b/bindings/typescript/test/batch.test.mjs index 88a17c3e..f5f86694 100644 --- a/bindings/typescript/test/batch.test.mjs +++ b/bindings/typescript/test/batch.test.mjs @@ -7,7 +7,10 @@ modelTest("batch returns one result per utterance", MODEL, async () => { try { const s = m.createSession(); const pcm = jfk(); - const items = await s.runBatch([pcm, pcm.subarray(0, pcm.length / 2)]); + const items = await s.runBatch( + [pcm, pcm.subarray(0, pcm.length / 2)], + { context: "Vocabulary: GGUF, ggml, Qwen3-ASR" }, + ); assert.equal(items.length, 2); assert.ok(items[0].ok && /ask not/i.test(items[0].result.text)); assert.ok(items[1].ok); diff --git a/bindings/typescript/test/streaming.test.mjs b/bindings/typescript/test/streaming.test.mjs index 7c38b171..812afa61 100644 --- a/bindings/typescript/test/streaming.test.mjs +++ b/bindings/typescript/test/streaming.test.mjs @@ -7,7 +7,10 @@ modelTest("streaming commits text and finalizes", STREAMING_MODEL, async () => { try { assert.equal(m.capabilities.supportsStreaming, true); const s = m.createSession(); - const stream = await s.stream({ commitPolicy: "stable_prefix" }); + const stream = await s.stream({ + context: "Vocabulary: Kennedy, Massachusetts", + commitPolicy: "stable_prefix", + }); assert.equal(stream.state, "active"); await feedChunks(stream, jfk()); const fin = await stream.finalize(); diff --git a/bindings/typescript/test/transcribe.test.mjs b/bindings/typescript/test/transcribe.test.mjs index f1b258e5..921019ca 100644 --- a/bindings/typescript/test/transcribe.test.mjs +++ b/bindings/typescript/test/transcribe.test.mjs @@ -5,7 +5,9 @@ import { TranscribeModel } from "../dist/index.js"; modelTest("offline transcription returns text + detected language", MODEL, async () => { const m = await TranscribeModel.load(MODEL); try { - const r = await m.transcribe(jfk()); + const r = await m.transcribe(jfk(), { + context: "Vocabulary: GGUF, ggml, Qwen3-ASR", + }); assert.match(r.text, /ask not what your country/i); assert.equal(r.language, "en"); assert.equal(r.aborted, false); @@ -35,6 +37,7 @@ modelTest("capabilities + identity", MODEL, async () => { const c = m.capabilities; assert.equal(c.nativeSampleRate, 16000); assert.ok(c.languages.length > 0); + assert.equal(typeof m.supports("context"), "boolean"); assert.equal(typeof m.arch, "string"); assert.ok(m.arch.length > 0); assert.ok(m.backend.length > 0); diff --git a/docs/input-limits.md b/docs/input-limits.md index fb7a6ae0..896db8fc 100644 --- a/docs/input-limits.md +++ b/docs/input-limits.md @@ -138,6 +138,13 @@ reports the model's default-context ceiling (`n_ctx == 0`); it is not re-derived for a session that narrows `n_ctx`. A session that lowers `n_ctx` may therefore reject audio shorter than the advertised `max_audio_ms`. +Per-run prompt material also shares this window. In particular, a non-empty +`transcribe_run_params::context` on Qwen3-ASR lowers the audio-token budget for +that request. Session limits remain context-free advisories because context is +not known until run time; the exact combined audio + prompt gate is enforced by +`transcribe_run` / `transcribe_run_batch`, which return +`TRANSCRIBE_ERR_INPUT_TOO_LONG` rather than truncating context. + Encoder-bound families are different. For cohere and canary, the input-audio limit is the encoder positional table, while `n_ctx` only bounds the decoder self-KV / output budget. In those families `transcribe_session_get_limits()` diff --git a/docs/models/qwen3-asr.md b/docs/models/qwen3-asr.md index a5553c73..dcdfb50b 100644 --- a/docs/models/qwen3-asr.md +++ b/docs/models/qwen3-asr.md @@ -46,7 +46,9 @@ family. That ceiling is there to bound memory and sits far beyond any normal clip; audio past it is rejected up front with `TRANSCRIBE_ERR_INPUT_TOO_LONG` rather than silently truncated. Lowering `--n-ctx` lowers the limit (and the KV-cache footprint), and `transcribe_session_get_limits()` reports the exact -per-session value. See the [input-length contract](../input-limits.md). +context-free per-session value. Recognition context also consumes tokens from +this window; a context-heavy request can therefore reject audio below that +advisory limit. See the [input-length contract](../input-limits.md). ## Quick start @@ -58,9 +60,17 @@ cmake --build build build/bin/transcribe-cli \ -m models/qwen3-asr-0.6b/qwen3-asr-0.6b-Q8_0.gguf \ + --context "Vocabulary: GGUF, ggml, Qwen3-ASR" \ samples/jfk.wav ``` +`--context` supplies background text that can improve recognition of names, +jargon, and other ambiguous terms. It is inserted into Qwen3-ASR's system turn, +but it is not an instruction prompt: task changes, formatting, translation, +punctuation, and style control are not guaranteed. C callers set +`transcribe_run_params::context`; other bindings expose the same `context` +option. Probe `TRANSCRIBE_FEATURE_CONTEXT` when selecting the option dynamically. + The repo doesn't ship the GGUFs — pull them from the corresponding `handy-computer/-gguf` repo on Hugging Face, or convert from the upstream Qwen checkpoint via the per-variant doc's reproduction @@ -72,6 +82,7 @@ All Qwen3-ASR variants support: - **Transcription** of 16 kHz mono WAV input. - **Auto language detection** across 30 languages. +- **Recognition context** through `transcribe_run_params::context`. What's not supported (consistent across the family): translation, real-time streaming, VAD, speaker diarization, timestamps. See the diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index f4cbbd05..acad42e8 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -211,6 +211,7 @@ struct cli_args { std::string model_path; std::string language; std::string target_language; // --target-language: target lang for translation + std::string context; // --context: recognition background text std::string batch_file; // --batch: one wav path per line int batch_size = 0; // --batch-size: >1 groups utterances into // transcribe_run_batch calls (offline only). @@ -314,6 +315,7 @@ void print_usage(const char * argv0) { " --batch-jsonl output one JSON line per file (for batch)\n" " --batch-size N group N utterances into one transcribe_run_batch\n" " call (offline only; 0/1 = per-file serial loop)\n" + " --context TEXT recognition background text (names, jargon, terminology)\n" " --initial-prompt TEXT (whisper) initial prompt text for context biasing\n" " --temperature F (whisper) tier-0 sampling temperature (default 0 = greedy)\n" " --condition-on-prev-tokens (whisper) carry prev-chunk tokens across chunks\n" @@ -559,6 +561,12 @@ bool parse_args(int argc, char ** argv, cli_args & out) { } out.initial_prompt = v; out.whisper_set = true; + } else if (a == "--context") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + out.context = v; } else if (a == "--temperature") { const char * v = take_value(a.c_str()); if (!v) { @@ -843,6 +851,9 @@ int main(int argc, char ** argv) { } rp.timestamps = args.timestamps; rp.spec_k_drafts = args.spec_k_drafts; + if (!args.context.empty()) { + rp.context = args.context.c_str(); + } if (args.itn_set) { rp.itn = args.use_itn ? TRANSCRIBE_ITN_MODE_ON : TRANSCRIBE_ITN_MODE_OFF; @@ -1266,6 +1277,9 @@ int main(int argc, char ** argv) { } rp.timestamps = args.timestamps; rp.spec_k_drafts = args.spec_k_drafts; + if (!args.context.empty()) { + rp.context = args.context.c_str(); + } if (args.itn_set) { rp.itn = args.use_itn ? TRANSCRIBE_ITN_MODE_ON : TRANSCRIBE_ITN_MODE_OFF; diff --git a/include/transcribe.abihash b/include/transcribe.abihash index b0e23c5d..c774f5de 100644 --- a/include/transcribe.abihash +++ b/include/transcribe.abihash @@ -1 +1 @@ -7df72bf9e667b8c2 +c5cab57f2043f3c9 diff --git a/include/transcribe.h b/include/transcribe.h index 702d5259..47dcb82d 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -1043,8 +1043,20 @@ TRANSCRIBE_API void transcribe_session_params_init(struct transcribe_session_par * * target_language: target language for translation tasks, or NULL. * - * String-pointer lifetime (language / target_language): caller-owned, and - * the library copies what it needs before the API call returns. This holds + * context: optional UTF-8 background text used to bias recognition of + * names, jargon, and other ambiguous terms. NULL (default) + * and the empty string both mean no context. The library + * preserves non-empty text byte-for-byte: it does not trim, + * normalize, or interpret special-token-looking substrings. + * This is not an instruction prompt and does not guarantee + * task, output-format, style, punctuation, or language + * control. Probe TRANSCRIBE_FEATURE_CONTEXT before use. A + * non-empty context on an unsupported model emits a WARN, is + * ignored, and the run proceeds. In transcribe_run_batch the + * one context value is shared by every utterance. + * + * String-pointer lifetime (language / target_language / context): caller-owned, + * and the library copies what it needs before the API call returns. This holds * for transcribe_run / transcribe_run_batch (synchronous) AND for * transcribe_stream_begin: the dispatcher copies these strings into * session-owned storage at begin, so the caller may free its params — @@ -1109,6 +1121,8 @@ struct transcribe_run_params { * to know whether the field will take effect. */ int32_t spec_k_drafts; + + const char * context; }; TRANSCRIBE_API void transcribe_run_params_init(struct transcribe_run_params * params); @@ -1300,6 +1314,11 @@ TRANSCRIBE_API transcribe_status transcribe_model_get_capabilities(const struct * prompt to bias decoding. Today: whisper * only; reached via transcribe_whisper_run_ext. * + * CONTEXT The model accepts transcribe_run_params::context as + * best-effort background information for recognition. + * It does not imply instruction-following behavior. + * Today: qwen3_asr. + * * TEMPERATURE_FALLBACK The model runs a multi-tier temperature loop * with metric-driven fallback. Today: whisper. * @@ -1347,6 +1366,7 @@ typedef enum { TRANSCRIBE_FEATURE_PNC = 4, TRANSCRIBE_FEATURE_ITN = 5, TRANSCRIBE_FEATURE_DIARIZATION = 6, + TRANSCRIBE_FEATURE_CONTEXT = 7, } transcribe_feature; TRANSCRIBE_API bool transcribe_model_supports(const struct transcribe_model * model, transcribe_feature feature); diff --git a/src/arch/qwen3_asr/capabilities.cpp b/src/arch/qwen3_asr/capabilities.cpp index 4737e588..18010d49 100644 --- a/src/arch/qwen3_asr/capabilities.cpp +++ b/src/arch/qwen3_asr/capabilities.cpp @@ -17,9 +17,11 @@ void apply_family_invariants(transcribe_model & model) { // lists the BCP-47 codes). Translation is not advertised. caps.supports_translate = false; - // Cancellation is wired at the per-run level. No PNC/ITN toggle; the - // Whisper-specific features do not apply here. + // Cancellation is wired at the per-run level. Free-text recognition + // context is inserted into the chat template's system turn. No PNC/ITN + // toggle; the Whisper-specific features do not apply here. transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CANCELLATION, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, true); } } // namespace transcribe::qwen3_asr diff --git a/src/arch/qwen3_asr/model.cpp b/src/arch/qwen3_asr/model.cpp index e48da627..30bf7b4b 100644 --- a/src/arch/qwen3_asr/model.cpp +++ b/src/arch/qwen3_asr/model.cpp @@ -361,19 +361,31 @@ transcribe_status resolve_chat_tokens(const transcribe::Tokenizer & tok, ChatTok return TRANSCRIBE_OK; } +transcribe_status encode_context(const transcribe::Tokenizer & tok, + const char * context, + std::vector & out_ids) { + out_ids.clear(); + if (context == nullptr || context[0] == '\0') { + return TRANSCRIBE_OK; + } + return tok.encode(context, out_ids); +} + +} // namespace + // Build the prompt token sequence + audio-position list, mirroring the // Qwen3-ASR chat template at the token level: // -// <|im_start|>system\n<|im_end|>\n +// <|im_start|>system\n[context]?<|im_end|>\n // <|im_start|>user\n<|audio_start|><|audio_pad|>*T_enc<|audio_end|><|im_end|>\n // <|im_start|>assistant\n[language {Name}]? // -// System prompt is empty. A non-null `lang_prefix_ids` (resolved via -// encode_language_prefix) is appended after the trailing newline to force an -// output language; kept out of here so this stays a pure token-id assembler. +// A non-null `lang_prefix_ids` (resolved via encode_language_prefix) is +// appended after the trailing newline to force an output language. void build_prompt_tokens(const QwenAsrHParams & hp, const ChatTokens & ct, - int T_enc, + int t_enc, + const std::vector * context_ids, const std::vector * lang_prefix_ids, std::vector & out_ids, std::vector & out_audio_positions) { @@ -383,6 +395,9 @@ void build_prompt_tokens(const QwenAsrHParams & hp, out_ids.push_back(ct.im_start); out_ids.push_back(ct.role_system); out_ids.push_back(ct.newline); + if (context_ids != nullptr && !context_ids->empty()) { + out_ids.insert(out_ids.end(), context_ids->begin(), context_ids->end()); + } out_ids.push_back(ct.im_end); out_ids.push_back(ct.newline); @@ -392,7 +407,7 @@ void build_prompt_tokens(const QwenAsrHParams & hp, out_ids.push_back(hp.audio_start_token_id); const int64_t audio_start_pos = static_cast(out_ids.size()); - for (int i = 0; i < T_enc; ++i) { + for (int i = 0; i < t_enc; ++i) { out_ids.push_back(hp.audio_token_id); out_audio_positions.push_back(audio_start_pos + i); } @@ -410,8 +425,6 @@ void build_prompt_tokens(const QwenAsrHParams & hp, } } -} // namespace - // below; encode_language_prefix matches the qwen3_asr.h declaration.) // BCP-47 → publisher canonical name ("English", "Chinese", ...), which the @@ -564,6 +577,14 @@ transcribe_status run(transcribe_session * session, lang_prefix_ptr = &lang_prefix_ids; } + std::vector context_ids; + if (const transcribe_status st = + encode_context(cm->tok, params != nullptr ? params->context : nullptr, context_ids); + st != TRANSCRIBE_OK) { + return st; + } + const std::vector * context_ptr = context_ids.empty() ? nullptr : &context_ids; + transcribe::debug::init(); // Mel front-end. @@ -725,7 +746,7 @@ transcribe_status run(transcribe_session * session, // Prompt construction. std::vector prompt_ids; std::vector audio_positions; - build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc, lang_prefix_ptr, prompt_ids, audio_positions); + build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc, context_ptr, lang_prefix_ptr, prompt_ids, audio_positions); const int T_prompt = static_cast(prompt_ids.size()); const int prefix_len = audio_positions.empty() ? 0 : static_cast(audio_positions.front()); const int suffix_len = T_prompt - prefix_len - T_enc; @@ -738,8 +759,8 @@ transcribe_status run(transcribe_session * session, transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "qwen3_asr run: input too long — %d audio + %d prompt tokens " "leave no room for output within the %d-token context (need %d). " - "Shorten the audio (see transcribe_capabilities.max_audio_ms) or " - "split it into segments.", + "Shorten the audio or recognition context (see " + "transcribe_capabilities.max_audio_ms), or split the audio into segments.", T_enc, prefix_len + suffix_len, ceiling, T_prompt + k_max_new); return TRANSCRIBE_ERR_INPUT_TOO_LONG; } @@ -1530,6 +1551,16 @@ transcribe_status run_batch(transcribe_session * session, lang_prefix_ptr = &lang_prefix_ids; } + // Shared recognition context (v1: one run_params across the batch), + // encoded once and inserted into every utterance's system turn. + std::vector context_ids; + if (const transcribe_status st = + encode_context(cm->tok, params != nullptr ? params->context : nullptr, context_ids); + st != TRANSCRIBE_OK) { + return st; + } + const std::vector * context_ptr = context_ids.empty() ? nullptr : &context_ids; + // Pass 1: per-utterance encoder + prefill into KV slabs. std::vector> generated(n); std::vector T_prompt(n, 0); @@ -1565,14 +1596,15 @@ transcribe_status run_batch(transcribe_session * session, continue; } std::vector ap; - build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc[b], lang_prefix_ptr, prompt_ids[b], ap); + build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc[b], context_ptr, lang_prefix_ptr, prompt_ids[b], ap); T_prompt[b] = static_cast(prompt_ids[b].size()); prefix_len = ap.empty() ? 0 : static_cast(ap.front()); // Same gate as single-shot run(); the rest of the batch still runs. if (T_prompt[b] + max_new > ceiling) { transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "qwen3_asr run_batch: utterance %d input too long — %d audio + " - "%d prompt tokens exceed the %d-token context. See " + "%d prompt tokens exceed the %d-token context. Shorten the audio " + "or recognition context; see " "transcribe_capabilities.max_audio_ms.", b, T_enc[b], T_prompt[b] - T_enc[b], ceiling); valid[b] = 0; diff --git a/src/arch/qwen3_asr/qwen3_asr.h b/src/arch/qwen3_asr/qwen3_asr.h index 91378235..f6782c63 100644 --- a/src/arch/qwen3_asr/qwen3_asr.h +++ b/src/arch/qwen3_asr/qwen3_asr.h @@ -70,6 +70,17 @@ struct ChatTokens { int32_t role_assistant = -1; }; +// Assemble the token-level Qwen3-ASR chat prompt. Plain-text context and the +// optional language prefix are encoded separately by the caller; this helper +// only places their token ids into the system and assistant turns. +void build_prompt_tokens(const QwenAsrHParams & hp, + const ChatTokens & ct, + int t_enc, + const std::vector * context_ids, + const std::vector * lang_prefix_ids, + std::vector & out_ids, + std::vector & out_audio_positions); + struct QwenAsrModel final : public transcribe_model { Tokenizer tok; QwenAsrHParams hparams; diff --git a/src/transcribe-session.h b/src/transcribe-session.h index c0887204..a4966e8d 100644 --- a/src/transcribe-session.h +++ b/src/transcribe-session.h @@ -256,13 +256,14 @@ struct transcribe_session { // Session-owned copies of the caller's run-params strings, refreshed // on every transcribe_stream_begin. The dispatcher hands the family - // hooks a params view whose language/target_language point HERE, so a + // hooks a params view whose language/target_language/context point HERE, so a // family that captures *run_params holds pointers into library-owned // storage (the public contract lets the caller free its params pointers // the moment begin returns). Stable for the stream's lifetime; only the // next begin mutates them. std::string stream_language_owned; std::string stream_target_language_owned; + std::string stream_context_owned; // UI-facing streaming text state. `full_text` above remains the raw // model hypothesis. `stream_committed_text` is the append-only public diff --git a/src/transcribe.cpp b/src/transcribe.cpp index c370fabc..50c14177 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -745,6 +745,7 @@ constexpr size_t k_min_context_params_size = TRANSCRIBE_FIELD_END(trans // pre-0.2 caller that bypasses the SONAME/abihash checks (dlopen by path, // stale static link, hand-rolled FFI) gets BAD_STRUCT_SIZE instead. constexpr size_t k_min_run_params_size = TRANSCRIBE_FIELD_END(transcribe_run_params, spec_k_drafts); +constexpr size_t k_run_params_context_size = TRANSCRIBE_FIELD_END(transcribe_run_params, context); constexpr size_t k_min_stream_params_size = TRANSCRIBE_FIELD_END(transcribe_stream_params, family); constexpr size_t k_stream_params_commit_policy_size = TRANSCRIBE_FIELD_END(transcribe_stream_params, commit_policy); constexpr size_t k_stream_params_agreement_n_size = @@ -780,6 +781,19 @@ static bool has_field(uint64_t struct_size, size_t field_end) { return struct_size >= static_cast(field_end); } +// Materialize a full library-sized view without reading past an older +// caller's declared prefix. Keep the caller's struct_size in the view so +// downstream has_field checks retain the original ABI truth. Empty context is +// normalized to the documented NULL form; non-empty text is not modified. +void stage_run_params(const transcribe_run_params * caller, transcribe_run_params & out) { + transcribe_run_params_init(&out); + std::memcpy(&out, caller, static_cast(std::min(caller->struct_size, sizeof(out)))); + if (!has_field(caller->struct_size, k_run_params_context_size) || out.context == nullptr || + out.context[0] == '\0') { + out.context = nullptr; + } +} + // Takes the RAW integer, not the enum: a C caller can store any int in the // struct's enum-typed field, and in C++ loading an out-of-range value through // the enum lvalue is UB (UBSan trips on it). Callers memcpy the bytes into an @@ -1364,14 +1378,14 @@ static void publish_observable_delta(transcribe_session * session, } } -// Advisory pnc/itn warning. Emits a WARN when a non-DEFAULT request hits a -// model that does not support runtime control of that axis, then returns so -// the dispatcher proceeds with best-effort semantics. The reserved +// Advisory run-param warning. Emits a WARN when an optional request hits a +// model that does not support the requested behavior, then returns so the +// dispatcher proceeds with best-effort semantics. The reserved // TRANSCRIBE_ERR_UNSUPPORTED_PNC / _ITN codes are NOT returned today // (placeholders for a future opt-in strict mode). The message includes the // arch + variant strings to pinpoint which model dropped the request. -void warn_unsupported_advisory(const struct transcribe_model * model, const struct transcribe_run_params * rp) { - if (model == nullptr || rp == nullptr) { +void apply_unsupported_advisories(const transcribe_model * model, transcribe_run_params & rp) { + if (model == nullptr) { return; } const char * arch_name = (model->arch != nullptr && model->arch->name != nullptr) ? model->arch->name : "(unknown)"; @@ -1383,9 +1397,9 @@ void warn_unsupported_advisory(const struct transcribe_model * model, const stru // Defense in depth: keep raw reads here even though every dispatcher // validates before warning. A future call-site reorder must not turn a // malformed C enum into a typed C++ load. - const int pnc_raw = enum_field_raw(&rp->pnc); - const int itn_raw = enum_field_raw(&rp->itn); - const int diarize_raw = enum_field_raw(&rp->diarize); + const int pnc_raw = enum_field_raw(&rp.pnc); + const int itn_raw = enum_field_raw(&rp.itn); + const int diarize_raw = enum_field_raw(&rp.diarize); char buf[512]; if (pnc_raw != TRANSCRIBE_PNC_MODE_DEFAULT && !transcribe::has_feature(model, TRANSCRIBE_FEATURE_PNC)) { @@ -1422,6 +1436,16 @@ void warn_unsupported_advisory(const struct transcribe_model * model, const stru req, arch_name, variant); transcribe_log_emit_or_stderr(TRANSCRIBE_LOG_LEVEL_WARN, buf); } + if (rp.context != nullptr && !transcribe::has_feature(model, TRANSCRIBE_FEATURE_CONTEXT)) { + std::snprintf(buf, sizeof(buf), + "transcribe_run: caller provided context but model '%s' " + "(variant '%s') does not support recognition context; the context " + "is ignored. Use transcribe_model_supports(model, " + "TRANSCRIBE_FEATURE_CONTEXT) to pre-check.", + arch_name, variant); + transcribe_log_emit_or_stderr(TRANSCRIBE_LOG_LEVEL_WARN, buf); + rp.context = nullptr; + } } } // namespace @@ -1735,6 +1759,9 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session if (const auto st = check_input_struct_size(run_params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + transcribe_run_params run_params_staged; + stage_run_params(run_params, run_params_staged); + run_params = &run_params_staged; if (const auto st = check_input_struct_size(stream_params->struct_size, k_min_stream_params_size); st != TRANSCRIBE_OK) { return st; @@ -1806,7 +1833,7 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session // expose the corresponding runtime toggle. Emitted before // clear_result so the pre-hook "snapshot preserved on rejection" // contract is undisturbed. - warn_unsupported_advisory(session->model, run_params); + apply_unsupported_advisories(session->model, run_params_staged); // Optional family preflight: validates extension field values // (e.g. parakeet's (L, C, R) menu) without mutating state. On @@ -1841,6 +1868,7 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session // wants a run-slot ext at stream begin must plumb it deliberately. session->stream_language_owned = run_params->language != nullptr ? run_params->language : ""; session->stream_target_language_owned = run_params->target_language != nullptr ? run_params->target_language : ""; + session->stream_context_owned = run_params->context != nullptr ? run_params->context : ""; // PREFIX copy, not struct assignment: the size gate above admits any // struct_size >= k_min_run_params_size, so a conforming caller's // allocation may be SHORTER than sizeof (fields past `family`, e.g. @@ -1855,7 +1883,8 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session run_params_owned.language = run_params->language != nullptr ? session->stream_language_owned.c_str() : nullptr; run_params_owned.target_language = run_params->target_language != nullptr ? session->stream_target_language_owned.c_str() : nullptr; - run_params_owned.family = nullptr; + run_params_owned.context = run_params->context != nullptr ? session->stream_context_owned.c_str() : nullptr; + run_params_owned.family = nullptr; const transcribe_status st = session->model->arch->stream_begin(session, &run_params_owned, stream_params); if (st != TRANSCRIBE_OK) { @@ -2102,6 +2131,9 @@ static transcribe_status run_one_inner(struct transcribe_session * sess if (const auto st = check_input_struct_size(params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + transcribe_run_params params_staged; + stage_run_params(params, params_staged); + params = ¶ms_staged; // A run cannot replace an active stream's results — that would // strand the in-flight stream's per-family state. Caller must // finalize or reset first. FINISHED and FAILED both fall through; @@ -2147,7 +2179,7 @@ static transcribe_status run_one_inner(struct transcribe_session * sess if (const transcribe_status st = validate_run_params_common(session, params); st != TRANSCRIBE_OK) { return st; } - warn_unsupported_advisory(session->model, params); + apply_unsupported_advisories(session->model, params_staged); if (params->task == TRANSCRIBE_TASK_TRANSLATE && !session->model->caps.supports_translate) { return TRANSCRIBE_ERR_UNSUPPORTED_TASK; } @@ -2283,6 +2315,9 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * if (const auto st = check_input_struct_size(params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + transcribe_run_params params_staged; + stage_run_params(params, params_staged); + params = ¶ms_staged; if (session->stream_state == TRANSCRIBE_STREAM_ACTIVE) { return TRANSCRIBE_ERR_INVALID_ARG; } @@ -2303,7 +2338,7 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * if (const transcribe_status st = validate_run_params_common(session, params); st != TRANSCRIBE_OK) { return st; } - warn_unsupported_advisory(session->model, params); + apply_unsupported_advisories(session->model, params_staged); if (params->task == TRANSCRIBE_TASK_TRANSLATE && !session->model->caps.supports_translate) { return TRANSCRIBE_ERR_UNSUPPORTED_TASK; } diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index a8dfdb07..e494c99f 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -54,6 +54,14 @@ target_link_libraries(transcribe_api_smoke PRIVATE transcribe) add_test(NAME transcribe_api_smoke COMMAND transcribe_api_smoke) +if(APPLE AND CMAKE_C_COMPILER_ID MATCHES "Clang") + add_test(NAME transcribe_objc_context_header_smoke + COMMAND ${CMAKE_C_COMPILER} + -fsyntax-only + -I${PROJECT_SOURCE_DIR}/include + ${CMAKE_CURRENT_SOURCE_DIR}/objc_context_header_smoke.m) +endif() + add_test( NAME transcribe_extension_umbrella_check COMMAND ${CMAKE_COMMAND} @@ -1306,6 +1314,21 @@ if(TRANSCRIBE_BUILD_EXAMPLES) set_tests_properties(transcribe_cli_smoke PROPERTIES PASS_REGULAR_EXPRESSION "duration:") + add_test(NAME transcribe_cli_context_parse_smoke + COMMAND $ + --context "Glossary: GGUF, ggml" + ${CMAKE_SOURCE_DIR}/samples/jfk.wav) + set_tests_properties(transcribe_cli_context_parse_smoke PROPERTIES + PASS_REGULAR_EXPRESSION "duration:") + + add_test(NAME transcribe_cli_context_model_smoke + COMMAND ${CMAKE_COMMAND} + -DCLI=$ + -DWAV=${CMAKE_SOURCE_DIR}/samples/jfk.wav + -P ${CMAKE_CURRENT_SOURCE_DIR}/cli_context_model_smoke.cmake) + set_tests_properties(transcribe_cli_context_model_smoke PROPERTIES + SKIP_REGULAR_EXPRESSION "SKIP:") + add_test(NAME transcribe_cli_output_smoke COMMAND ${CMAKE_COMMAND} -DCLI=$ diff --git a/tests/api_smoke.c b/tests/api_smoke.c index ce776b72..8f60aa6b 100644 --- a/tests/api_smoke.c +++ b/tests/api_smoke.c @@ -206,6 +206,12 @@ static void test_log_level_values(void) { CHECK(TRANSCRIBE_LOG_LEVEL_CONT == 5); } +static void test_feature_values(void) { + /* Feature values are append-only because model feature bits use them as + * stable bit positions. */ + CHECK(TRANSCRIBE_FEATURE_CONTEXT == 7); +} + static void test_log_set_null(void) { /* Disabling the log sink must not crash. */ transcribe_log_set(NULL, NULL); @@ -239,6 +245,7 @@ static void test_init_macros(void) { CHECK(rp_macro.target_language == NULL); CHECK(rp_macro.keep_special_tags == false); CHECK(rp_macro.family == NULL); + CHECK(rp_macro.context == NULL); struct transcribe_stream_params sp_macro; transcribe_stream_params_init(&sp_macro); @@ -532,6 +539,7 @@ static void test_model_introspection_null(void) { CHECK(transcribe_model_supports(NULL, TRANSCRIBE_FEATURE_PNC) == false); CHECK(transcribe_model_supports(NULL, TRANSCRIBE_FEATURE_ITN) == false); CHECK(transcribe_model_supports(NULL, TRANSCRIBE_FEATURE_DIARIZATION) == false); + CHECK(transcribe_model_supports(NULL, TRANSCRIBE_FEATURE_CONTEXT) == false); CHECK(transcribe_model_supports(NULL, (transcribe_feature) 9999) == false); } @@ -813,6 +821,7 @@ int main(void) { test_abi_metadata(); test_backend_devices(); test_log_level_values(); + test_feature_values(); test_log_set_null(); test_init_macros(); test_log_set_publication(); diff --git a/tests/cli_context_model_smoke.cmake b/tests/cli_context_model_smoke.cmake new file mode 100644 index 00000000..45547111 --- /dev/null +++ b/tests/cli_context_model_smoke.cmake @@ -0,0 +1,24 @@ +foreach(_var CLI WAV) + if(NOT DEFINED ${_var} OR "${${_var}}" STREQUAL "") + message(FATAL_ERROR "${_var} is required") + endif() +endforeach() + +set(_model "$ENV{TRANSCRIBE_QWEN3_ASR_GGUF}") +if(_model STREQUAL "" OR NOT EXISTS "${_model}") + message("SKIP: TRANSCRIBE_QWEN3_ASR_GGUF unset or missing") + return() +endif() + +execute_process( + COMMAND "${CLI}" + -q + -m "${_model}" + --context "Vocabulary: GGUF, ggml, Qwen3-ASR" + "${WAV}" + RESULT_VARIABLE _result + OUTPUT_VARIABLE _stdout + ERROR_VARIABLE _stderr) +if(NOT _result EQUAL 0) + message(FATAL_ERROR "context command failed (${_result}):\n${_stdout}\n${_stderr}") +endif() diff --git a/tests/objc_context_header_smoke.m b/tests/objc_context_header_smoke.m new file mode 100644 index 00000000..219f3409 --- /dev/null +++ b/tests/objc_context_header_smoke.m @@ -0,0 +1,14 @@ +#import +#import "transcribe.h" + +static void configure_context(struct transcribe_run_params * params) { + NSString * context = @"Vocabulary: GGUF, ggml, Qwen3-ASR"; + transcribe_run_params_init(params); + params->context = context.UTF8String; +} + +int main(void) { + struct transcribe_run_params params; + configure_context(¶ms); + return params.context == NULL || TRANSCRIBE_FEATURE_CONTEXT != 7; +} diff --git a/tests/qwen3_asr_bpe_parity.cpp b/tests/qwen3_asr_bpe_parity.cpp index 8d6f09f1..93123e70 100644 --- a/tests/qwen3_asr_bpe_parity.cpp +++ b/tests/qwen3_asr_bpe_parity.cpp @@ -151,6 +151,34 @@ int main() { check_ids_equal(label, f.pub_name, f.ids, f.n_ids, got); } + // Section 3: full chat-prompt parity. These literals come from the + // publisher tokenizer/template: context is plain text in the system turn, + // audio pads occupy the user turn, and language remains an independent + // assistant prefix. + { + std::vector context_ids; + std::vector language_ids; + if (tok.encode("line1\nline2", context_ids) != TRANSCRIBE_OK || + transcribe::qwen3_asr::encode_language_prefix(tok, "en", language_ids) != TRANSCRIBE_OK) { + std::fprintf(stderr, "FAIL[prompt] could not encode context or language prefix\n"); + ++g_failures; + } + + std::vector prompt_ids; + std::vector audio_positions; + transcribe::qwen3_asr::build_prompt_tokens(qm->hparams, qm->chat_tokens, 2, &context_ids, &language_ids, + prompt_ids, audio_positions); + const int32_t expected[] = { + 151644, 8948, 198, 1056, 16, 198, 1056, 17, 151645, 198, 151644, 872, 198, + 151669, 151676, 151676, 151670, 151645, 198, 151644, 77091, 198, 11528, 6364, 151704, + }; + check_ids_equal("prompt", "line1\\nline2 + en", expected, sizeof(expected) / sizeof(expected[0]), prompt_ids); + if (audio_positions != std::vector({ 14, 15 })) { + std::fprintf(stderr, "FAIL[prompt] audio positions differ\n"); + ++g_failures; + } + } + transcribe_model_free(model); if (g_failures > 0) { diff --git a/tests/qwen3_asr_e2e_smoke.cpp b/tests/qwen3_asr_e2e_smoke.cpp index f88f5fe0..9efca43e 100644 --- a/tests/qwen3_asr_e2e_smoke.cpp +++ b/tests/qwen3_asr_e2e_smoke.cpp @@ -241,6 +241,8 @@ int main() { } CHECK_STR_EQ(transcribe_model_arch_string(model), "qwen3_asr"); + CHECK(transcribe_model_supports(model, TRANSCRIBE_FEATURE_CONTEXT)); + CHECK(!transcribe_model_supports(model, TRANSCRIBE_FEATURE_INITIAL_PROMPT)); // Language hinting contract: // - A BCP-47 code advertised in caps.languages must be accepted @@ -315,9 +317,54 @@ int main() { ++g_failures; } + // Context and language occupy independent chat-template fields and + // may be used together through the public ABI. + rp_lang.language = "en"; + rp_lang.context = "Vocabulary: fellow Americans, country."; + const transcribe_status st_context = transcribe_run(ctx, pcm.data(), static_cast(pcm.size()), &rp_lang); + if (st_context != TRANSCRIBE_OK || std::strlen(transcribe_full_text(ctx)) == 0) { + std::fprintf(stderr, "FAIL: context + language run returned %s\n", transcribe_status_string(st_context)); + ++g_failures; + } + transcribe_session_free(ctx); } + // Context tokens share the decoder window with audio and generation. A + // request that cannot fit is rejected rather than silently truncating the + // caller's background text. The same run-level context is broadcast to all + // batch rows. + { + transcribe_session_params cp; + transcribe_session_params_init(&cp); + cp.n_ctx = 512; + transcribe_session * ctx = nullptr; + if (transcribe_session_init(model, &cp, &ctx) != TRANSCRIBE_OK || ctx == nullptr) { + std::fprintf(stderr, "FAIL: context-limit session init failed\n"); + ++g_failures; + } else { + std::string long_context; + for (int i = 0; i < 700; ++i) { + long_context += "context "; + } + transcribe_run_params rp; + transcribe_run_params_init(&rp); + rp.context = long_context.c_str(); + + const std::vector silence(1600, 0.0f); + CHECK(transcribe_run(ctx, silence.data(), static_cast(silence.size()), &rp) == + TRANSCRIBE_ERR_INPUT_TOO_LONG); + + const float * batch_pcm[] = { silence.data(), silence.data() }; + const int batch_n[] = { static_cast(silence.size()), static_cast(silence.size()) }; + CHECK(transcribe_run_batch(ctx, batch_pcm, batch_n, 2, &rp) == TRANSCRIBE_OK); + CHECK(transcribe_batch_n_results(ctx) == 2); + CHECK(transcribe_batch_status(ctx, 0) == TRANSCRIBE_ERR_INPUT_TOO_LONG); + CHECK(transcribe_batch_status(ctx, 1) == TRANSCRIBE_ERR_INPUT_TOO_LONG); + transcribe_session_free(ctx); + } + } + // Case 1: jfk.wav (single-chunk). run_case(model, "jfk", jfk_wav, k_jfk_reference_text, k_jfk_max_edit_distance); diff --git a/tests/qwen3_asr_smoke.cpp b/tests/qwen3_asr_smoke.cpp index 98ac5ba2..b72993d7 100644 --- a/tests/qwen3_asr_smoke.cpp +++ b/tests/qwen3_asr_smoke.cpp @@ -169,6 +169,8 @@ int main() { CHECK_EQ_INT(caps->n_languages, 4); CHECK(caps->languages != nullptr); CHECK_EQ_INT(caps->max_timestamp_kind, TRANSCRIBE_TIMESTAMPS_NONE); + CHECK(transcribe_model_supports(model, TRANSCRIBE_FEATURE_CONTEXT)); + CHECK(!transcribe_model_supports(model, TRANSCRIBE_FEATURE_INITIAL_PROMPT)); } // Internal-view assertions. @@ -237,6 +239,24 @@ int main() { CHECK_EQ_INT(ct.role_system, 5); CHECK_EQ_INT(ct.role_user, 6); CHECK_EQ_INT(ct.role_assistant, 7); + + const std::vector context_ids = { 10, 11 }; + const std::vector language_ids = { 12 }; + std::vector prompt_ids; + std::vector audio_positions; + transcribe::qwen3_asr::build_prompt_tokens(hp, ct, 2, &context_ids, &language_ids, prompt_ids, audio_positions); + const std::vector expected = { + 2, 5, 4, 10, 11, 3, 4, 2, 6, 4, 16, 18, 18, 17, 3, 4, 2, 7, 4, 12, + }; + CHECK(prompt_ids == expected); + CHECK(audio_positions == std::vector({ 11, 12 })); + + transcribe::qwen3_asr::build_prompt_tokens(hp, ct, 1, nullptr, nullptr, prompt_ids, audio_positions); + const std::vector expected_empty = { + 2, 5, 4, 3, 4, 2, 6, 4, 16, 18, 17, 3, 4, 2, 7, 4, + }; + CHECK(prompt_ids == expected_empty); + CHECK(audio_positions == std::vector({ 9 })); } // ----- Encoder subsample slots populated ----- diff --git a/tests/run_dispatch_unit.cpp b/tests/run_dispatch_unit.cpp index e2ae9487..a07fd3af 100644 --- a/tests/run_dispatch_unit.cpp +++ b/tests/run_dispatch_unit.cpp @@ -60,6 +60,10 @@ constexpr uint32_t kFakeRunKind = 0xF00D; bool g_run_called = false; transcribe_status g_run_validate_status = TRANSCRIBE_OK; +std::string g_context_seen; +int g_context_run_calls = 0; + +const transcribe::Arch & run_validate_arch(); transcribe_status fake_run(transcribe_session * session, const float * pcm, @@ -67,7 +71,8 @@ transcribe_status fake_run(transcribe_session * session, const transcribe_run_params * params) { (void) pcm; (void) n_samples; - (void) params; + g_context_seen = params != nullptr && params->context != nullptr ? params->context : ""; + ++g_context_run_calls; g_run_called = true; // A successful run installs a fresh result. session->full_text = "fresh result"; @@ -75,6 +80,96 @@ transcribe_status fake_run(transcribe_session * session, return TRANSCRIBE_OK; } +int g_context_warns = 0; + +void context_log_cb(transcribe_log_level level, const char * msg, void * userdata) { + (void) userdata; + if (level == TRANSCRIBE_LOG_LEVEL_WARN && msg != nullptr && std::strstr(msg, "context") != nullptr && + std::strstr(msg, "ignored") != nullptr) { + ++g_context_warns; + } +} + +void test_context_preserved_or_ignored_once() { + transcribe_model model; + model.arch = &run_validate_arch(); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, true); + + transcribe_session session; + session.model = &model; + + transcribe_run_params params; + transcribe_run_params_init(¶ms); + params.context = + " Glossary: GGUF\n<|im_start|> \xe6" + "\x97" + "\xa5" + "\xe6" + "\x9c" + "\xac" + "\xe8" + "\xaa" + "\x9e" + " "; + + float sample = 0.0f; + g_context_seen.clear(); + g_context_run_calls = 0; + CHECK(transcribe_run(&session, &sample, 1, ¶ms) == TRANSCRIBE_OK); + CHECK(g_context_seen == params.context); + CHECK(g_context_run_calls == 1); + + const float * batch_pcm[] = { &sample, &sample }; + const int batch_n[] = { 1, 1 }; + g_context_run_calls = 0; + CHECK(transcribe_run_batch(&session, batch_pcm, batch_n, 2, ¶ms) == TRANSCRIBE_OK); + CHECK(g_context_seen == params.context); + CHECK(g_context_run_calls == 2); + + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, false); + transcribe_log_set(context_log_cb, nullptr); + g_context_warns = 0; + g_context_run_calls = 0; + CHECK(transcribe_run_batch(&session, batch_pcm, batch_n, 2, ¶ms) == TRANSCRIBE_OK); + transcribe_log_set(nullptr, nullptr); + CHECK(g_context_warns == 1); + CHECK(g_context_run_calls == 2); + CHECK(g_context_seen == ""); + + params.context = ""; + g_context_warns = 0; + g_context_run_calls = 0; + transcribe_log_set(context_log_cb, nullptr); + CHECK(transcribe_run(&session, &sample, 1, ¶ms) == TRANSCRIBE_OK); + transcribe_log_set(nullptr, nullptr); + CHECK(g_context_warns == 0); + CHECK(g_context_run_calls == 1); + CHECK(g_context_seen == ""); +} + +void test_context_absent_from_old_run_params_prefix() { + transcribe_model model; + model.arch = &run_validate_arch(); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, true); + + transcribe_session session; + session.model = &model; + + const size_t min_size = + offsetof(transcribe_run_params, spec_k_drafts) + sizeof(((transcribe_run_params *) nullptr)->spec_k_drafts); + transcribe_run_params staged; + transcribe_run_params_init(&staged); + staged.struct_size = min_size; + auto * params = static_cast(std::malloc(min_size)); + std::memcpy(params, &staged, min_size); + + float sample = 0.0f; + g_context_seen.clear(); + CHECK(transcribe_run(&session, &sample, 1, params) == TRANSCRIBE_OK); + std::free(params); + CHECK(g_context_seen == ""); +} + bool fake_accepts_run_kind(const transcribe_model * model, transcribe_ext_slot slot, uint32_t kind) { (void) model; return slot == TRANSCRIBE_EXT_SLOT_RUN && kind == kFakeRunKind; @@ -445,6 +540,8 @@ int main() { test_no_run_hook_clears_and_not_implemented(); test_run_validate_failure_preserves_snapshot(); test_run_validate_success_clears_and_runs(); + test_context_preserved_or_ignored_once(); + test_context_absent_from_old_run_params_prefix(); test_advisory_enum_validation(); test_batch_abort_pads_missing_to_n(); test_batch_fastpath_abort_pads_missing_to_n(); diff --git a/tests/stream_dispatch_unit.cpp b/tests/stream_dispatch_unit.cpp index fa0df5e8..032e325b 100644 --- a/tests/stream_dispatch_unit.cpp +++ b/tests/stream_dispatch_unit.cpp @@ -1395,6 +1395,7 @@ void test_two_sessions_independent_streams() { transcribe_run_params g_retained_rp; // the family-side shallow copy std::string g_language_seen_at_feed; +std::string g_context_seen_at_feed; transcribe_status retain_stream_begin(transcribe_session * session, const transcribe_run_params * run_params, @@ -1414,6 +1415,7 @@ transcribe_status retain_stream_feed(transcribe_session * session, (void) n_samples; (void) update; g_language_seen_at_feed = g_retained_rp.language != nullptr ? g_retained_rp.language : ""; + g_context_seen_at_feed = g_retained_rp.context != nullptr ? g_retained_rp.context : ""; return TRANSCRIBE_OK; } @@ -1436,6 +1438,7 @@ void test_begin_copies_param_strings_out() { model.arch = &arch; model.caps.supports_streaming = true; model.caps.max_timestamp_kind = TRANSCRIBE_TIMESTAMPS_NONE; + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, true); transcribe_session session; session.model = &model; @@ -1447,18 +1450,34 @@ void test_begin_copies_param_strings_out() { transcribe_run_params_init(rp); char * lang = static_cast(std::malloc(16)); std::snprintf(lang, 16, "en-US"); - rp->language = lang; + rp->language = lang; + char * context = static_cast(std::malloc(32)); + std::snprintf(context, 32, + " context \xe6" + "\x97" + "\xa5" + "\xe6" + "\x9c" + "\xac" + "\xe8" + "\xaa" + "\x9e" + " "); + rp->context = context; transcribe_stream_params sp; transcribe_stream_params_init(&sp); g_language_seen_at_feed.clear(); + g_context_seen_at_feed.clear(); CHECK(transcribe_stream_begin(&session, rp, &sp) == TRANSCRIBE_OK); // Caller destroys its storage: scribble (still owned, so well-defined), // then free. std::memset(lang, 'X', 15); + std::memset(context, 'Y', 31); std::memset(rp, 0x5A, sizeof(*rp)); std::free(lang); + std::free(context); std::free(rp); transcribe_stream_update up; @@ -1466,6 +1485,17 @@ void test_begin_copies_param_strings_out() { const float pcm[160] = {}; CHECK(transcribe_stream_feed(&session, pcm, 160, &up) == TRANSCRIBE_OK); CHECK(g_language_seen_at_feed == "en-US"); + CHECK(g_context_seen_at_feed == + " context \xe6" + "\x97" + "\xa5" + "\xe6" + "\x9c" + "\xac" + "\xe8" + "\xaa" + "\x9e" + " "); // The run-slot family ext pointer must not survive into the retained // copy either: it is consumed during begin per its copy-out contract, // and a stale pointer would dangle just like the strings. @@ -1504,6 +1534,7 @@ void test_begin_accepts_min_prefix_run_params() { model.arch = &arch; model.caps.supports_streaming = true; model.caps.max_timestamp_kind = TRANSCRIBE_TIMESTAMPS_NONE; + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT, true); transcribe_session session; session.model = &model; @@ -1525,6 +1556,7 @@ void test_begin_accepts_min_prefix_run_params() { transcribe_stream_params_init(&sp); g_retained_rp = transcribe_run_params{}; g_language_seen_at_feed.clear(); + g_context_seen_at_feed.clear(); CHECK(transcribe_stream_begin(&session, rp, &sp) == TRANSCRIBE_OK); std::free(rp); @@ -1540,6 +1572,7 @@ void test_begin_accepts_min_prefix_run_params() { const float pcm[160] = {}; CHECK(transcribe_stream_feed(&session, pcm, 160, &up) == TRANSCRIBE_OK); CHECK(g_language_seen_at_feed == "en"); + CHECK(g_context_seen_at_feed == ""); transcribe_stream_reset(&session); }