use std::{
cell::Cell,
path::PathBuf,
sync::{Mutex, atomic::AtomicBool},
};
use super::*;
use crate::audio::whisper::{
backend::{InferenceBackend, mock::MockBackend},
decode::sampler::GreedyTokenSampler,
options::DecodingOptions,
result::TranscriptionTimings,
tokenizer::{SpecialTokens, WhisperTokenizer},
};
fn tiny_tokenizer() -> WhisperTokenizer {
let root = std::env::var_os("WHISPERKIT_TEST_MODELS")
.map_or_else(crate::tests::models_root, PathBuf::from);
WhisperTokenizer::from_folder(root.join("tokenizers/whisper-tiny")).unwrap()
}
fn special() -> SpecialTokens {
SpecialTokens::whisper_defaults()
}
fn default_prompt(s: &SpecialTokens) -> Vec<u32> {
vec![
s.start_of_transcript_token(),
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
]
}
fn run_mock(
mock: &MockBackend,
prompt: &[u32],
options: &DecodingOptions,
tokenizer: &WhisperTokenizer,
) -> crate::audio::whisper::result::DecodingResult {
let encoded = mock
.encode(&mock.extract_features(&[0.0; 16]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let mut sampler = GreedyTokenSampler::new(options.temperature(), special().end_token(), options);
let mut timings = TranscriptionTimings::new();
decode_text(
mock,
&encoded,
&mut state,
prompt,
&mut sampler,
options,
tokenizer,
&mut timings,
&AtomicBool::new(false),
&Cell::new(None),
None,
)
.unwrap()
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn prefill_tokens_multilingual_default_shape() {
let t = tiny_tokenizer();
let s = special();
let options = DecodingOptions::new();
assert_eq!(prefill_tokens(&options, &t, true), default_prompt(&s));
let options = DecodingOptions::new().with_without_timestamps();
assert_eq!(
prefill_tokens(&options, &t, true).last(),
Some(&s.no_timestamps_token())
);
let options = DecodingOptions::new();
assert_eq!(
prefill_tokens(&options, &t, false),
vec![s.start_of_transcript_token(), s.time_token_begin()]
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn prefill_prompt_tokens_truncate_and_prepend_previous() {
let t = tiny_tokenizer();
let s = special();
let long_prompt: Vec<u32> = (0..200u32).chain([s.end_token()]).collect();
let options = DecodingOptions::new().with_prompt_tokens(long_prompt);
let tokens = prefill_tokens(&options, &t, true);
assert_eq!(tokens[0], s.start_of_previous_token());
assert_eq!(tokens[1], 90);
assert_eq!(tokens[110], 199);
assert_eq!(tokens[111], s.start_of_transcript_token());
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn loop_forces_prompt_then_samples_to_eot() {
let t = tiny_tokenizer();
let s = special();
let prompt = default_prompt(&s);
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(), s.transcribe_token(), s.time_token_begin(), 2425, 1002,
s.time_token_begin() + 50,
s.end_token(),
]);
let result = run_mock(&mock, &prompt, &DecodingOptions::new(), &t);
let expected: Vec<u32> = prompt
.iter()
.copied()
.chain([2425, 1002, s.time_token_begin() + 50, s.end_token()])
.collect();
assert_eq!(result.tokens_slice(), expected.as_slice());
assert!(result.avg_logprob() < 0.0); assert_eq!(result.temperature(), 0.0);
let counters = mock.counters();
assert_eq!(counters.decode_steps(), 7);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn negative_temperature_decode_with_timestamp_filter_does_not_panic() {
let t = tiny_tokenizer();
let s = special();
let mut mock = MockBackend::new();
mock.push_token_steps(&[100u32; 24]);
let options = DecodingOptions::new()
.with_temperature(-0.2)
.with_sample_length(8);
let result = run_mock(&mock, &default_prompt(&s), &options, &t);
assert!(
result.avg_logprob().is_finite(),
"a negative-temperature decode over masked logits must finish with a finite avg log-prob"
);
assert!(
(result.temperature() - (-0.2)).abs() < 1e-6,
"the accepted temperature is the configured negative one"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn language_observed_only_for_a_predicted_language_token() {
let t = tiny_tokenizer();
let s = special();
let es = t.token_to_id("<|es|>").unwrap();
let mut mock = MockBackend::new();
mock.push_token_steps(&[
es,
s.transcribe_token(),
s.time_token_begin(),
2425,
1002,
s.time_token_begin() + 50,
s.end_token(),
]);
let prompt_es = vec![
s.start_of_transcript_token(),
es,
s.transcribe_token(),
s.time_token_begin(),
];
let configured = run_mock(
&mock,
&prompt_es,
&DecodingOptions::new().with_language("es"),
&t,
);
assert_eq!(configured.language(), "es");
assert_eq!(
configured.observed_language(),
None,
"a configured language is copied, not detected"
);
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(), s.transcribe_token(), s.time_token_begin(), 2425, 1002,
s.time_token_begin() + 50,
s.end_token(),
]);
let forced = run_mock(&mock, &default_prompt(&s), &DecodingOptions::new(), &t);
assert_eq!(
forced.language(),
"en",
"forced <|en|> is still the display language"
);
assert_eq!(
forced.observed_language(),
None,
"a FORCED prefill <|en|> is an input, not a detection -- never observed"
);
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(), s.transcribe_token(), s.no_timestamps_token(), es, 2425,
s.end_token(),
]);
let forced_en_predicts_es = run_mock(
&mock,
&[
s.start_of_transcript_token(),
s.english_token(),
s.transcribe_token(),
s.no_timestamps_token(),
],
&DecodingOptions::new().with_without_timestamps(),
&t,
);
assert_eq!(
forced_en_predicts_es.language(),
"en",
"the DISPLAY language is the forced-prefill <|en|>, first in the whole slice"
);
assert_eq!(
forced_en_predicts_es.observed_language(),
Some("es"),
"the OBSERVATION is the PREDICTED <|es|>, never the forced display <|en|>"
);
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(), s.transcribe_token(), s.no_timestamps_token(), es, 2425,
s.end_token(),
]);
let configured_en_predicts_es = run_mock(
&mock,
&[
s.start_of_transcript_token(),
s.english_token(),
s.transcribe_token(),
s.no_timestamps_token(),
],
&DecodingOptions::new()
.with_without_timestamps()
.with_language("en"),
&t,
);
assert_eq!(
configured_en_predicts_es.language(),
"en",
"the DISPLAY language is the configured/forced <|en|>, unchanged by the fix"
);
assert_eq!(
configured_en_predicts_es.observed_language(),
Some("es"),
"a configured language must NOT suppress a genuinely PREDICTED <|es|> observation"
);
let mut mock = MockBackend::new();
mock.push_token_steps(&[es, 2425, s.end_token()]);
let predicted = run_mock(
&mock,
&[s.start_of_transcript_token()],
&DecodingOptions::new().with_without_timestamps(),
&t,
);
assert_eq!(
predicted.language(),
"es",
"the predicted <|es|> is the display language"
);
assert_eq!(
predicted.observed_language(),
Some("es"),
"a <|lang|> token PREDICTED after the prompt is a genuine detection"
);
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.transcribe_token(),
s.time_token_begin(),
2425,
1002,
s.time_token_begin() + 50,
s.end_token(),
]);
let prompt_no_lang = vec![
s.start_of_transcript_token(),
s.transcribe_token(),
s.time_token_begin(),
];
let fallback = run_mock(&mock, &prompt_no_lang, &DecodingOptions::new(), &t);
assert_eq!(
fallback.language(),
crate::audio::whisper::constants::DEFAULT_LANGUAGE_CODE
);
assert_eq!(
fallback.observed_language(),
None,
"the \"en\" fallback is a default, not a detection"
);
let mut mock = MockBackend::new();
let mut low_confidence_es = vec![0.0_f32; mock.dims().vocab()];
low_confidence_es[es as usize] = 1.0; mock.push_step(low_confidence_es);
let options = DecodingOptions::new()
.with_without_timestamps()
.with_temperature_fallback_count(0);
let encoded = mock
.encode(&mock.extract_features(&[0.0; 16]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let mut sampler = GreedyTokenSampler::new(options.temperature(), s.end_token(), &options);
let mut timings = TranscriptionTimings::new();
let observed_cell: Cell<Option<u32>> = Cell::new(None);
let low_first = decode_text(
&mock,
&encoded,
&mut state,
&[s.start_of_transcript_token()],
&mut sampler,
&options,
&t,
&mut timings,
&AtomicBool::new(false),
&observed_cell,
None,
)
.unwrap();
assert!(
low_first.first_token_log_prob() < -1.5,
"the sampled <|es|> is below the -1.5 first-token threshold, got {}",
low_first.first_token_log_prob(),
);
assert_eq!(
mock.counters().decode_steps(),
1,
"the below-threshold first token completes the decode on the first step",
);
assert_eq!(
low_first.observed_language(),
Some("es"),
"a first token below the threshold still latches its PREDICTED language onto the DecodingResult",
);
assert_eq!(
observed_cell.get(),
Some(es),
"and into the cell the attempt sink carries into the task facts -- latched BEFORE the break",
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn zero_iteration_decode_forces_english_but_observes_nothing() {
let t = tiny_tokenizer();
let s = special();
let mut mock = MockBackend::new();
mock.push_token_step(2425); let result = run_mock(
&mock,
&default_prompt(&s),
&DecodingOptions::new().with_sample_length(0),
&t,
);
assert_eq!(mock.counters().decode_steps(), 0, "zero decoder steps ran");
assert_eq!(
result.language(),
"en",
"the display language is the forced <|en|>"
);
assert_eq!(
result.observed_language(),
None,
"a zero-iteration decode predicts nothing, so it observes no language"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn last_prefill_timestamp_keeps_model_prediction() {
let t = tiny_tokenizer();
let s = special();
let prompt = default_prompt(&s); let predicted_ts = s.time_token_begin() + 25; let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
predicted_ts, 100, s.end_token(),
]);
let result = run_mock(&mock, &prompt, &DecodingOptions::new(), &t);
assert_eq!(
result.tokens_slice()[3],
predicted_ts,
"model timestamp kept"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn first_token_logprob_below_threshold_stops_immediately() {
let t = tiny_tokenizer();
let s = special();
let mut mock = MockBackend::new();
mock.push_step(vec![0.0; 51865]);
let result = run_mock(&mock, &default_prompt(&s), &DecodingOptions::new(), &t);
assert_eq!(
mock.counters().decode_steps(),
1,
"stopped after first step"
);
assert!(result.first_token_log_prob() < -1.5);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn early_stop_flag_breaks_loop_and_callback_sets_it() {
let t = tiny_tokenizer();
let s = special();
let mut mock = MockBackend::new();
mock.push_token_steps(&[
s.english_token(),
s.transcribe_token(),
s.time_token_begin(),
100,
101,
102,
103,
104,
s.end_token(),
]);
let steps_seen = Mutex::new(0usize);
let encoded = mock
.encode(&mock.extract_features(&[0.0; 4]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let options = DecodingOptions::new();
let mut sampler = GreedyTokenSampler::new(0.0, s.end_token(), &options);
let mut timings = TranscriptionTimings::new();
let callback: &(
dyn Fn(&crate::audio::whisper::result::TranscriptionProgress) -> Option<bool> + Sync
) = &|_progress| {
let mut seen = steps_seen.lock().unwrap();
*seen += 1;
Some(*seen < 5)
};
let result = decode_text(
&mock,
&encoded,
&mut state,
&default_prompt(&s),
&mut sampler,
&options,
&t,
&mut timings,
&AtomicBool::new(false),
&Cell::new(None),
Some(callback),
)
.unwrap();
assert!(
result.tokens_slice().len() < 4 + 5 + 1,
"stopped before scripted EOT"
);
assert_eq!(
*result.tokens_slice().last().unwrap(),
s.end_token(),
"finalize appends EOT"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn detect_language_single_step_and_resets_state() {
let t = tiny_tokenizer();
let es_token = t.token_to_id("<|es|>").unwrap();
let mut mock = MockBackend::new();
mock.push_token_step(es_token);
mock.push_token_step(es_token); let encoded = mock
.encode(&mock.extract_features(&[0.0; 4]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let mut timings = TranscriptionTimings::new();
let mut sampler = GreedyTokenSampler::new(
0.0,
SpecialTokens::whisper_defaults().end_token(),
&DecodingOptions::new(),
);
let result =
detect_language(&mock, &encoded, &mut state, &t, &mut sampler, &mut timings).unwrap();
assert_eq!(result.language(), "es");
assert!(
result
.language_probs_slice()
.iter()
.any(|(code, _)| code == "es")
);
assert!(result.tokens_slice().is_empty()); assert_eq!(mock.counters().resets(), 1, "state reset after probe");
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn detect_language_resets_state_even_when_the_step_fails() {
let t = tiny_tokenizer();
let mock = MockBackend::new(); let encoded = mock
.encode(&mock.extract_features(&[0.0; 4]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let mut timings = TranscriptionTimings::new();
let s = SpecialTokens::whisper_defaults();
let mut sampler = GreedyTokenSampler::new(0.0, s.end_token(), &DecodingOptions::new());
let err =
detect_language(&mock, &encoded, &mut state, &t, &mut sampler, &mut timings).unwrap_err();
assert!(matches!(err, DecodeError::Backend(_)));
assert_eq!(mock.counters().resets(), 1, "state reset despite the error");
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn detect_language_samples_through_the_callers_sampler() {
let t = tiny_tokenizer();
let es = t.token_to_id("<|es|>").unwrap();
let de = t.token_to_id("<|de|>").unwrap();
let mut mock = MockBackend::new();
let mut logits = vec![0.0f32; crate::audio::whisper::backend::ModelDims::new().vocab()];
logits[es as usize] = 10.0;
logits[de as usize] = 9.5;
mock.push_step(logits.clone());
let encoded = mock
.encode(&mock.extract_features(&[0.0; 4]).unwrap())
.unwrap();
let mut state = mock.new_decoder_state().unwrap();
let mut timings = TranscriptionTimings::new();
let s = SpecialTokens::whisper_defaults();
let mut reference =
GreedyTokenSampler::new(0.7, s.end_token(), &DecodingOptions::new()).with_seed(7);
let filter =
crate::audio::whisper::decode::filter::LanguageLogitsFilter::new(t.all_language_tokens(), 0);
let mut reference_logits = logits;
filter
.filter(&mut reference_logits, &[s.start_of_transcript_token()])
.expect("every id here is inside the vocabulary");
let expected = reference.sample(&reference_logits);
let expected_language = t
.language_for_token(expected.token())
.expect("draw lands on a language token");
let mut probe_sampler =
GreedyTokenSampler::new(0.7, s.end_token(), &DecodingOptions::new()).with_seed(7);
let result = detect_language(
&mock,
&encoded,
&mut state,
&t,
&mut probe_sampler,
&mut timings,
)
.unwrap();
assert_eq!(
result.language(),
expected_language,
"probe = the caller's draw"
);
let plain = [1.0f32, 2.0, 3.0, 2.5];
for _ in 0..5 {
assert_eq!(
probe_sampler.sample(&plain).token(),
reference.sample(&plain).token(),
"streams diverged: the probe consumed a different number of draws"
);
}
}