use std::path::PathBuf;
use super::*;
use crate::audio::whisper::{
backend::AlignmentView,
constants::{APPEND_PUNCTUATION, MAX_TOKEN_CONTEXT, PREPEND_PUNCTUATION},
options::DecodingOptions,
result::{DecodingResult, WordTiming},
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 ts(index: u32) -> u32 {
SpecialTokens::whisper_defaults().time_token_begin() + index
}
fn result_with_tokens(tokens: Vec<u32>, no_speech: f32, avg_logprob: f32) -> DecodingResult {
let log_probs: Vec<(u32, f32)> = tokens.iter().map(|&t| (t, -0.1)).collect();
let mut r = DecodingResult::new();
r.set_tokens(tokens)
.set_token_log_probs(log_probs)
.set_no_speech_prob(no_speech)
.set_avg_logprob(avg_logprob);
r
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn silence_skips_full_segment() {
let t = tiny_tokenizer();
let r = result_with_tokens(vec![], 0.9, -1.5);
let (seek, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, 16_000, 480_000, &t).unwrap();
assert_eq!(seek, 16_000 + 480_000);
assert!(segments.is_none());
let r = result_with_tokens(vec![50258, 100, ts(0), ts(50)], 0.9, -0.2);
let (_, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, 0, 480_000, &t).unwrap();
assert!(segments.is_some());
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn consecutive_timestamps_slice_into_segments_and_seek_to_last() {
let t = tiny_tokenizer();
let hello = 15947u32; let world = 1002u32;
let tokens = vec![
ts(0),
hello,
ts(50),
ts(50),
world,
ts(100),
t.special_tokens().end_token(),
];
let r = result_with_tokens(tokens, 0.0, -0.2);
let (seek, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 3, 32_000, 480_000, &t).unwrap();
let segments = segments.unwrap();
assert_eq!(segments.len(), 2);
assert_eq!(segments[0].id(), 3); assert!((segments[0].start() - 2.0).abs() < 1e-4);
assert!((segments[0].end() - 3.0).abs() < 1e-4);
assert!((segments[1].start() - 3.0).abs() < 1e-4);
assert!((segments[1].end() - 4.0).abs() < 1e-4);
assert_eq!(segments[0].seek(), 32_000);
assert_eq!(seek, 32_000 + 32_000);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn single_timestamp_ending_appends_final_slice() {
let t = tiny_tokenizer();
let tokens = vec![ts(0), 100, ts(50), ts(50), 101, ts(75), 102];
let r = result_with_tokens(tokens, 0.0, -0.2);
let (seek, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, 0, 480_000, &t).unwrap();
let segments = segments.unwrap();
assert_eq!(segments.len(), 2);
assert!((segments[1].end() - 1.5).abs() < 1e-4);
assert_eq!(seek, (1.5 * 16_000.0) as usize);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn no_consecutive_timestamps_lumps_window_and_refines_duration() {
let t = tiny_tokenizer();
let tokens = vec![ts(0), 100, 101, ts(150)]; let r = result_with_tokens(tokens, 0.0, -0.2);
let (seek, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, 0, 160_000, &t).unwrap();
let segments = segments.unwrap();
assert_eq!(segments.len(), 1);
assert!((segments[0].start() - 0.0).abs() < 1e-4);
assert!(
(segments[0].end() - 3.0).abs() < 1e-4,
"refined by trailing timestamp"
);
assert_eq!(seek, 160_000);
let r = result_with_tokens(vec![ts(0), 100, 101], 0.0, -0.2);
let (_, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, 0, 160_000, &t).unwrap();
assert!((segments.unwrap()[0].end() - 10.0).abs() < 1e-4);
}
#[test]
fn dtw_diagonal_identity() {
#[rustfmt::skip]
let matrix = [
1.0f32, 0.0, 0.0,
0.0, 1.0, 0.0,
0.0, 0.0, 1.0,
];
let view = AlignmentView::new(&matrix, 3, 3);
let path = dynamic_time_warping(&view).unwrap();
assert_eq!(path.text_indices_slice(), &[0, 1, 1, 2, 2]);
assert_eq!(path.time_indices_slice(), &[0, 0, 1, 1, 2]);
}
#[test]
fn dtw_wide_matrix_repeats_text_indices() {
#[rustfmt::skip]
let matrix = [
0.9f32, 0.9, 0.1, 0.1,
0.1, 0.1, 0.9, 0.9,
];
let view = AlignmentView::new(&matrix, 2, 4);
let path = dynamic_time_warping(&view).unwrap();
assert_eq!(path.text_indices_slice(), &[0, 0, 1, 1, 1]);
assert_eq!(path.time_indices_slice(), &[0, 1, 1, 2, 3]);
}
#[test]
fn dtw_matches_swift_unit_test_ground_truth() {
#[rustfmt::skip]
let matrix = [
1.0f32, 1.0, 1.0,
5.0, 2.0, 1.0,
1.0, 5.0, 2.0,
];
let view = AlignmentView::new(&matrix, 3, 3);
let path = dynamic_time_warping(&view).unwrap();
assert_eq!(path.text_indices_slice(), &[0, 1, 1, 2, 2]);
assert_eq!(path.time_indices_slice(), &[0, 0, 1, 1, 2]);
}
#[test]
fn dtw_rejects_empty() {
let view = AlignmentView::new(&[], 0, 0);
assert!(dynamic_time_warping(&view).is_err());
}
fn word(text: &str, start: f32, end: f32) -> WordTiming {
WordTiming::new(text, vec![1], start, end, 0.9)
}
#[test]
fn merge_punctuations_english() {
let alignment = [
word(" Hey", 0.0, 0.2),
word(",", 0.2, 0.3),
word(" you", 0.3, 0.6),
word("!", 0.6, 0.7),
];
let merged = merge_punctuations(&alignment, PREPEND_PUNCTUATION, APPEND_PUNCTUATION);
let words: Vec<&str> = merged.iter().map(|w| w.word()).collect();
assert_eq!(words, vec![" Hey,", " you!"]);
assert_eq!(merged[0].tokens_slice().len(), 2, "tokens concatenated");
}
#[test]
fn merge_punctuations_prepended() {
let alignment = [
word(" \u{00bf}", 0.0, 0.1),
word("Que", 0.1, 0.4),
word("?", 0.4, 0.5),
];
let merged = merge_punctuations(&alignment, PREPEND_PUNCTUATION, APPEND_PUNCTUATION);
let words: Vec<&str> = merged.iter().map(|w| w.word()).collect();
assert_eq!(words, vec![" \u{00bf}Que?"]);
}
#[test]
fn merge_punctuations_ignores_whitespace_only_words() {
let alignment = [
word(" a", 0.0, 0.2),
word(" ", 0.2, 0.3),
word(" b", 0.3, 0.6),
];
let merged = merge_punctuations(&alignment, PREPEND_PUNCTUATION, APPEND_PUNCTUATION);
let words: Vec<&str> = merged.iter().map(|w| w.word()).collect();
assert_eq!(words, vec![" a", " ", " b"], "no merges, nothing dropped");
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn find_alignment_produces_monotonic_word_timings() {
let t = tiny_tokenizer();
let ids = t.encode(" Hello world again").unwrap();
let cols = 100usize;
let mut matrix = vec![0.0f32; ids.len() * cols];
for (i, row) in matrix.chunks_mut(cols).enumerate() {
row[i * 10] = 1.0;
}
let view = AlignmentView::new(&matrix, ids.len(), cols);
let log_probs = vec![-0.2f32; ids.len()];
let words = find_alignment(&ids, &view, &log_probs, &t, "en", WordGrouping::FineGrained).unwrap();
assert!(!words.is_empty());
for pair in words.windows(2) {
assert!(pair[0].end() <= pair[1].start() + 1e-4, "monotonic timings");
}
for w in &words {
assert!((0.0..=1.0).contains(&w.probability()));
}
}
#[test]
fn dtw_add_before_compare_matches_swift_rounding_ties() {
#[rustfmt::skip]
let matrix = [
0.0f32, 1.0,
0.0, 1.0e30,
];
let view = AlignmentView::new(&matrix, 2, 2);
let path = dynamic_time_warping(&view).unwrap();
assert_eq!(path.text_indices_slice(), &[0, 1, 1]);
assert_eq!(path.time_indices_slice(), &[0, 0, 1]);
}
#[test]
fn duration_constraints_take_the_capped_upper_median() {
let alignment = [
word("a", 0.0, 0.2),
word("b", 0.2, 0.6),
word("c", 0.6, 1.2),
word("z", 1.2, 1.2), ];
let constraints = calculate_word_duration_constraints(&alignment);
assert!((constraints.median() - 0.4).abs() < 1e-6);
assert!((constraints.max_duration() - 0.8).abs() < 1e-6);
let long = [word("a", 0.0, 1.0), word("b", 1.0, 2.0)];
let constraints = calculate_word_duration_constraints(&long);
assert!((constraints.median() - 0.7).abs() < 1e-6);
assert!((constraints.max_duration() - 1.4).abs() < 1e-6);
let constraints = calculate_word_duration_constraints(&[]);
assert_eq!(constraints.median(), 0.0);
assert_eq!(constraints.max_duration(), 0.0);
}
#[test]
fn even_count_median_takes_the_upper_middle_value() {
let alignment = [
word("a", 0.0, 0.2),
word("b", 0.2, 0.6),
word("c", 0.6, 1.2),
word("d", 1.2, 2.0),
];
let constraints = calculate_word_duration_constraints(&alignment);
assert!((constraints.median() - 0.6).abs() < 1e-6);
assert!((constraints.max_duration() - 1.2).abs() < 1e-6);
}
#[test]
fn truncation_fires_only_at_sentence_boundaries() {
let alignment = vec![word(" ok", 0.0, 0.3), word(".", 0.3, 2.0)];
let out = truncate_long_words_at_sentence_boundaries(alignment, 0.5);
assert!((out[1].end() - 0.8).abs() < 1e-6);
let alignment = vec![word("!", 0.0, 0.1), word(" Next", 0.1, 2.0)];
let out = truncate_long_words_at_sentence_boundaries(alignment, 0.5);
assert!((out[1].start() - 1.5).abs() < 1e-6);
let alignment = vec![
word(" a", 0.0, 0.1),
word(" long", 0.1, 3.0),
word(" .", 3.0, 6.0),
];
let out = truncate_long_words_at_sentence_boundaries(alignment, 0.5);
assert_eq!(out[1].end(), 3.0);
assert_eq!(out[2].end(), 6.0);
let alignment = vec![word("!", 0.0, 0.1), word(".", 0.1, 2.0)];
let out = truncate_long_words_at_sentence_boundaries(alignment, 0.5);
assert!(
(out[1].end() - 0.6).abs() < 1e-6,
"first branch: end = start + max"
);
assert_eq!(out[1].start(), 0.1, "second branch must not also fire");
let alignment = vec![word(".", 0.0, 5.0)];
let out = truncate_long_words_at_sentence_boundaries(alignment, 0.5);
assert_eq!(out[0].end(), 5.0);
}
#[test]
fn rounded_to_places_matches_swift_rounding() {
assert_eq!(rounded_to_places(1.234, 2), 1.23);
assert_eq!(rounded_to_places(1.235, 2), 1.24);
assert_eq!(rounded_to_places(-1.235, 2), -1.24);
}
fn plain_segment(tokens: Vec<u32>, start: f32, end: f32) -> TranscriptionSegment {
let mut segment = TranscriptionSegment::new();
segment.set_tokens(tokens).set_start(start).set_end(end);
segment
}
fn aligned(text: &str, tokens: Vec<u32>, start: f32, end: f32) -> WordTiming {
WordTiming::new(text, tokens, start, end, 0.9)
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn word_walk_assigns_words_and_pulls_short_words_back() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let segments = [plain_segment(vec![hello, world], 0.0, 2.0)];
let alignment = [
aligned(" Hello", vec![hello], 0.0, 0.5),
aligned(" world", vec![world], 1.9, 2.0),
];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
assert_eq!(updated.len(), 1);
let words = updated[0].words_slice();
assert_eq!(words.len(), 2);
assert!(
(words[1].start() - 1.6).abs() < 1e-4,
"0.1s word pulled back by median/2"
);
assert!((words[1].end() - 2.0).abs() < 1e-4);
assert!((updated[0].start() - 0.0).abs() < 1e-4);
assert!((updated[0].end() - 2.0).abs() < 1e-4);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn special_only_alignment_entries_are_skipped() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let segments = [plain_segment(vec![hello], 0.0, 1.0)];
let alignment = [
aligned("<|0.00|>", vec![s.time_token_begin()], 0.0, 0.0),
aligned(" Hello", vec![hello], 0.0, 0.5),
];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
let words = updated[0].words_slice();
assert_eq!(words.len(), 1);
assert_eq!(words[0].word(), " Hello");
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn seek_offset_and_pause_hack_apply() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let segments = [plain_segment(vec![hello], 2.0, 5.0)];
let alignment = [aligned(" Hello", vec![hello], 0.0, 3.0)];
let updated =
update_segments_with_word_timings(&segments, &alignment, 32_000, 0.0, 0.6, 1.2, &t).unwrap();
let words = updated[0].words_slice();
assert!((words[0].end() - 5.0).abs() < 1e-4, "offset applied");
assert!(
(words[0].start() - 3.8).abs() < 1e-4,
"pause-hack clamped the first word"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn word_index_cursor_is_shared_and_previous_segment_gap_pulls_first_word_back() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let segments = [
plain_segment(vec![hello], 0.0, 1.0),
plain_segment(vec![world], 2.0, 3.0),
];
let alignment = [
aligned(" Hello", vec![hello], 0.0, 1.0),
aligned(" world", vec![world], 2.4, 2.45),
];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
assert_eq!(updated.len(), 2);
assert!((updated[0].end() - 1.0).abs() < 1e-4);
let second_words = updated[1].words_slice();
assert_eq!(
second_words.len(),
1,
"cursor advanced past segment 0's word, not reused"
);
assert!(
(second_words[0].start() - 2.1).abs() < 1e-4,
"first word of segment 1 pulled back against segment 0's end"
);
assert!((second_words[0].end() - 2.45).abs() < 1e-4);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn second_word_too_long_triggers_boundary_resplit() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let segments = [plain_segment(vec![hello, world], 0.0, 10.0)];
let alignment = [
aligned(" Hello", vec![hello], 3.0, 3.3),
aligned(" world", vec![world], 3.3, 6.0),
];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
let words = updated[0].words_slice();
assert_eq!(words.len(), 2);
assert!((words[0].end() - 4.8).abs() < 1e-4, "resplit boundary");
assert!((words[1].start() - 4.8).abs() < 1e-4, "resplit boundary");
assert!(
(words[0].start() - 3.6).abs() < 1e-4,
"first word clamped after resplit"
);
assert!(
(words[1].end() - 6.0).abs() < 1e-4,
"second word's end untouched by resplit"
);
assert!((updated[0].start() - 3.6).abs() < 1e-4);
assert!((updated[0].end() - 6.0).abs() < 1e-4);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn segment_level_bounds_preferred_when_words_drift_far_from_segment() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let segments = [plain_segment(vec![hello], 3.0, 5.0)];
let alignment = [aligned(" Hello", vec![hello], 2.0, 10.0)];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 9.9, 2.5, 1.2, &t).unwrap();
let words = updated[0].words_slice();
assert_eq!(words.len(), 1);
assert!(
(words[0].start() - 3.0).abs() < 1e-4,
"segment start preferred"
);
assert!(
(words[0].end() - 5.5).abs() < 1e-4,
"word-anchored clamp term reads the live start"
);
assert!((updated[0].start() - 3.0).abs() < 1e-4);
assert!(
(updated[0].end() - 5.0).abs() < 1e-4,
"IF branch leaves segment.end"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn mixed_special_and_word_tokens_retokenize_the_surviving_ones() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let segments = [plain_segment(vec![hello], 0.0, 1.0)];
let alignment = [aligned(
"WRONG",
vec![s.time_token_begin(), hello],
0.0,
0.5,
)];
let updated =
update_segments_with_word_timings(&segments, &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
let words = updated[0].words_slice();
assert_eq!(words.len(), 1);
assert_eq!(
words[0].word(),
" Hello",
"retokenized from the surviving token, not `.word`"
);
assert_eq!(
words[0].tokens_slice().to_vec(),
vec![hello],
"special token filtered out of stored tokens too"
);
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn empty_segments_returns_empty_without_consuming_alignment() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let alignment = [aligned(" Hello", vec![hello], 0.0, 0.5)];
let updated = update_segments_with_word_timings(&[], &alignment, 0, 0.0, 0.6, 1.2, &t).unwrap();
assert!(updated.is_empty());
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn add_word_timestamps_attaches_merged_monotonic_words() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let tokens = vec![
s.time_token_begin(),
hello,
world,
s.time_token_begin() + 100,
];
let log_probs: Vec<(u32, f32)> = tokens.iter().map(|&tok| (tok, -0.2)).collect();
let mut segment = TranscriptionSegment::new();
segment
.set_tokens(tokens)
.set_token_log_probs(log_probs)
.set_start(0.0)
.set_end(2.0);
let cols = 150usize;
let mut weights = vec![0.0f32; 4 * cols];
for (i, row) in weights.chunks_mut(cols).enumerate() {
row[i * 25] = 1.0;
}
let view = AlignmentView::new(&weights, 4, cols);
let updated = add_word_timestamps(
&[segment],
&view,
&t,
"en",
WordGrouping::FineGrained,
AlignmentGather::Complete,
MAX_TOKEN_CONTEXT,
0,
crate::audio::whisper::constants::PREPEND_PUNCTUATION,
crate::audio::whisper::constants::APPEND_PUNCTUATION,
0.0,
)
.unwrap();
assert_eq!(updated.len(), 1);
let words = updated[0].words_slice();
assert!(!words.is_empty(), "text tokens produced word timings");
let joined: String = words.iter().map(|w| w.word()).collect();
assert_eq!(
crate::audio::whisper::text::normalized(&joined),
"hello world"
);
for pair in words.windows(2) {
assert!(
pair[0].start() <= pair[1].start() + 1e-4,
"monotonic starts"
);
}
for word in words {
assert!(word.end() >= word.start());
assert!((0.0..=1.0).contains(&word.probability()));
}
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn add_word_timestamps_zero_pads_missing_rows() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let mut segment = TranscriptionSegment::new();
segment
.set_tokens(vec![hello, world])
.set_token_log_probs(vec![(hello, -0.1), (world, -0.1)])
.set_start(0.0)
.set_end(1.0);
let weights = vec![1.0f32; 3]; let view = AlignmentView::new(&weights, 1, 3);
let updated = add_word_timestamps(
&[segment],
&view,
&t,
"en",
WordGrouping::FineGrained,
AlignmentGather::Complete,
MAX_TOKEN_CONTEXT,
0,
crate::audio::whisper::constants::PREPEND_PUNCTUATION,
crate::audio::whisper::constants::APPEND_PUNCTUATION,
0.0,
)
.unwrap();
assert_eq!(updated.len(), 1); }
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn add_word_timestamps_errors_on_empty_segments() {
let t = tiny_tokenizer();
let view = AlignmentView::new(&[], 0, 3);
let err = add_word_timestamps(
&[],
&view,
&t,
"en",
WordGrouping::FineGrained,
AlignmentGather::Complete,
MAX_TOKEN_CONTEXT,
0,
crate::audio::whisper::constants::PREPEND_PUNCTUATION,
crate::audio::whisper::constants::APPEND_PUNCTUATION,
0.0,
)
.unwrap_err();
assert!(matches!(
err,
SegmentError::InvalidAlignmentShape(ref shape) if shape.rows() == 0
));
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn add_word_timestamps_errors_on_zero_columns() {
let t = tiny_tokenizer();
let hello = t.encode(" Hello").unwrap()[0];
let segments = [plain_segment(vec![hello], 0.0, 1.0)];
let view = AlignmentView::new(&[], 5, 0);
let err = add_word_timestamps(
&segments,
&view,
&t,
"en",
WordGrouping::FineGrained,
AlignmentGather::Complete,
MAX_TOKEN_CONTEXT,
0,
"",
"",
0.0,
)
.unwrap_err();
assert!(matches!(
err,
SegmentError::InvalidAlignmentShape(ref shape) if shape.cols() == 0
));
}
const PROBE_WIDTHS: [usize; 13] = [1, 3, 8, 9, 37, 63, 64, 100, 511, 1496, 1500, 1504, 4096];
#[test]
fn coreml_f16_row_pitch_answers_with_a_usable_pitch_or_a_typed_refusal() {
for cols in PROBE_WIDTHS {
for rows in [1usize, 3, 120, 224, 225] {
match coreml_f16_row_pitch(rows, cols) {
Ok(pitch) => assert!(
pitch >= cols,
"rows={rows} cols={cols}: a row pitch below the logical width ({pitch}) would make \
the gather reproduction nonsense, and must have been refused instead"
),
Err(SegmentError::AlignmentPitchUnavailable(ref pitch)) => {
assert_eq!(
(pitch.rows(), pitch.cols()),
(rows, cols),
"the refusal must name the shape it refused"
);
}
Err(SegmentError::AlignmentPitchUnexpectedLayout(ref layout)) => {
assert_eq!(
(layout.rows(), layout.cols()),
(rows, cols),
"the refusal must name the shape it refused"
);
}
Err(other) => panic!("rows={rows} cols={cols}: unexpected error {other}"),
}
}
}
}
#[test]
#[ignore = "host allocator diagnostic: asserts this machine's CoreVideo answers one shape with \
one layout, which nothing in production requires (SwiftParity fails closed on a host \
that does otherwise, and Complete never asks)"]
fn coreml_f16_row_pitch_reports_this_hosts_live_allocation_layout() {
for cols in PROBE_WIDTHS {
for rows in [1usize, 3, 120, 224, 225] {
let pitch = coreml_f16_row_pitch(rows, cols).unwrap();
let independent = MultiArray::f16_surface(&[rows, cols]).unwrap();
assert_eq!(
independent.strides(),
[pitch, 1],
"rows={rows} cols={cols}: the helper must return the surface's own strides"
);
assert_eq!(
pitch,
coreml_f16_row_pitch(rows, cols).unwrap(),
"rows={rows} cols={cols}: the pitch moved between two allocations of the same shape, so \
assumption A1 of `AlignmentGather::SwiftParity` does not hold here"
);
}
}
}
fn replay_swift_gather(
source_rows: &[f32],
src_rows: usize,
needed: usize,
cols: usize,
src_pitch: usize,
dst_pitch: usize,
) -> Vec<f32> {
let mut source_storage = vec![0.0f32; src_rows * src_pitch];
for row in 0..src_rows {
source_storage[row * src_pitch..row * src_pitch + cols]
.copy_from_slice(&source_rows[row * cols..(row + 1) * cols]);
}
let mut destination_storage = vec![f32::NAN; needed * dst_pitch];
for (offset, cell) in destination_storage
.iter_mut()
.enumerate()
.take(needed * cols)
{
*cell = source_storage.get(offset).copied().unwrap_or(0.0);
}
(0..needed)
.flat_map(|row| {
destination_storage[row * dst_pitch..row * dst_pitch + cols]
.iter()
.copied()
.map(|value| if value.is_nan() { 0.0 } else { value })
.collect::<Vec<_>>()
})
.collect()
}
#[test]
fn swift_gather_keeps_only_the_final_rows_prefix() {
fn kept_per_row(rows: usize, cols: usize, pitch: usize) -> Vec<usize> {
let ones = vec![1.0f32; rows * cols];
let gathered = replay_swift_gather(&ones, rows, rows, cols, pitch, pitch);
(0..rows)
.map(|row| {
gathered[row * cols..(row + 1) * cols]
.iter()
.take_while(|&&value| value == 1.0)
.count()
})
.collect()
}
let kept = kept_per_row(120, 1500, 1504);
assert_eq!(kept[119], 1024, "the final row keeps 1024 of 1500 columns");
assert_eq!(1500 - kept[119], 476, "and reads 476 zeros after them");
assert!(
kept[..119].iter().all(|&columns| columns == 1500),
"no row but the last is touched"
);
assert_eq!(*kept_per_row(31, 1500, 1504).last().unwrap(), 1380);
assert_eq!(*kept_per_row(2, 1500, 1504).last().unwrap(), 1496);
assert_eq!(
kept_per_row(1, 1500, 1504),
vec![1500],
"a lone row is never truncated: it starts at storage 0"
);
assert_eq!(kept_per_row(120, 1500, 1500), vec![1500; 120]);
for (rows, cols, pitch) in [
(120, 1500, 1504),
(31, 1500, 1504),
(2, 1500, 1504),
(1, 1500, 1504),
(3, 100, 128),
(7, 40, 64),
(120, 1500, 1536),
(120, 1500, 2048),
(7, 40, 40),
] {
let mut expected_data = vec![0.0f32; rows * cols];
for (row, &kept) in kept_per_row(rows, cols, pitch).iter().enumerate() {
expected_data[row * cols..row * cols + kept].fill(1.0);
}
let source = vec![1.0f32; rows * cols];
let mut data = vec![0.0f32; rows * cols];
gather_swift_rows(&mut data, &source, rows, rows, cols, pitch, pitch);
assert_eq!(data, expected_data, "rows={rows} cols={cols} pitch={pitch}");
}
}
#[test]
fn swift_gather_reproduces_the_copy_when_the_two_surfaces_pitch_differently() {
fn ramp(rows: usize, cols: usize) -> Vec<f32> {
(0..rows * cols).map(|index| index as f32 + 1.0).collect()
}
for (src_rows, needed, cols, src_pitch, dst_pitch) in [
(225usize, 120usize, 1500usize, 1504usize, 1504usize),
(225, 120, 1500, 1504, 1536),
(225, 120, 1500, 1536, 1504),
(225, 120, 1500, 1500, 1504),
(225, 120, 1500, 1504, 1500),
(225, 120, 1500, 1500, 1500),
(8, 5, 4, 8, 4),
(8, 5, 4, 4, 8),
(8, 5, 4, 8, 8),
(8, 8, 4, 6, 5),
(225, 1, 1500, 1504, 1504),
(225, 1, 1500, 1500, 1504),
(4, 4, 1, 1, 1),
(4, 4, 1, 3, 2),
(1, 1, 1, 1, 1),
(1, 1, 4, 7, 5),
(3, 6, 4, 5, 5),
(3, 6, 4, 4, 7),
(0, 3, 4, 4, 6),
] {
let source = ramp(src_rows, cols);
let expected = replay_swift_gather(&source, src_rows, needed, cols, src_pitch, dst_pitch);
let mut data = vec![0.0f32; needed * cols];
gather_swift_rows(
&mut data, &source, src_rows, needed, cols, src_pitch, dst_pitch,
);
assert_eq!(
data, expected,
"src_rows={src_rows} needed={needed} cols={cols} src_pitch={src_pitch} \
dst_pitch={dst_pitch}"
);
}
}
#[test]
fn swift_gather_at_equal_unpadded_pitches_is_the_plain_prefix_take() {
for (src_rows, needed, cols) in [
(225usize, 120usize, 1500usize),
(8, 5, 4),
(1, 1, 1),
(3, 6, 4),
] {
let source: Vec<f32> = (0..src_rows * cols).map(|i| i as f32 + 1.0).collect();
let mut prefix = vec![0.0f32; needed * cols];
let copied = src_rows.min(needed) * cols;
prefix[..copied].copy_from_slice(&source[..copied]);
let mut data = vec![0.0f32; needed * cols];
gather_swift_rows(&mut data, &source, src_rows, needed, cols, cols, cols);
assert_eq!(
data, prefix,
"src_rows={src_rows} needed={needed} cols={cols}: an unpadded host must gather every row \
whole"
);
}
}
#[test]
fn swift_gather_reads_the_sources_padding_as_zero() {
let source: Vec<f32> = (0..3 * 4).map(|i| i as f32 + 1.0).collect();
let mut data = vec![0.0f32; 3 * 4];
gather_swift_rows(&mut data, &source, 3, 3, 4, 6, 4);
assert_eq!(
data,
vec![
1.0, 2.0, 3.0, 4.0, 0.0, 0.0, 5.0, 6.0, 7.0, 8.0, 0.0, 0.0, ],
"the source's padding must read as zero, not as a neighbouring row's weights"
);
}
#[test]
fn swift_parity_probes_swifts_source_height_not_this_ports_commit_headroom() {
const COLS: usize = 8;
const SWIFT_ROWS: usize = MAX_TOKEN_CONTEXT; const VIEW_ROWS: usize = SWIFT_ROWS + 1; const NEEDED: usize = 5;
let asked = std::cell::RefCell::new(Vec::new());
let injected = |rows: usize, cols: usize| -> Result<usize, SegmentError> {
asked.borrow_mut().push((rows, cols));
Ok(if rows == VIEW_ROWS { cols * 2 } else { cols })
};
let source: Vec<f32> = (0..VIEW_ROWS * COLS).map(|i| i as f32 + 1.0).collect();
let view = AlignmentView::new(&source, VIEW_ROWS, COLS);
let mut data = vec![0.0f32; NEEDED * COLS];
gather_swift_parity_into(&mut data, &view, NEEDED, COLS, SWIFT_ROWS, &injected).unwrap();
let asked_shapes: Vec<(usize, usize)> = asked.borrow().clone();
assert!(
asked_shapes.contains(&(SWIFT_ROWS, COLS)),
"the source probe must ask CoreVideo about Swift's own {SWIFT_ROWS}-row array; asked \
{asked_shapes:?}"
);
assert!(
!asked_shapes.contains(&(VIEW_ROWS, COLS)),
"the source probe must NOT ask about this port's {VIEW_ROWS}-row commit accumulator; asked \
{asked_shapes:?}"
);
assert!(
asked_shapes.contains(&(NEEDED, COLS)),
"and the destination probe must ask about the per-call `[needed, cols]` surface; asked \
{asked_shapes:?}"
);
let at_swift_height = replay_swift_gather(&source, SWIFT_ROWS, NEEDED, COLS, COLS, COLS);
assert_eq!(
data, at_swift_height,
"the reproduction must decode Swift's storage at Swift's own source pitch"
);
let at_port_height = replay_swift_gather(&source, SWIFT_ROWS, NEEDED, COLS, COLS * 2, COLS);
assert_ne!(
at_swift_height, at_port_height,
"the fixture proves nothing unless the two heights disagree"
);
}
const RECORDED_REFERENCE_HOST_PITCH: [(usize, usize); 6] = [
(8, 32),
(9, 32),
(100, 128),
(1496, 1504),
(1500, 1504),
(1504, 1504),
];
#[test]
#[ignore = "reference-host layout probe: asserts the #41 capture host's CoreVideo pitches, which \
no production path depends on (run explicitly when re-capturing the Swift probe)"]
fn reference_host_pitch_table() {
for (cols, recorded) in RECORDED_REFERENCE_HOST_PITCH {
assert_eq!(
coreml_f16_row_pitch(224, cols).unwrap(),
recorded,
"cols={cols}: this host's CoreVideo Float16 row pitch differs from the reference host \
the whisper #41 probe and long-form parity numbers were captured on. The shipping \
gather is UNAFFECTED -- it measures this host rather than assuming a quantum -- but \
the hand-computed columns in the gather fixtures, and the recorded 1417 s/1042 s \
parity results, describe the reference layout only"
);
}
for (cols, recorded) in RECORDED_REFERENCE_HOST_PITCH {
for rows in [1usize, 3, 120, 225] {
assert_eq!(
coreml_f16_row_pitch(rows, cols).unwrap(),
recorded,
"rows={rows} cols={cols}: this host's pitch varies with the row count, unlike the \
reference host's"
);
}
}
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn swift_parity_gather_truncates_final_alignment_row() {
let t = tiny_tokenizer();
let s = SpecialTokens::whisper_defaults();
let hello = t.encode(" Hello").unwrap()[0];
let world = t.encode(" world").unwrap()[0];
let tokens = vec![hello, world, s.time_token_begin() + 100];
let log_probs: Vec<(u32, f32)> = tokens.iter().map(|&token| (token, -0.2)).collect();
let mut segment = TranscriptionSegment::new();
segment
.set_tokens(tokens)
.set_token_log_probs(log_probs)
.set_start(0.0)
.set_end(3.0);
let cols = 100usize;
let rows = 3usize;
let pitch = coreml_f16_row_pitch(rows, cols).unwrap();
let truncated_at = (rows * cols).saturating_sub((rows - 1) * pitch).min(cols);
assert_eq!(
truncated_at, 44,
"at the reference host's pitch of 128 the final row keeps 44 of {cols} columns; this host \
measured a pitch of {pitch}, so the plateau/cutoff columns below no longer straddle the \
truncation and this fixture would not discriminate (see \
`reference_host_pitch_table`)"
);
let row_one_cutoff = truncated_at + 36;
let mut weights = vec![0.0f32; rows * cols];
weights[0] = 3.0; for column in 0..cols {
weights[cols + column] = if column < row_one_cutoff { 0.5 } else { -0.1 };
weights[2 * cols + column] = if column < truncated_at { 0.0 } else { 1.0 };
}
let mut pre_truncated = weights.clone();
pre_truncated[2 * cols + truncated_at..rows * cols].fill(0.0);
let last_word_end = |weights: &[f32], gather| {
let view = AlignmentView::new(weights, 3, cols);
let updated = add_word_timestamps(
std::slice::from_ref(&segment),
&view,
&t,
"en",
WordGrouping::FineGrained,
gather,
MAX_TOKEN_CONTEXT,
0,
PREPEND_PUNCTUATION,
APPEND_PUNCTUATION,
0.0,
)
.unwrap();
let words = updated[0].words_slice();
assert_eq!(
words.iter().map(WordTiming::word).collect::<Vec<_>>(),
vec![" Hello", " world"],
"the trailing timestamp token is dropped, so ` world` is the last word"
);
(words.last().unwrap().end(), updated[0].end())
};
let (complete_end, complete_segment_end) = last_word_end(&weights, AlignmentGather::Complete);
let (parity_end, parity_segment_end) = last_word_end(&weights, AlignmentGather::SwiftParity);
let (reference_end, _) = last_word_end(&pre_truncated, AlignmentGather::Complete);
assert_ne!(
complete_end, parity_end,
"the gather modes must disagree, or this fixture proves nothing"
);
assert_eq!(
complete_end, 0.88,
"row 2's tail pulls the boundary to col {truncated_at}"
);
assert_eq!(
parity_end,
1.58,
"with that tail zeroed, row 1 keeps the path to col {}",
row_one_cutoff - 1
);
assert_eq!(
parity_end, reference_end,
"SwiftParity over the full matrix == Complete over the hand-truncated one"
);
assert_ne!(
complete_end, reference_end,
"the hand truncation must actually move the boundary"
);
assert_eq!(complete_segment_end, complete_end);
assert_eq!(parity_segment_end, parity_end);
}
#[test]
fn window_span_never_collapses_a_short_final_window_onto_its_start() {
let (start, end) = window_span(16_800_000, 1);
assert_eq!(start, 1050.0, "1050 s is exactly representable in f32");
assert!(
end > start,
"a window that held audio cannot end where it began: {start} .. {end}"
);
assert_eq!(end, start.next_up(), "one ulp above the start, not on it");
}
#[test]
fn window_span_nudges_an_end_the_addition_itself_absorbs() {
let seek: usize = 32_768_000;
let start = seek as f32 / SAMPLE_RATE as f32;
assert_eq!(
start + 1.0 / SAMPLE_RATE as f32,
start,
"the unguarded addition absorbs a one-sample duration at 2048 s"
);
let (guarded_start, end) = window_span(seek, 1);
assert_eq!(guarded_start, 2048.0);
assert_eq!(
end,
start.next_up(),
"absorbed, so nudged rather than empty"
);
}
#[test]
fn window_span_of_a_large_ordinary_window_is_its_own_duration() {
let seek = (1usize << 24) + 1;
let (start, end) = window_span(seek, 480_000);
let ulp = start.next_up() - start;
assert!(
(end - start - 30.0).abs() <= ulp,
"a 30 s window spans 30 s within one ulp ({ulp}): {start} .. {end}"
);
}
#[test]
fn window_span_of_an_empty_window_is_empty() {
let (start, end) = window_span(16_800_000, 0);
assert_eq!(start, end, "no audio, no extent");
}
#[test]
fn shift_span_never_collapses_a_span_the_shift_absorbs() {
let (start, end) = window_span(0, 1);
assert!(end > start, "the local span is valid before the shift");
let offset_seconds = 32_768_000f32 / SAMPLE_RATE as f32;
let (shifted_start, shifted_end) = shift_span(start, end, offset_seconds);
assert_eq!(
shifted_start, 2048.0,
"2048 s is exactly representable in f32"
);
assert!(
shifted_end > shifted_start,
"a span that had an extent cannot lose it to re-anchoring: \
{shifted_start} .. {shifted_end}"
);
assert_eq!(
shifted_end,
shifted_start.next_up(),
"one ulp above the shifted start, not on it"
);
}
#[test]
fn shift_span_of_an_ordinary_shift_is_exact() {
assert_eq!(shift_span(1.0, 2.0, 2.0), (3.0, 4.0));
assert_eq!(shift_span(0.5, 0.75, 0.25), (0.75, 1.0));
assert_eq!(
shift_span(1.0, 2.0, 0.0),
(1.0, 2.0),
"a zero offset moves nothing"
);
}
#[test]
fn shift_span_leaves_an_originally_empty_span_empty() {
assert_eq!(shift_span(1.0, 1.0, 2.0), (3.0, 3.0), "an ordinary shift");
let offset_seconds = 32_768_000f32 / SAMPLE_RATE as f32;
let (start, end) = shift_span(1.0, 1.0, offset_seconds);
assert_eq!(
start, end,
"no extent going in, no spurious extent coming out: {start} .. {end}"
);
assert_eq!(start, 2049.0, "1 s into a chunk that begins at 2048 s");
}
#[test]
#[ignore = "requires local tokenizer (WHISPERKIT_TEST_MODELS)"]
fn lump_segment_carries_the_shared_window_span() {
let t = tiny_tokenizer();
let seek = 16_800_000usize;
let r = result_with_tokens(vec![50258, 100], 0.0, -0.2);
let (next_seek, segments) =
find_seek_point_and_segments(&r, &DecodingOptions::new(), 0, seek, 1, &t).unwrap();
let segments = segments.expect("a confident window is not silence-skipped");
assert_eq!(
segments.len(),
1,
"no timestamps at all -> one lump segment"
);
assert_eq!(next_seek, seek + 1, "the lump branch consumes the window");
assert_eq!(
(segments[0].start(), segments[0].end()),
window_span(seek, 1),
"the segment is timed through the observation's own helper"
);
assert!(
segments[0].end() > segments[0].start(),
"the window held a sample, so its segment has an extent"
);
}