mod common;
use core::sync::atomic::AtomicBool;
use std::collections::BTreeMap;
use core::num::NonZeroU32;
use coremlit::audio::align::{
ANALYSIS_TIMEBASE, AcousticContract, AcousticGeometry, AlignError, Aligner, AlignerError,
AlignerOptions, EnglishNormalizer, Granularity, Lang, LetterCase, OovEvent, OovKind, OutputClock,
OutputKind, Tokenization, Vocabulary, Word, WordDelimiter, default_oov_policy,
};
fn contract(blank: u32, geometry: AcousticGeometry) -> AcousticContract {
AcousticContract::new(
blank,
geometry,
Tokenization::new(
WordDelimiter::Pipe,
LetterCase::Upper,
Granularity::Character,
&[],
),
OutputKind::LogProbabilities,
)
}
fn align_jfk(samples: &[f32]) -> Vec<Word> {
let aligner = Aligner::from_paths(
Lang::En,
&common::model_path(),
Box::new(EnglishNormalizer::new()),
)
.expect(
"build the En aligner from the CoreML model + bundled tokenizer (set ALIGNKIT_TEST_MODELS \
to the model directory)",
);
let text = common::JFK_TRANSCRIPT;
let resolution = aligner
.detect_oov(text)
.expect("detect_oov")
.decide(default_oov_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
aligner
.align_chunk(samples, &[], text, clock, &abort, resolution)
.expect("align_chunk succeeds end-to-end")
.words()
.to_vec()
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn align_chunk_produces_monotonic_word_timings() {
let samples = common::load_wav_mono_f32(&common::jfk_wav_path());
assert!(!samples.is_empty(), "fixture decoded to no samples");
let words = &align_jfk(&samples);
assert!(
!words.is_empty(),
"a real transcript over matching audio must produce words"
);
let bound = samples.len() as i64;
let mut prev_start = 0_i64;
for word in words {
let range = word.range();
let (start, end) = (range.start_pts(), range.end_pts());
assert!(
start <= end,
"word `{}`: start {start} exceeds end {end}",
word.text()
);
assert!(
start >= 0 && end <= bound,
"word `{}`: range [{start}, {end}] escapes the audio [0, {bound}]",
word.text()
);
assert!(
start >= prev_start,
"word `{}`: start {start} precedes the previous word's start {prev_start} (not monotonic)",
word.text()
);
let score = word.score();
assert!(
(0.0..=1.0).contains(&score),
"word `{}`: score {score} outside [0, 1]",
word.text()
);
prev_start = start;
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn align_chunk_is_bit_identical_across_runs() {
let samples = common::load_wav_mono_f32(&common::jfk_wav_path());
let first = align_jfk(&samples);
let second = align_jfk(&samples);
assert!(
!first.is_empty(),
"a real transcript over matching audio must produce words"
);
assert_eq!(
first.len(),
second.len(),
"two runs over identical input produced different word counts"
);
for (a, b) in first.iter().zip(&second) {
assert_eq!(a.text(), b.text(), "word text differs between runs");
assert_eq!(
(a.range().start_pts(), a.range().end_pts()),
(b.range().start_pts(), b.range().end_pts()),
"word `{}`: timing differs between two runs over identical input",
a.text()
);
assert_eq!(
a.score().to_bits(),
b.score().to_bits(),
"word `{}`: score differs between two runs over identical input ({} vs {})",
a.text(),
a.score(),
b.score()
);
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn align_chunk_641_abc_is_a_named_no_alignment_path() {
let aligner = Aligner::from_paths(
Lang::En,
&common::model_path(),
Box::new(EnglishNormalizer::new()),
)
.expect("build the En aligner (set ALIGNKIT_TEST_MODELS to the model directory)");
let jfk = common::load_wav_mono_f32(&common::jfk_wav_path());
let samples = &jfk[80_000..80_641];
assert_eq!(
samples.len(),
641,
"the fence case is exactly 641 real samples"
);
let text = "ABC";
let detection = aligner.detect_oov(text).expect("detect_oov");
assert!(
detection.events().is_empty(),
"A, B, C must be in-vocab, or the OOV path — not the frame count — would drive the result"
);
let resolution = detection.decide(default_oov_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
match aligner.align_chunk(samples, &[], text, clock, &abort, resolution) {
Err(AlignError::NoAlignmentPath(_)) => {}
Ok(result) => panic!(
"one frame cannot carry three distinct tokens, yet align_chunk returned words {:?}",
result.words().iter().map(Word::text).collect::<Vec<_>>()
),
Err(err) => panic!("expected the named AlignError::NoAlignmentPath, got {err:?}"),
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_fail_closed_refusal_is_named_through_align_chunk() {
let aligner = Aligner::from_paths(
Lang::En,
&common::model_path(),
Box::new(EnglishNormalizer::new()),
)
.expect("build the En aligner (set ALIGNKIT_TEST_MODELS to the model directory)");
let samples = common::load_wav_mono_f32(&common::jfk_wav_path());
let text = "ask not what your country can do for you, AT&T";
let resolution = aligner
.detect_oov(text)
.expect("detect_oov")
.decide(default_oov_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
let err = aligner
.align_chunk(&samples, &[], text, clock, &abort, resolution)
.expect_err("the default policy fails closed on `&`");
let AlignError::Refused(refusal) = err else {
panic!("the refusal must be named, got {err:?}");
};
assert_eq!(
refusal
.events()
.iter()
.map(|event| event.kind().clone())
.collect::<Vec<_>>(),
[OovKind::Symbol('&')],
"the refusal names the one position the policy refused"
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_character_the_vocabulary_cannot_spell_is_an_event_through_the_aligner() {
let aligner = Aligner::from_paths(
Lang::En,
&common::model_path(),
Box::new(EnglishNormalizer::new()),
)
.expect("build the En aligner (set ALIGNKIT_TEST_MODELS to the model directory)");
let samples = common::load_wav_mono_f32(&common::jfk_wav_path());
let text = common::JFK_TRANSCRIPT.replacen("Americans", "Américans", 1);
let detection = aligner
.detect_oov(&text)
.expect("a character the vocabulary cannot spell is an event, never an error");
let events = detection.events();
let symbols: Vec<&OovEvent> = events
.iter()
.filter(|event| matches!(event.kind(), OovKind::Symbol(_)))
.collect();
assert_eq!(symbols.len(), 1, "one unspellable character: {events:?}");
assert_eq!(symbols[0].kind(), &OovKind::Symbol('é'));
assert_eq!(symbols[0].word_index(), 4, "`Américans` is the fifth word");
let resolution = detection.decide(default_oov_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
let words = aligner
.align_chunk(&samples, &[], &text, clock, &abort, resolution)
.expect("the wildcarded character is aligned around")
.words()
.to_vec();
assert_eq!(
words.len(),
align_jfk(&samples).len(),
"every word of the transcript aligns, the one holding `é` included"
);
assert!(
words
.iter()
.any(|word| word.text().eq_ignore_ascii_case("américans")),
"{:?}",
words.iter().map(Word::text).collect::<Vec<_>>()
);
}
fn align_jfk_with(
samples: &[f32],
vocabulary: &Vocabulary,
contract: &AcousticContract,
) -> Vec<Word> {
let aligner = Aligner::from_paths_with_vocabulary(
Lang::En,
&common::model_path(),
vocabulary,
contract,
Box::new(EnglishNormalizer::new()),
AlignerOptions::new(),
)
.expect("the model loads with its own vocabulary and contract");
let text = common::JFK_TRANSCRIPT;
let resolution = aligner
.detect_oov(text)
.expect("detect_oov")
.decide(default_oov_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
aligner
.align_chunk(samples, &[], text, clock, &abort, resolution)
.expect("align_chunk through the model's own vocabulary")
.words()
.to_vec()
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn the_staged_model_aligns_identically_through_its_own_vocabulary() {
let samples = common::load_wav_mono_f32(&common::jfk_wav_path());
let bundled = align_jfk(&samples);
let vocabulary = Vocabulary::from_file(common::dict_path())
.expect("read base960h_dict.json (set ALIGNKIT_TEST_MODELS to the model directory)");
assert_eq!(vocabulary.size().get(), 29);
let generic = contract(0, AcousticGeometry::WAV2VEC2);
for contract in [AcousticContract::BASE960H, generic] {
let own = align_jfk_with(&samples, &vocabulary, &contract);
assert!(!own.is_empty(), "jfk.wav aligns to words");
assert_eq!(own.len(), bundled.len(), "the same number of words");
for (a, b) in own.iter().zip(&bundled) {
assert_eq!(a.text(), b.text());
assert_eq!(
(a.range().start_pts(), a.range().end_pts()),
(b.range().start_pts(), b.range().end_pts()),
"word `{}` under {contract:?}: the two roads time it differently",
a.text()
);
assert_eq!(
a.score().to_bits(),
b.score().to_bits(),
"word `{}` under {contract:?}: the two roads score it differently",
a.text()
);
}
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_vocabulary_of_another_width_is_refused_by_name_at_load() {
let staged = std::fs::read(common::dict_path())
.expect("read base960h_dict.json (set ALIGNKIT_TEST_MODELS to the model directory)");
let table: BTreeMap<String, u32> =
serde_json::from_slice(&staged).expect("the staged table is `{token: id}` JSON");
let mut wider = table.clone();
wider.insert("É".to_owned(), 29);
let mut narrower = table;
assert_eq!(narrower.remove("Z"), Some(28), "`Z` is the last id");
for (synthetic, entries) in [(wider, 30), (narrower, 28)] {
let json = serde_json::to_vec(&synthetic).expect("a table serializes");
let vocabulary = Vocabulary::from_json(&json).expect("the synthetic table is well formed");
assert_eq!(vocabulary.size().get(), entries);
let result = Aligner::from_paths_with_vocabulary(
Lang::En,
&common::model_path(),
&vocabulary,
&AcousticContract::BASE960H,
Box::new(EnglishNormalizer::new()),
AlignerOptions::new(),
);
let Err(AlignerError::VocabularyMismatch(mismatch)) = result else {
panic!(
"a {entries}-entry table on the 29-class model must be refused by name, got {:?}",
result.err()
);
};
assert_eq!((mismatch.vocabulary(), mismatch.model()), (entries, 29));
}
}
fn load_with(vocabulary: &Vocabulary, contract: &AcousticContract) -> AlignerError {
match Aligner::from_paths_with_vocabulary(
Lang::En,
&common::model_path(),
vocabulary,
contract,
Box::new(EnglishNormalizer::new()),
AlignerOptions::new(),
) {
Ok(_) => panic!("{contract:?} must be refused at load"),
Err(err) => err,
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_geometry_that_disagrees_with_the_declared_frame_count_is_refused_at_load() {
let vocabulary = Vocabulary::from_file(common::dict_path())
.expect("read base960h_dict.json (set ALIGNKIT_TEST_MODELS to the model directory)");
for (stride, derived) in [(160u32, 5_998usize), (321, 2_990)] {
let geometry = AcousticGeometry::new(
16_000,
NonZeroU32::new(400).expect("nonzero"),
NonZeroU32::new(stride).expect("nonzero"),
)
.expect("a geometry");
let AlignerError::FrameCountMismatch(mismatch) = load_with(&vocabulary, &contract(0, geometry))
else {
panic!("a {stride}-sample stride must be a FrameCountMismatch");
};
assert_eq!(
(mismatch.window(), mismatch.declared(), mismatch.derived()),
(960_000, 2_999, derived)
);
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn an_explicit_blank_outside_the_table_is_refused_at_load() {
let vocabulary = Vocabulary::from_file(common::dict_path())
.expect("read base960h_dict.json (set ALIGNKIT_TEST_MODELS to the model directory)");
let AlignerError::BlankOutOfVocabulary(refused) =
load_with(&vocabulary, &contract(29, AcousticGeometry::WAV2VEC2))
else {
panic!("id 29 is no id of the 29-entry table");
};
assert_eq!((refused.blank(), refused.entries()), (29, 29));
}