From a2b3e9ed68c1dc9ec08475d41e33db34442f7c5f Mon Sep 17 00:00:00 2001 From: Paul Scheduikat Date: Wed, 23 Sep 2026 21:51:29 +0200 Subject: [PATCH] feat: constrain Nemotron automatic language detection --- .../python/src/transcribe_cpp/_generated.py | 6 +- bindings/rust/sys/src/transcribe_sys.rs | 12 ++- bindings/rust/transcribe-cpp/src/session.rs | 49 +++++++++++- bindings/typescript/src/_generated.ts | 6 +- include/transcribe.abihash | 2 +- include/transcribe.h | 8 ++ src/arch/parakeet/decoder.cpp | 62 +++++++++++---- src/arch/parakeet/decoder.h | 35 +++++---- src/arch/parakeet/model.cpp | 73 +++++++++++++++-- src/arch/parakeet/parakeet.h | 5 ++ src/transcribe-session.h | 6 +- src/transcribe.cpp | 49 +++++++++--- tests/CMakeLists.txt | 18 +++++ tests/api_smoke.c | 2 + tests/nemotron_allowlist_real_stream.cpp | 78 +++++++++++++++++++ tests/parakeet_language_mask_unit.cpp | 60 ++++++++++++++ tests/parakeet_stream_ext_reject_unit.cpp | 20 +++++ 17 files changed, 429 insertions(+), 62 deletions(-) create mode 100644 tests/nemotron_allowlist_real_stream.cpp create mode 100644 tests/parakeet_language_mask_unit.cpp diff --git a/bindings/python/src/transcribe_cpp/_generated.py b/bindings/python/src/transcribe_cpp/_generated.py index ac9763c2..e14ea9c8 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 = "8622da005d6a6fcc" # === enum constants === TRANSCRIBE_OK = 0 @@ -167,7 +167,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), ("n_allowed_languages", _c.c_int32), ("allowed_languages", _c.POINTER(_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 +212,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, 'n_allowed_languages': 68, 'allowed_languages': 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/rust/sys/src/transcribe_sys.rs b/bindings/rust/sys/src/transcribe_sys.rs index 363cd6e3..de537f77 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 = 8622da005d6a6fcc /// 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 = "8622da005d6a6fcc"; /* automatically generated by rust-bindgen 0.72.1 */ @@ -346,10 +346,12 @@ pub struct transcribe_run_params { pub keep_special_tags: bool, pub family: *const transcribe_ext, pub spec_k_drafts: i32, + pub n_allowed_languages: i32, + pub allowed_languages: *const *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 +376,10 @@ 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::n_allowed_languages"] + [::std::mem::offset_of!(transcribe_run_params, n_allowed_languages) - 68usize]; + ["Offset of field: transcribe_run_params::allowed_languages"] + [::std::mem::offset_of!(transcribe_run_params, allowed_languages) - 72usize]; }; unsafe extern "C" { pub fn transcribe_run_params_init(params: *mut transcribe_run_params); diff --git a/bindings/rust/transcribe-cpp/src/session.rs b/bindings/rust/transcribe-cpp/src/session.rs index 220f2b63..4c81b7c1 100644 --- a/bindings/rust/transcribe-cpp/src/session.rs +++ b/bindings/rust/transcribe-cpp/src/session.rs @@ -35,6 +35,8 @@ pub struct RunOptions { pub diarize: Diarize, /// Source language hint (ISO code), or `None` to autodetect. pub language: Option, + /// Limit decoded language tags in automatic mode. Empty is unrestricted. + pub allowed_languages: Vec, /// Target language for translation, or `None`. pub target_language: Option, /// Keep special vocab tags (e.g. `<|...|>`) in the returned text. @@ -54,6 +56,7 @@ impl Default for RunOptions { itn: Itn::Default, diarize: Diarize::Default, language: None, + allowed_languages: Vec::new(), target_language: None, keep_special_tags: false, spec_k_drafts: -1, @@ -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, _family, _allowed, _allowed_ptrs) = 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, _family, _allowed, _allowed_ptrs) = 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, _family, _allowed, _allowed_ptrs) = 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 @@ -423,6 +426,8 @@ type RunParamsBundle = ( Option, Option, Option, + Vec, + Vec<*const std::os::raw::c_char>, ); /// Build `transcribe_run_params` from options. The returned keepalives own the @@ -444,6 +449,20 @@ fn build_run_params(o: &RunOptions) -> Result { let target = o.target_language.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()); + let allowed: Vec = o + .allowed_languages + .iter() + .map(|code| CString::new(code.as_str())) + .collect::>()?; + let allowed_ptrs: Vec<*const std::os::raw::c_char> = + allowed.iter().map(|code| code.as_ptr()).collect(); + params.n_allowed_languages = i32::try_from(allowed_ptrs.len()) + .map_err(|_| Error::InvalidArgument("too many allowed languages".into()))?; + params.allowed_languages = if allowed_ptrs.is_empty() { + std::ptr::null() + } else { + allowed_ptrs.as_ptr() + }; let family = o .family @@ -452,7 +471,29 @@ 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, family, allowed, allowed_ptrs)) +} + +#[cfg(test)] +mod allowlist_tests { + use super::*; + + #[test] + fn run_options_keep_allowed_language_strings_alive() { + let options = RunOptions { + allowed_languages: vec!["en-US".into(), "de-DE".into()], + ..Default::default() + }; + let (params, _, _, _, codes, pointers) = build_run_params(&options).unwrap(); + assert_eq!(params.n_allowed_languages, 2); + assert_eq!(params.allowed_languages, pointers.as_ptr()); + assert_eq!(codes[0].to_str().unwrap(), "en-US"); + assert_eq!(codes[1].to_str().unwrap(), "de-DE"); + + let (empty, _, _, _, _, _) = build_run_params(&RunOptions::default()).unwrap(); + assert_eq!(empty.n_allowed_languages, 0); + assert!(empty.allowed_languages.is_null()); + } } /// PCM/utterance lengths cross the ABI as `int`; reject anything that overflows. diff --git a/bindings/typescript/src/_generated.ts b/bindings/typescript/src/_generated.ts index fff2ca59..0d27e077 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 = "8622da005d6a6fcc"; // === enum constants === export const TRANSCRIBE_OK = 0; @@ -121,7 +121,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, 'n_allowed_languages': 68, 'allowed_languages': 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 +166,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', n_allowed_languages: 'int32_t', allowed_languages: 'void *' }); 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/include/transcribe.abihash b/include/transcribe.abihash index b0e23c5d..8e4a1691 100644 --- a/include/transcribe.abihash +++ b/include/transcribe.abihash @@ -1 +1 @@ -7df72bf9e667b8c2 +8622da005d6a6fcc diff --git a/include/transcribe.h b/include/transcribe.h index 6d74a75f..2e0700a0 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -1118,6 +1118,14 @@ struct transcribe_run_params { * to know whether the field will take effect. */ int32_t spec_k_drafts; + + /* Optional language allowlist for automatic detection. Codes must match + * model-advertised languages; unknown codes return UNSUPPORTED_LANGUAGE. + * An empty list leaves detection unrestricted. Currently used by + * Nemotron streaming. Ignored when language is explicitly set. The + * library copies strings for the duration of a streaming run. */ + int32_t n_allowed_languages; + const char * const * allowed_languages; }; TRANSCRIBE_API void transcribe_run_params_init(struct transcribe_run_params * params); diff --git a/src/arch/parakeet/decoder.cpp b/src/arch/parakeet/decoder.cpp index 27a83c2f..10fbe380 100644 --- a/src/arch/parakeet/decoder.cpp +++ b/src/arch/parakeet/decoder.cpp @@ -956,6 +956,26 @@ int argmax_range(const float * data, int n) { return best_i; } +} // namespace + +int argmax_language_masked(const float * data, int n, const std::vector & blocked_language_tokens) { + if (blocked_language_tokens.empty()) { + return argmax_range(data, n); + } + int best = -1; + for (int i = 0; i < n; ++i) { + if (i < static_cast(blocked_language_tokens.size()) && blocked_language_tokens[i]) { + continue; + } + if (best < 0 || data[i] > data[best]) { + best = i; + } + } + return best; +} + +namespace { + // Entropy-based confidence over a token-logit slice. Mirrors the // reference ParakeetTDT.decode_greedy path: // @@ -1224,12 +1244,13 @@ transcribe_status decode_tdt_greedy(const HostDecoderWeights & w, // duration). Matches the reference dump points (`dec.embed.0`, // `dec.lstm.{layer}.{h,c}.0`, `dec.joint.0`) emitted on iter 1. -transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, - const float * enc_out, - int T_enc, - int d_enc, - int n_threads, - std::vector & out_tokens) { +transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, + const float * enc_out, + int T_enc, + int d_enc, + int n_threads, + const std::vector & blocked_language_tokens, + std::vector & out_tokens) { if (enc_out == nullptr || T_enc <= 0 || d_enc <= 0) { return TRANSCRIBE_ERR_INVALID_ARG; } @@ -1312,7 +1333,10 @@ transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, // RNNT joint output is just `n_token_cls` floats (no duration extras). const float * token_logits = logits.data(); - const int pred_token = argmax_range(token_logits, n_token_cls); + const int pred_token = argmax_language_masked(token_logits, n_token_cls, blocked_language_tokens); + if (pred_token < 0) { + return TRANSCRIBE_ERR_BACKEND; + } if (iter == 1 && transcribe::debug::enabled()) { const long long s_h = H; @@ -1381,15 +1405,16 @@ transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, // LSTM state (state_io) and previous token (last_token_io). The // chunk's encoder frames are decoded in stream-wide coordinates // (step_at_emit = frame_offset + local_step). No timing log. -transcribe_status decode_rnnt_greedy_streaming(const HostDecoderWeights & w, - const float * enc_out, - int T_enc_new, - int d_enc, - LstmState & state_io, - int & last_token_io, - int frame_offset, - int n_threads, - std::vector & out_tokens) { +transcribe_status decode_rnnt_greedy_streaming(const HostDecoderWeights & w, + const float * enc_out, + int T_enc_new, + int d_enc, + LstmState & state_io, + int & last_token_io, + int frame_offset, + int n_threads, + const std::vector & blocked_language_tokens, + std::vector & out_tokens) { if (enc_out == nullptr || T_enc_new <= 0 || d_enc <= 0) { return TRANSCRIBE_ERR_INVALID_ARG; } @@ -1485,7 +1510,10 @@ transcribe_status decode_rnnt_greedy_streaming(const HostDecoderWeights & w, joint_step(w.joint, jg, enc_proj, decoder_out, logits); const float * token_logits = logits.data(); - const int pred_token = argmax_range(token_logits, n_token_cls); + const int pred_token = argmax_language_masked(token_logits, n_token_cls, blocked_language_tokens); + if (pred_token < 0) { + return TRANSCRIBE_ERR_BACKEND; + } const bool is_blank = (pred_token == blank_id); if (is_blank) { diff --git a/src/arch/parakeet/decoder.h b/src/arch/parakeet/decoder.h index e4de9dfa..623aa12a 100644 --- a/src/arch/parakeet/decoder.h +++ b/src/arch/parakeet/decoder.h @@ -228,12 +228,16 @@ transcribe_status decode_tdt_greedy(const HostDecoderWeights & w, // rule is "blank → advance one frame, non-blank → emit + stay, capped by // tdt_max_symbols". Per-emit duration_frames is fixed at 1. Same I/O // contract and TdtToken result type as decode_tdt_greedy. -transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, - const float * enc_out, - int T_enc, - int d_enc, - int n_threads, - std::vector & out_tokens); +transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, + const float * enc_out, + int T_enc, + int d_enc, + int n_threads, + const std::vector & blocked_language_tokens, + std::vector & out_tokens); + +// The mask marks only disallowed locale-tag IDs; every other score is unchanged. +int argmax_language_masked(const float * data, int n, const std::vector & blocked_language_tokens); // Streaming variant of RNN-T greedy decode. Consumes T_enc_new encoder // frames (the chunk just produced) and APPENDS emitted tokens to @@ -242,15 +246,16 @@ transcribe_status decode_rnnt_greedy(const HostDecoderWeights & w, // index of this chunk's first frame (so step_at_emit lands in // stream-wide coordinates). state_io must have been reset to a fresh // start-of-sequence state at stream_begin (last_token_io = -1). -transcribe_status decode_rnnt_greedy_streaming(const HostDecoderWeights & w, - const float * enc_out, - int T_enc_new, - int d_enc, - LstmState & state_io, - int & last_token_io, - int frame_offset, - int n_threads, - std::vector & out_tokens); +transcribe_status decode_rnnt_greedy_streaming(const HostDecoderWeights & w, + const float * enc_out, + int T_enc_new, + int d_enc, + LstmState & state_io, + int & last_token_io, + int frame_offset, + int n_threads, + const std::vector & blocked_language_tokens, + std::vector & out_tokens); // Run CTC greedy decode end-to-end. Per-frame: logits = W @ enc[t] + b, // argmax; collapse rule "drop adjacent duplicates, then drop blanks" diff --git a/src/arch/parakeet/model.cpp b/src/arch/parakeet/model.cpp index bc426aa7..b3723d08 100644 --- a/src/arch/parakeet/model.cpp +++ b/src/arch/parakeet/model.cpp @@ -678,6 +678,54 @@ static bool is_lang_tag_piece(const std::string & p) { return i == end; // interior consumed exactly up to '>' } +} // namespace + +transcribe_status resolve_language_block_mask(const ParakeetModel * pm, + const transcribe_run_params * params, + std::vector & mask) { + mask.clear(); + if (params == nullptr || params->language != nullptr || + params->struct_size < offsetof(transcribe_run_params, allowed_languages) + sizeof(params->allowed_languages) || + params->n_allowed_languages == 0) { + return TRANSCRIBE_OK; + } + if (params->n_allowed_languages < 0 || params->allowed_languages == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (!pm->hparams.has_prompt || pm->host_decoder.head_kind != HostHeadKind::RNNT) { + return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; + } + std::vector allowed_ids; + for (int i = 0; i < params->n_allowed_languages; ++i) { + const char * code = params->allowed_languages[i]; + if (code == nullptr || *code == '\0') { + return TRANSCRIBE_ERR_INVALID_ARG; + } + bool supported = false; + for (int j = 0; pm->caps.languages != nullptr && j < pm->caps.n_languages; ++j) { + if (pm->caps.languages[j] != nullptr && std::strcmp(code, pm->caps.languages[j]) == 0) { + supported = true; + break; + } + } + const int id = pm->tok.find("<" + std::string(code) + ">"); + if (!supported || id < 0 || id >= pm->host_decoder.n_vocab || !is_lang_tag_piece(pm->tok.token(id))) { + return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; + } + allowed_ids.push_back(id); + } + mask.resize(static_cast(pm->tok.n_tokens()), 0); + for (int id = 0; id < pm->tok.n_tokens(); ++id) { + if (is_lang_tag_piece(pm->tok.token(id)) && + std::find(allowed_ids.begin(), allowed_ids.end(), id) == allowed_ids.end()) { + mask[static_cast(id)] = 1; + } + } + return TRANSCRIBE_OK; +} + +namespace { + // Drop this piece from the public result when keep_special_tags is off: // stripped if CONTROL-typed or matching the locale-tag pattern // (transitional fallback). Shared by the offline and streaming builders. @@ -730,6 +778,11 @@ static transcribe_status decode_and_populate(ParakeetSession * pc, int d_enc, int utt_index, const char * enc_dump_name_override = nullptr) { + std::vector language_block_mask; + if (const transcribe_status st = resolve_language_block_mask(pm, params, language_block_mask); + st != TRANSCRIBE_OK) { + return st; + } // Default dump name is "dec.enc_out"; the prompt-conditioned path // overrides to "dec.enc_out_prompted" so the comparator sees the // post-prompt tensor under its expected filename. @@ -758,7 +811,8 @@ static transcribe_status decode_and_populate(ParakeetSession * pc, st = decode_tdt_greedy(pm->host_decoder, enc, T_enc, d_enc, pc->n_threads, pc->raw_tokens); break; case HostHeadKind::RNNT: - st = decode_rnnt_greedy(pm->host_decoder, enc, T_enc, d_enc, pc->n_threads, pc->raw_tokens); + st = decode_rnnt_greedy(pm->host_decoder, enc, T_enc, d_enc, pc->n_threads, language_block_mask, + pc->raw_tokens); break; case HostHeadKind::CTC: st = decode_ctc_greedy(pm->host_decoder, enc, T_enc, d_enc, pc->n_threads, pc->raw_tokens); @@ -2217,7 +2271,7 @@ transcribe_status emit_streaming_chunk(ParakeetSession * pc, if (const transcribe_status st = decode_rnnt_greedy_streaming( pm->host_decoder, pc->enc_host.data(), T_q_new, d_enc, pc->stream_dec_state.lstm_state, pc->stream_dec_state.prev_token_id, static_cast(pc->stream_dec_state.frame_offset), pc->n_threads, - pc->raw_tokens); + pc->stream_language_block_mask, pc->raw_tokens); st != TRANSCRIBE_OK) { return st; } @@ -2610,7 +2664,7 @@ transcribe_status emit_buffered_chunk(ParakeetSession * pc, if (const transcribe_status st = decode_rnnt_greedy_streaming( pm->host_decoder, enc_chunk, T_to_decode, d_enc, pc->stream_dec_state.lstm_state, pc->stream_dec_state.prev_token_id, static_cast(pc->stream_dec_state.frame_offset), pc->n_threads, - pc->raw_tokens); + pc->stream_language_block_mask, pc->raw_tokens); st != TRANSCRIBE_OK) { return st; } @@ -2784,14 +2838,19 @@ transcribe_status resolve_cache_aware_stream_geom(const ParakeetModel * // Pre-flight: validate caller extension fields without mutating state. // Called by the dispatcher before clear_result, so a rejection leaves the // previous snapshot intact. stream_begin re-runs the same resolvers. -transcribe_status stream_validate(const transcribe_session * session, - const transcribe_run_params * /*run_params*/, +transcribe_status stream_validate(const transcribe_session * session, + const transcribe_run_params * run_params, const transcribe_stream_params * stream_params) { const auto * pc = static_cast(session); const auto * pm = static_cast(pc->model); if (pm == nullptr || pm->plan.scheduler_list.empty()) { return TRANSCRIBE_ERR_INVALID_ARG; } + std::vector language_block_mask; + if (const transcribe_status st = resolve_language_block_mask(pm, run_params, language_block_mask); + st != TRANSCRIBE_OK) { + return st; + } const bool is_chunked_limited = (pm->hparams.enc_att_context_style == ParakeetHParams::AttContextStyle::ChunkedLimited); @@ -2818,6 +2877,10 @@ transcribe_status stream_begin(transcribe_session * session, if (pm == nullptr || pm->plan.scheduler_list.empty()) { return TRANSCRIBE_ERR_INVALID_ARG; } + if (const transcribe_status st = resolve_language_block_mask(pm, run_params, pc->stream_language_block_mask); + st != TRANSCRIBE_OK) { + return st; + } // Streaming may run without going through run(), so init the dumper here. transcribe::debug::init(); diff --git a/src/arch/parakeet/parakeet.h b/src/arch/parakeet/parakeet.h index ca8479b7..8e29be48 100644 --- a/src/arch/parakeet/parakeet.h +++ b/src/arch/parakeet/parakeet.h @@ -135,6 +135,10 @@ struct ParakeetModel final : public transcribe_model { const transcribe::Tokenizer * tokenizer() const override { return &tok; } }; +transcribe_status resolve_language_block_mask(const ParakeetModel * pm, + const transcribe_run_params * params, + std::vector & mask); + // Per-context streaming encoder state, allocated lazily on the first // stream_begin for ChunkedLimited variants. Mirrors NeMo's // get_initial_cache_state tuple. Layout cheat sheet: @@ -264,6 +268,7 @@ struct ParakeetSession final : public transcribe_session { // clear_result (the family owns its per-utterance audio scratch). std::vector stream_pcm_buffer; transcribe_run_params stream_run_params{}; + std::vector stream_language_block_mask; // Convert from cumulative samples so odd feed sizes do not lose fractions. int64_t stream_audio_input_samples = 0; diff --git a/src/transcribe-session.h b/src/transcribe-session.h index 3aa67b69..e8eda28a 100644 --- a/src/transcribe-session.h +++ b/src/transcribe-session.h @@ -269,8 +269,10 @@ struct transcribe_session { // 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_language_owned; + std::string stream_target_language_owned; + std::vector stream_allowed_languages_owned; + std::vector stream_allowed_language_ptrs_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 dfeb5fb4..08cf8e64 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -298,6 +298,18 @@ int timestamp_rank(transcribe_timestamp_kind k) { // streaming-begin rejects unconditionally in v1), so each caller // applies its own translate check before reaching this helper. transcribe_status validate_run_params_common(const transcribe_session * session, const transcribe_run_params * params) { + if (params->language == nullptr && + params->struct_size >= offsetof(transcribe_run_params, allowed_languages) + sizeof(params->allowed_languages)) { + if (params->n_allowed_languages < 0 || + (params->n_allowed_languages > 0 && params->allowed_languages == nullptr)) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + for (int i = 0; i < params->n_allowed_languages; ++i) { + if (params->allowed_languages[i] == nullptr || *params->allowed_languages[i] == '\0') { + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + } // Raw-validate every enum field before its first enum-typed load (see // enum_field_raw). Once a field passes here, downstream typed reads — // including the per-family handlers' — are defined. @@ -578,11 +590,13 @@ extern "C" void transcribe_run_params_init(struct transcribe_run_params * p) { return; } std::memset(p, 0, sizeof(*p)); - p->struct_size = sizeof(*p); - p->spec_k_drafts = -1; // family default + p->struct_size = sizeof(*p); + p->spec_k_drafts = -1; // family default + p->n_allowed_languages = 0; + p->allowed_languages = nullptr; // Default to AUTO (richest output compatible with the model and selected // run tasks, resolved per-family) rather than the memset NONE. - p->timestamps = TRANSCRIBE_TIMESTAMPS_AUTO; + p->timestamps = TRANSCRIBE_TIMESTAMPS_AUTO; } extern "C" void transcribe_stream_params_init(struct transcribe_stream_params * p) { @@ -735,8 +749,8 @@ namespace { // library-side prefix do NOT raise this value. #define TRANSCRIBE_FIELD_END(type, field) (offsetof(type, field) + sizeof(((type *) 0)->field)) -constexpr size_t k_min_model_params_size = TRANSCRIBE_FIELD_END(transcribe_model_load_params, device); -constexpr size_t k_min_context_params_size = TRANSCRIBE_FIELD_END(transcribe_session_params, kv_type); +constexpr size_t k_min_model_params_size = TRANSCRIBE_FIELD_END(transcribe_model_load_params, device); +constexpr size_t k_min_context_params_size = TRANSCRIBE_FIELD_END(transcribe_session_params, kv_type); // run_params is the one 0.2.0 exception to the append-only rule: `diarize` // was inserted mid-struct, shifting every field from `language` on by 8 // bytes. A 0.1-layout caller's sizeof (64) equals FIELD_END(family) in the @@ -744,9 +758,10 @@ constexpr size_t k_min_context_params_size = TRANSCRIBE_FIELD_END(trans // read their pointers as enums. Require through spec_k_drafts so every // 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_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_min_run_params_size = TRANSCRIBE_FIELD_END(transcribe_run_params, spec_k_drafts); +constexpr size_t k_run_params_allowed_languages_size = TRANSCRIBE_FIELD_END(transcribe_run_params, allowed_languages); +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 = TRANSCRIBE_FIELD_END(transcribe_stream_params, stable_prefix_agreement_n); constexpr size_t k_min_stream_update_size = TRANSCRIBE_FIELD_END(transcribe_stream_update, buffered_ms); @@ -1841,6 +1856,14 @@ 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_allowed_languages_owned.clear(); + if (run_params->language == nullptr && has_field(run_params->struct_size, k_run_params_allowed_languages_size) && + run_params->n_allowed_languages > 0 && run_params->allowed_languages != nullptr) { + session->stream_allowed_languages_owned.reserve(static_cast(run_params->n_allowed_languages)); + for (int i = 0; i < run_params->n_allowed_languages; ++i) { + session->stream_allowed_languages_owned.emplace_back(run_params->allowed_languages[i]); + } + } // 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 +1878,15 @@ 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; + auto & allowed_language_ptrs = session->stream_allowed_language_ptrs_owned; + allowed_language_ptrs.clear(); + allowed_language_ptrs.reserve(session->stream_allowed_languages_owned.size()); + for (const auto & lang : session->stream_allowed_languages_owned) { + allowed_language_ptrs.push_back(lang.c_str()); + } + run_params_owned.n_allowed_languages = static_cast(allowed_language_ptrs.size()); + run_params_owned.allowed_languages = allowed_language_ptrs.empty() ? nullptr : allowed_language_ptrs.data(); + run_params_owned.family = nullptr; const transcribe_status st = session->model->arch->stream_begin(session, &run_params_owned, stream_params); if (st != TRANSCRIBE_OK) { diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 1e2318a0..a7312b7a 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -170,6 +170,12 @@ transcribe_apply_warnings(transcribe_parakeet_chunked_limited_with_rc_mask_unit) add_test(NAME transcribe_parakeet_chunked_limited_with_rc_mask_unit COMMAND transcribe_parakeet_chunked_limited_with_rc_mask_unit) +add_executable(transcribe_parakeet_language_mask_unit parakeet_language_mask_unit.cpp) +target_link_libraries(transcribe_parakeet_language_mask_unit PRIVATE transcribe ggml) +target_include_directories(transcribe_parakeet_language_mask_unit PRIVATE ${CMAKE_SOURCE_DIR}/src) +transcribe_apply_warnings(transcribe_parakeet_language_mask_unit) +add_test(NAME transcribe_parakeet_language_mask_unit COMMAND transcribe_parakeet_language_mask_unit) + # ----------------------------------------------------------------------------- # Variable-length batch mask and causal subsampling unit test # ----------------------------------------------------------------------------- @@ -709,6 +715,18 @@ endif() # of failing the suite). if(TRANSCRIBE_BUILD_REAL_MODEL_TESTS) + add_executable(transcribe_nemotron_allowlist_real_stream + nemotron_allowlist_real_stream.cpp + ${CMAKE_SOURCE_DIR}/examples/common/wav.cpp) + target_link_libraries(transcribe_nemotron_allowlist_real_stream PRIVATE transcribe) + target_include_directories(transcribe_nemotron_allowlist_real_stream PRIVATE + ${CMAKE_SOURCE_DIR}/examples/common) + target_compile_definitions(transcribe_nemotron_allowlist_real_stream PRIVATE + "TRANSCRIBE_TEST_WAV_FILE=\"${CMAKE_SOURCE_DIR}/samples/jfk.wav\"") + transcribe_apply_warnings(transcribe_nemotron_allowlist_real_stream) + add_test(NAME transcribe_nemotron_allowlist_real_stream COMMAND transcribe_nemotron_allowlist_real_stream) + set_tests_properties(transcribe_nemotron_allowlist_real_stream PROPERTIES SKIP_RETURN_CODE 77) + add_executable(transcribe_parakeet_real_smoke parakeet_real_smoke.cpp) diff --git a/tests/api_smoke.c b/tests/api_smoke.c index ce776b72..d0947805 100644 --- a/tests/api_smoke.c +++ b/tests/api_smoke.c @@ -237,6 +237,8 @@ static void test_init_macros(void) { CHECK(rp_macro.itn == TRANSCRIBE_ITN_MODE_DEFAULT); CHECK(rp_macro.language == NULL); CHECK(rp_macro.target_language == NULL); + CHECK(rp_macro.n_allowed_languages == 0); + CHECK(rp_macro.allowed_languages == NULL); CHECK(rp_macro.keep_special_tags == false); CHECK(rp_macro.family == NULL); diff --git a/tests/nemotron_allowlist_real_stream.cpp b/tests/nemotron_allowlist_real_stream.cpp new file mode 100644 index 00000000..cf1166e1 --- /dev/null +++ b/tests/nemotron_allowlist_real_stream.cpp @@ -0,0 +1,78 @@ +#include "transcribe.h" +#include "wav.h" + +#include +#include +#include +#include +#include + +#ifndef TRANSCRIBE_TEST_WAV_FILE +# error "TRANSCRIBE_TEST_WAV_FILE must be defined" +#endif + +int main() { + const char * path = std::getenv("TRANSCRIBE_NEMOTRON_GGUF"); + if (path == nullptr || *path == '\0') { + return 77; + } + + transcribe_model_load_params load_params; + transcribe_model_load_params_init(&load_params); + transcribe_model * model = nullptr; + if (transcribe_model_load_file(path, &load_params, &model) != TRANSCRIBE_OK) { + return 1; + } + transcribe_session * session = nullptr; + if (transcribe_session_init(model, nullptr, &session) != TRANSCRIBE_OK) { + transcribe_model_free(model); + return 2; + } + + std::vector pcm; + std::string error; + if (!transcribe_cli::load_wav_mono_16k(TRANSCRIBE_TEST_WAV_FILE, pcm, error)) { + std::fprintf(stderr, "%s\n", error.c_str()); + transcribe_session_free(session); + transcribe_model_free(model); + return 3; + } + + transcribe_run_params run; + transcribe_run_params_init(&run); + const char * allowed[] = { "en-US", "de-DE" }; + run.allowed_languages = allowed; + run.n_allowed_languages = 2; + run.keep_special_tags = true; + + transcribe_status status = transcribe_stream_begin(session, &run, nullptr); + for (size_t offset = 0; status == TRANSCRIBE_OK && offset < pcm.size(); offset += 16000) { + const int count = static_cast(std::min(16000, pcm.size() - offset)); + status = transcribe_stream_feed(session, pcm.data() + offset, count, nullptr); + } + if (status == TRANSCRIBE_OK) { + status = transcribe_stream_finalize(session, nullptr); + } + const std::string raw = status == TRANSCRIBE_OK ? transcribe_raw_text(session) : ""; + if (status != TRANSCRIBE_OK || + (raw.find("") == std::string::npos && raw.find("") == std::string::npos)) { + std::fprintf(stderr, "stream failed or emitted no allowed language tag: %s, raw=%s\n", + transcribe_status_string(status), raw.c_str()); + status = TRANSCRIBE_ERR_BACKEND; + } + + transcribe_capabilities caps; + transcribe_capabilities_init(&caps); + if (transcribe_model_get_capabilities(model, &caps) == TRANSCRIBE_OK) { + for (int i = 0; i < caps.n_languages; ++i) { + const std::string code = caps.languages[i]; + if (code != "en-US" && code != "de-DE" && raw.find("<" + code + ">") != std::string::npos) { + status = TRANSCRIBE_ERR_BACKEND; + } + } + } + + transcribe_session_free(session); + transcribe_model_free(model); + return status == TRANSCRIBE_OK ? 0 : 4; +} diff --git a/tests/parakeet_language_mask_unit.cpp b/tests/parakeet_language_mask_unit.cpp new file mode 100644 index 00000000..8c9c9429 --- /dev/null +++ b/tests/parakeet_language_mask_unit.cpp @@ -0,0 +1,60 @@ +#include "arch/parakeet/decoder.h" +#include "arch/parakeet/parakeet.h" + +#include +#include + +int main() { + transcribe::parakeet::ParakeetModel model; + model.hparams.has_prompt = true; + model.host_decoder.head_kind = transcribe::parakeet::HostHeadKind::RNNT; + model.host_decoder.n_vocab = 6; + model.set_languages({ "en-US", "de-DE", "fr-FR" }); + if (model.tok.load_decode_only_raw_bytes({ "word", "", "", "", "", "" }) != + TRANSCRIBE_OK) return 10; + + transcribe_run_params params; + transcribe_run_params_init(¶ms); + std::vector mask; + using transcribe::parakeet::resolve_language_block_mask; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK || !mask.empty()) return 11; + const char * allowed[] = { "en-US", "de-DE" }; + params.allowed_languages = allowed; + params.n_allowed_languages = 2; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK || + mask.size() != 6 || mask[0] || mask[1] || mask[2] || !mask[3] || !mask[4] || mask[5]) return 12; + params.language = "de-DE"; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK || !mask.empty()) return 13; + params.language = "en-US"; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK || !mask.empty()) return 14; + params.language = nullptr; + const char * unknown[] = { "xx-ZZ" }; + params.allowed_languages = unknown; + params.n_allowed_languages = 1; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE) return 15; + params.n_allowed_languages = 0; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK || !mask.empty()) return 16; + + params.allowed_languages = allowed; + params.n_allowed_languages = 2; + if (resolve_language_block_mask(&model, ¶ms, mask) != TRANSCRIBE_OK) return 17; + + // 0 = ordinary token, 1 = English tag, 2 = German tag, + // 3 = French tag, 4 = unadvertised language tag, 5 = blank. + const float scores[] = { 1.0f, 6.0f, 7.0f, 9.0f, 0.0f, 0.0f }; + using transcribe::parakeet::argmax_language_masked; + if (argmax_language_masked(scores, 6, {}) != 3) return 1; + + if (argmax_language_masked(scores, 6, mask) != 2) return 2; + // Apply the same constraint on a later streaming decoder step. + const float later_scores[] = { 2.0f, 8.0f, 7.0f, 10.0f, 0.0f, 0.0f }; + if (argmax_language_masked(later_scores, 6, mask) != 1) return 3; + + // An ordinary token or blank still wins on its own raw score. + const float ordinary_scores[] = { 11.0f, 6.0f, 7.0f, 9.0f, 0.0f, 0.0f }; + if (argmax_language_masked(ordinary_scores, 6, mask) != 0) return 4; + const float blank_scores[] = { 1.0f, 6.0f, 7.0f, 9.0f, 0.0f, 12.0f }; + if (argmax_language_masked(blank_scores, 6, mask) != 5) return 5; + if (argmax_language_masked(scores, 6, { 1, 1, 1, 1, 1, 1 }) != -1) return 6; + return 0; +} diff --git a/tests/parakeet_stream_ext_reject_unit.cpp b/tests/parakeet_stream_ext_reject_unit.cpp index 3aa8a5a8..ca608994 100644 --- a/tests/parakeet_stream_ext_reject_unit.cpp +++ b/tests/parakeet_stream_ext_reject_unit.cpp @@ -133,6 +133,25 @@ void test_cache_aware_rejects_sub_sentinel() { transcribe_model_free(model); } +void test_stream_allowlist_rejects_model_without_language_tags() { + struct transcribe_model * model = nullptr; + struct transcribe_session * ctx = nullptr; + if (!load_and_init("tokenizer_minimal_streaming_cache_aware.gguf", &model, &ctx)) { + return; + } + + transcribe_run_params rp; + transcribe_run_params_init(&rp); + const char * allowed[] = { "en-US", "de-DE" }; + rp.allowed_languages = allowed; + rp.n_allowed_languages = 2; + CHECK(transcribe_stream_begin(ctx, &rp, nullptr) == TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE); + CHECK(transcribe_stream_get_state(ctx) == TRANSCRIBE_STREAM_IDLE); + + transcribe_session_free(ctx); + transcribe_model_free(model); +} + // Buffered path. Each of {left,chunk,right}_ms < -1 must return // INVALID_ARG. Tests each field independently so a bug that misses one // of the three rejects is caught. @@ -225,6 +244,7 @@ void test_buffered_zero_is_real_value() { int main() { test_cache_aware_rejects_sub_sentinel(); + test_stream_allowlist_rejects_model_without_language_tags(); test_buffered_rejects_sub_sentinel(); test_buffered_zero_is_real_value();