use super::*;
use asry::{
emissions::{EmissionsAligner, EmissionsError, EnglishNormalizer},
time::ANALYSIS_TIMEBASE,
};
use crate::audio::align::error::{Refusal, RefusedOov};
fn default_policy(event: &SetOovEvent<'_>) -> OovDecision {
event.default_decision()
}
fn fail_closed_policy(_event: &SetOovEvent<'_>) -> OovDecision {
OovDecision::FailClosed
}
#[test]
fn aligner_key_distinguishes_lang_from_any() {
assert_ne!(AlignerKey::Lang(Lang::En), AlignerKey::Any);
assert_eq!(AlignerKey::Lang(Lang::En), AlignerKey::Lang(Lang::En));
assert_ne!(AlignerKey::Lang(Lang::En), AlignerKey::Lang(Lang::Zh));
}
#[test]
fn aligner_key_hashes_consistently() {
use std::collections::HashSet;
let mut set = HashSet::new();
set.insert(AlignerKey::Lang(Lang::En));
set.insert(AlignerKey::Any);
set.insert(AlignerKey::Lang(Lang::En)); assert_eq!(set.len(), 2);
}
#[test]
fn fallback_default_is_skip_chunk() {
assert_eq!(AlignmentFallback::default(), AlignmentFallback::SkipChunk);
}
#[test]
fn fallback_round_trips_through_its_own_text_form() {
for &fallback in AlignmentFallback::ALL {
let text = fallback.as_str();
assert_eq!(
text.parse::<AlignmentFallback>(),
Ok(fallback),
"`{text}` must parse back to the variant that produced it"
);
assert_eq!(fallback.to_string(), text);
}
}
#[test]
fn fallback_from_str_names_the_snake_case_spelling() {
for &fallback in AlignmentFallback::ALL {
let spelling = match fallback {
AlignmentFallback::SkipChunk => "skip_chunk",
AlignmentFallback::Error => "error",
};
assert_eq!(
fallback.as_str(),
spelling,
"as_str spelling for {fallback:?} drifted from its pinned wire form"
);
assert_eq!(
spelling.parse::<AlignmentFallback>(),
Ok(fallback),
"`{spelling}` must parse back to {fallback:?}"
);
}
}
#[test]
fn fallback_from_str_is_total_and_rejects_everything_else() {
for unknown in [
"",
"SkipChunk",
"skip-chunk",
"Error",
"skip_chunk ",
"fail",
] {
assert!(
unknown.parse::<AlignmentFallback>().is_err(),
"`{unknown}` must not parse"
);
}
}
#[test]
fn fallback_is_variant_predicates() {
assert!(AlignmentFallback::SkipChunk.is_skip_chunk());
assert!(!AlignmentFallback::SkipChunk.is_error());
assert!(AlignmentFallback::Error.is_error());
}
#[test]
fn aligner_key_is_variant_predicates() {
assert!(AlignerKey::Lang(Lang::En).is_lang());
assert!(!AlignerKey::Lang(Lang::En).is_any());
assert!(AlignerKey::Any.is_any());
}
#[cfg(feature = "serde")]
#[test]
fn fallback_serde_uses_the_same_snake_case_spelling() {
for &fallback in AlignmentFallback::ALL {
let expected_json = match fallback {
AlignmentFallback::SkipChunk => r#""skip_chunk""#,
AlignmentFallback::Error => r#""error""#,
};
let json = serde_json::to_string(&fallback).unwrap();
assert_eq!(
json, expected_json,
"serde spelling for {fallback:?} drifted"
);
assert_eq!(json, format!("\"{}\"", fallback.as_str()));
let back: AlignmentFallback = serde_json::from_str(&json).unwrap();
assert_eq!(
back, fallback,
"{fallback:?} must deserialize back from its own JSON"
);
}
assert!(serde_json::from_str::<AlignmentFallback>(r#""SkipChunk""#).is_err());
}
#[test]
fn empty_set_misses_with_default_fallback() {
let set = AlignmentSetBuilder::new().build();
assert!(set.is_empty());
assert_eq!(set.len(), 0);
match set.lookup(&Lang::En) {
AlignmentLookup::Miss(fallback) => assert_eq!(fallback, AlignmentFallback::SkipChunk),
_ => panic!("expected Miss"),
}
}
#[test]
fn empty_set_misses_with_error_fallback() {
let set = AlignmentSetBuilder::new()
.with_fallback(AlignmentFallback::Error)
.build();
assert_eq!(set.fallback(), AlignmentFallback::Error);
match set.lookup(&Lang::Zh) {
AlignmentLookup::Miss(fallback) => assert_eq!(fallback, AlignmentFallback::Error),
_ => panic!("expected Miss"),
}
}
#[test]
fn builder_set_fallback_in_place() {
let mut builder = AlignmentSetBuilder::new();
assert!(builder.is_empty());
builder.set_fallback(AlignmentFallback::Error);
assert_eq!(builder.build().fallback(), AlignmentFallback::Error);
}
#[test]
fn empty_set_detect_oov_on_miss_reads_nothing() {
let set = AlignmentSetBuilder::new().build();
let detection = set.detect_oov("anything", &Lang::En).unwrap();
assert_eq!(detection.language(), &Lang::En);
assert!(detection.events().is_none());
let resolution = detection.decide(fail_closed_policy);
assert_eq!(resolution.language(), &Lang::En);
assert!(resolution.resolved().is_none());
}
#[test]
fn a_resolution_decided_for_another_language_is_refused_on_a_miss_too() {
for fallback in [AlignmentFallback::SkipChunk, AlignmentFallback::Error] {
let set = AlignmentSetBuilder::new().with_fallback(fallback).build();
let resolution = set
.detect_oov("anything", &Lang::Zh)
.expect("detect_oov")
.decide(default_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let err = set
.align_chunk(&Lang::En, &[], &[], "anything", clock, &abort, resolution)
.expect_err("decisions made for Zh do not answer an En request");
assert!(
matches!(
err,
AlignError::DecisionLanguage(ref e)
if *e.requested() == Lang::En && *e.found() == Lang::Zh
),
"{fallback}: {err:?}"
);
}
}
fn missed(set: &AlignmentSet, language: &Lang) -> SetResolution {
set
.detect_oov("anything", language)
.expect("a miss detects nothing, and fails at nothing")
.decide(default_policy)
}
#[test]
fn align_chunk_miss_skip_chunk_returns_empty_words() {
let set = AlignmentSetBuilder::new().build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let alignment = set
.align_chunk(
&Lang::Zh,
&[],
&[],
"anything",
clock,
&abort,
missed(&set, &Lang::Zh),
)
.expect("a SkipChunk miss is not an error");
assert!(alignment.words().is_empty());
assert!(matches!(alignment.cause(), Some(UnalignedCause::Skipped)));
}
#[test]
fn align_chunk_miss_error_returns_language_unsupported() {
let set = AlignmentSetBuilder::new()
.with_fallback(AlignmentFallback::Error)
.build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let err = set
.align_chunk(
&Lang::Zh,
&[],
&[],
"anything",
clock,
&abort,
missed(&set, &Lang::Zh),
)
.unwrap_err();
assert!(matches!(
err,
AlignError::LanguageUnsupported(ref language) if *language == Lang::Zh
));
}
#[test]
fn resolve_binds_the_requested_language_and_reports_the_miss() {
let set = AlignmentSetBuilder::new().build();
let handle = set.resolve(&Lang::Zh);
assert_eq!(handle.language(), &Lang::Zh);
assert_eq!(
handle.binding(),
AlignmentBinding::Miss(AlignmentFallback::SkipChunk)
);
}
#[test]
fn handle_align_chunk_is_the_guarded_set_align_chunk() {
let set = AlignmentSetBuilder::new().build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let handle = set.resolve(&Lang::Zh);
let resolution = handle
.detect_oov("anything")
.expect("detect_oov")
.decide(default_policy);
let alignment = handle
.align_chunk(&[], &[], "anything", clock, &abort, resolution)
.expect("a SkipChunk miss is not an error");
assert!(alignment.words().is_empty());
assert!(matches!(alignment.cause(), Some(UnalignedCause::Skipped)));
}
fn english_seam() -> EmissionsAligner {
EmissionsAligner::builder(Lang::En, crate::audio::align::vocab::tokenizer_json_bytes())
.normalizer(Box::new(EnglishNormalizer::new()))
.blank_token_id(crate::audio::align::vocab::BLANK_ID)
.build()
.expect("the bundled document builds a seam")
}
fn english_wildcard_korean_fail_closed(event: &SetOovEvent<'_>) -> OovDecision {
match event.language() {
Lang::Ko => OovDecision::FailClosed,
_ => OovDecision::Wildcard,
}
}
fn korean_wildcard_english_fail_closed(event: &SetOovEvent<'_>) -> OovDecision {
match event.language() {
Lang::Ko => OovDecision::Wildcard,
_ => OovDecision::FailClosed,
}
}
fn decided_under(
requested: Lang,
text: &str,
policy: fn(&SetOovEvent<'_>) -> OovDecision,
) -> Vec<(Lang, OovDecision)> {
let detection = english_seam().detect_oov(text).expect("detect_oov");
assert!(
detection
.events()
.iter()
.all(|event| event.language() == &Lang::En),
"asry stamps the language of the aligner that read the text"
);
let held = SetDetection {
made_by: SetId::next(),
language: requested.clone(),
detection: Some(detection),
};
let shown = held.events().expect("an aligner read the text");
assert!(shown.iter().all(|event| event.language() == &requested));
let resolution = held.decide(policy);
resolution
.resolved()
.expect("decided")
.iter()
.map(|resolved| (resolved.event().language().clone(), resolved.decision()))
.collect()
}
#[test]
fn a_fallback_detection_is_judged_under_the_requested_language() {
let text = "Café AT&T b4d";
assert_eq!(
decided_under(Lang::Ko, text, english_wildcard_korean_fail_closed),
vec![(Lang::Ko, OovDecision::FailClosed); 3]
);
assert_eq!(
decided_under(Lang::Ko, text, korean_wildcard_english_fail_closed),
vec![(Lang::Ko, OovDecision::Wildcard); 3]
);
}
#[test]
fn an_exact_language_detection_is_judged_as_before() {
let text = "Café AT&T b4d";
assert_eq!(
decided_under(Lang::En, text, english_wildcard_korean_fail_closed),
vec![(Lang::En, OovDecision::Wildcard); 3]
);
assert_eq!(
decided_under(Lang::En, text, korean_wildcard_english_fail_closed),
vec![(Lang::En, OovDecision::FailClosed); 3]
);
}
#[test]
fn the_event_view_compares_and_prints_the_language_shown() {
let detection = english_seam().detect_oov("b4d").expect("detect_oov");
let held = SetDetection {
made_by: SetId::next(),
language: Lang::Ko,
detection: Some(detection),
};
let shown = held.events().expect("read");
let [four] = shown.as_slice() else {
panic!("one event: {shown:?}");
};
assert_eq!(
(four.char(), four.char_index(), four.word_index()),
(Some('4'), 1, 0)
);
assert_eq!(four.default_decision(), OovDecision::Wildcard);
let rendered = format!("{four:?}");
assert!(rendered.contains("language: Ko"), "{rendered}");
assert!(!rendered.contains("En"), "{rendered}");
assert_eq!(*four, held.events().expect("read")[0]);
}
#[test]
fn a_detection_and_its_resolution_print_the_requested_language_alone() {
let held = SetDetection {
made_by: SetId::next(),
language: Lang::Ko,
detection: Some(
english_seam()
.detect_oov("Café AT&T b4d")
.expect("detect_oov"),
),
};
let detection = format!("{held:?}");
assert!(detection.contains("language: Ko"), "{detection}");
assert!(detection.contains("Symbol('&')"), "{detection}");
assert!(!detection.contains("En"), "{detection}");
let resolution = format!("{:?}", held.decide(english_wildcard_korean_fail_closed));
assert!(resolution.contains("language: Ko"), "{resolution}");
assert!(resolution.contains("FailClosed"), "{resolution}");
assert!(!resolution.contains("En"), "{resolution}");
}
#[test]
fn the_registry_restates_a_refusal_under_the_requested_language() {
let read = english_seam()
.detect_oov("Café AT&T b4d")
.expect("detect_oov");
let refusal = Refusal::new(read.events().iter().map(RefusedOov::detected).collect());
assert!(
refusal
.events()
.iter()
.all(|event| event.language() == &Lang::En)
);
let restated = for_request(AlignError::Refused(refusal), &Lang::Ko);
let shown = format!("{restated:?}");
let AlignError::Refused(refusal) = restated else {
panic!("a refusal stays a refusal: {shown}");
};
assert_eq!(refusal.events().len(), 3);
assert!(
refusal
.events()
.iter()
.all(|event| event.language() == &Lang::Ko)
);
assert!(!shown.contains("En"), "{shown}");
let other = for_request(AlignError::LanguageUnsupported(Lang::Zh), &Lang::Ko);
assert!(matches!(other, AlignError::LanguageUnsupported(Lang::Zh)));
}
fn decided_by(set: &AlignmentSet, language: &Lang, text: &str) -> SetResolution {
set
.detect_oov(text, language)
.expect("detect_oov")
.decide(default_policy)
}
fn on_a_miss(set: &AlignmentSet, resolution: SetResolution) -> Result<UnitAlignment, AlignError> {
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
set.align_chunk(&Lang::En, &[], &[], "anything", clock, &abort, resolution)
}
#[test]
fn a_set_and_what_it_makes_share_one_identity() {
let one = AlignmentSetBuilder::new().build();
let two = AlignmentSetBuilder::new().build();
assert_ne!(one.id(), two.id());
assert_ne!(one.id().to_string(), two.id().to_string());
assert!(one.id().to_string().starts_with("alignment set #"));
let detection = one.detect_oov("anything", &Lang::En).expect("detect_oov");
assert_eq!(detection.made_by(), one.id());
assert_eq!(detection.decide(default_policy).made_by(), one.id());
}
#[test]
fn a_foreign_resolution_is_refused_on_a_miss_under_either_policy() {
let a = AlignmentSetBuilder::new().build();
let outcomes: Vec<_> = [AlignmentFallback::SkipChunk, AlignmentFallback::Error]
.into_iter()
.map(|policy| {
let b = AlignmentSetBuilder::new().with_fallback(policy).build();
let outcome = on_a_miss(&b, decided_by(&a, &Lang::En, "anything"));
(policy, b.id(), outcome)
})
.collect();
for (policy, b, outcome) in &outcomes {
assert!(
matches!(
outcome,
Err(AlignError::ForeignResolution(foreign))
if foreign.made_by() == a.id() && foreign.asked() == *b
),
"{policy}: every outcome {outcomes:?}"
);
let shown = outcome.as_ref().expect_err("refused").to_string();
assert!(shown.contains(&a.id().to_string()), "{shown}");
assert!(shown.contains(&b.to_string()), "{shown}");
}
}
#[test]
fn a_miss_with_no_resolution_skips_or_errors_as_before() {
let skip = AlignmentSetBuilder::new().build();
let skipped = on_a_miss(&skip, decided_by(&skip, &Lang::En, "anything"))
.expect("a SkipChunk miss is not an error");
assert!(matches!(skipped.cause(), Some(UnalignedCause::Skipped)));
let error = AlignmentSetBuilder::new()
.with_fallback(AlignmentFallback::Error)
.build();
assert!(matches!(
on_a_miss(&error, decided_by(&error, &Lang::En, "anything")),
Err(AlignError::LanguageUnsupported(Lang::En))
));
}
#[test]
fn a_resolution_off_its_route_is_refused_by_name() {
let set = AlignmentSetBuilder::new().build();
let resolution = SetResolution {
made_by: set.id(),
language: Lang::En,
resolution: Some(
english_seam()
.detect_oov("AT&T")
.expect("detect_oov")
.decide(asry::emissions::default_oov_policy),
),
};
let err = on_a_miss(&set, resolution).expect_err("decisions on a miss");
assert!(
matches!(
err,
AlignError::MisroutedResolution(ref misrouted)
if misrouted.decided()
&& *misrouted.route() == AlignmentBinding::Miss(AlignmentFallback::SkipChunk)
),
"{err:?}"
);
let shown = err.to_string();
assert!(shown.contains("holds an aligner's decisions"), "{shown}");
assert!(shown.contains("Miss(SkipChunk)"), "{shown}");
}
fn models_dir() -> std::path::PathBuf {
std::env::var_os("ALIGNKIT_TEST_MODELS").map_or_else(
|| crate::tests::models_root().join("alignkit"),
std::path::PathBuf::from,
)
}
fn en_aligner() -> Aligner {
Aligner::from_paths(
Lang::En,
&models_dir().join("base960h_aligner.mlmodelc"),
Box::new(EnglishNormalizer::new()),
)
.expect("load base960h_aligner.mlmodelc as an En aligner (set ALIGNKIT_TEST_MODELS)")
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn lookup_hits_registered_language() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
assert_eq!(set.len(), 1);
assert!(matches!(set.lookup(&Lang::En), AlignmentLookup::Hit(_)));
assert_eq!(set.resolve(&Lang::En).binding(), AlignmentBinding::Exact);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn lookup_misses_unregistered_language_without_any() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
assert!(matches!(
set.lookup(&Lang::Zh),
AlignmentLookup::Miss(AlignmentFallback::SkipChunk)
));
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn strict_lookup_prefers_lang_over_any() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.register(AlignerKey::Any, en_aligner())
.build();
assert!(matches!(set.lookup(&Lang::En), AlignmentLookup::Hit(_)));
assert!(matches!(
set.lookup(&Lang::Zh),
AlignmentLookup::AnyFallback(_)
));
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn any_fallback_can_match_the_requested_language() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
assert!(matches!(
set.lookup(&Lang::En),
AlignmentLookup::AnyFallback(_)
));
assert_eq!(
set.resolve(&Lang::En).binding(),
AlignmentBinding::AnyFallback(Lang::En)
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
#[should_panic(expected = "cannot accept an aligner built for")]
fn register_panics_on_language_mismatch() {
let _ = AlignmentSetBuilder::new().register(AlignerKey::Lang(Lang::Zh), en_aligner());
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn any_fallback_detection_names_the_requested_language() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let detection = set
.detect_oov("hello AT&T", &Lang::Zh)
.expect("detect_oov on the Any-fallback aligner");
assert_eq!(detection.language(), &Lang::Zh);
let events = detection.events().expect("the Any aligner read the text");
assert!(!events.is_empty(), "the `&` is an OOV event");
assert!(events.iter().all(|event| event.language() == &Lang::Zh));
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn any_fallback_aligns_a_cross_language_request_with_an_oov_decision() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let samples = load_jfk_wav();
let text = JFK_TRANSCRIPT.replace("Americans", "Américans");
let detection = set
.detect_oov(&text, &Lang::Zh)
.expect("detect_oov on the Any fallback");
assert_eq!(detection.language(), &Lang::Zh);
assert!(
detection
.events()
.is_some_and(|events| events.iter().any(|event| event.char() == Some('é'))),
"the `é` of `Américans` is an OOV event"
);
let resolution = detection.decide(default_policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let alignment = set
.align_chunk(
&Lang::Zh,
&samples,
&whole_chunk_is_speech(&samples),
&text,
clock,
&abort,
resolution,
)
.expect("Any-fallback alignment must not fail on the decisions");
assert!(
!alignment.words().is_empty(),
"the English Any aligner must align English speech to English words"
);
}
fn aligned_with(
set: &AlignmentSet,
requested: &Lang,
text: &str,
policy: fn(&SetOovEvent<'_>) -> OovDecision,
) -> Result<UnitAlignment, AlignError> {
let samples = load_jfk_wav();
let resolution = set
.detect_oov(text, requested)
.expect("detect_oov")
.decide(policy);
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
set.align_chunk(
requested,
&samples,
&whole_chunk_is_speech(&samples),
text,
clock,
&abort,
resolution,
)
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn an_english_fallback_judges_a_korean_requests_oov_under_korean() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let text = JFK_TRANSCRIPT.replace("Americans", "Américans");
let refused = aligned_with(&set, &Lang::Ko, &text, english_wildcard_korean_fail_closed)
.expect_err("the Korean policy fails the `é` closed");
let AlignError::Refused(refusal) = refused else {
panic!("the refusal must be named, got {refused:?}");
};
assert_eq!(
refusal
.events()
.iter()
.map(RefusedOov::char)
.collect::<Vec<_>>(),
[Some('é')]
);
let aligned = aligned_with(&set, &Lang::Ko, &text, korean_wildcard_english_fail_closed)
.expect("the Korean policy wildcards the `é`");
assert!(!aligned.words().is_empty(), "the chunk aligns");
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_korean_refusal_through_the_english_fallback_is_reported_under_korean() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let text = "ask not what your country can do for you, AT&T b4d";
let err = aligned_with(&set, &Lang::Ko, text, |_| OovDecision::FailClosed)
.expect_err("a policy failing every event closed refuses the chunk");
let shown = format!("{err:?}");
let AlignError::Refused(refusal) = err else {
panic!("the refusal must be named, got {shown}");
};
assert!(
refusal.events().len() >= 2,
"the `&` and the `4` are both refused: {shown}"
);
assert!(
refusal
.events()
.iter()
.all(|event| event.language() == &Lang::Ko),
"{shown}"
);
assert!(!shown.contains("En"), "{shown}");
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn an_exact_hit_is_judged_under_its_own_language() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
let text = JFK_TRANSCRIPT.replace("Americans", "Américans");
let aligned = aligned_with(&set, &Lang::En, &text, english_wildcard_korean_fail_closed)
.expect("English wildcards the `é`");
assert!(!aligned.words().is_empty());
assert!(matches!(
aligned_with(&set, &Lang::En, &text, korean_wildcard_english_fail_closed),
Err(AlignError::Refused(_))
));
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_fallback_resolution_stays_bound_to_its_text_and_its_aligner() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let other = AlignmentSetBuilder::new()
.register(AlignerKey::Any, en_aligner())
.build();
let samples = load_jfk_wav();
let speech = whole_chunk_is_speech(&samples);
let text = JFK_TRANSCRIPT.replace("Americans", "Américans");
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let decided = || {
set
.detect_oov(&text, &Lang::Ko)
.expect("detect_oov")
.decide(korean_wildcard_english_fail_closed)
};
let another_text = text.replace("country", "county");
let err = set
.align_chunk(
&Lang::Ko,
&samples,
&speech,
&another_text,
clock,
&abort,
decided(),
)
.expect_err("decided in another text");
assert!(
matches!(err, AlignError::Alignment(EmissionsError::Tokenization(_))),
"{err:?}"
);
let err = other
.align_chunk(
&Lang::Ko,
&samples,
&speech,
&text,
clock,
&abort,
decided(),
)
.expect_err("detected by another aligner");
assert!(
matches!(
err,
AlignError::ForeignResolution(ref foreign)
if foreign.made_by() == set.id() && foreign.asked() == other.id()
),
"{err:?}"
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_resolution_decided_for_another_request_is_refused_before_dispatch() {
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.register(AlignerKey::Any, en_aligner())
.build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let assert_decision_language =
|requested: Lang, decided_for: Lang, samples: &[f32], case: &str| {
let resolution = set
.detect_oov("test", &decided_for)
.expect("detect_oov")
.decide(default_policy);
let err = set
.align_chunk(&requested, samples, &[], "test", clock, &abort, resolution)
.expect_err("decisions made for another request must be refused");
assert!(
matches!(
err,
AlignError::DecisionLanguage(ref e)
if *e.requested() == requested && *e.found() == decided_for
),
"{case}: must be the typed DecisionLanguage, got {err:?}"
);
};
let in_window = vec![0.0f32; 16_000];
assert_decision_language(
Lang::Fr,
Lang::Zh,
&in_window,
"one Any aligner, two requests",
);
assert_decision_language(Lang::En, Lang::Zh, &in_window, "in-window");
let oversized = vec![0.0f32; crate::audio::align::encode::ENCODER_WINDOW_SAMPLES + 1];
assert_decision_language(Lang::En, Lang::Zh, &oversized, "oversized");
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn another_registrys_miss_resolution_is_refused_on_a_hit() {
let empty = AlignmentSetBuilder::new().build();
let resolution = empty
.detect_oov("test", &Lang::En)
.expect("a miss detects nothing")
.decide(default_policy);
assert!(resolution.resolved().is_none());
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let err = set
.align_chunk(
&Lang::En,
&[0.0f32; 16_000],
&[],
"test",
clock,
&abort,
resolution,
)
.expect_err("no aligner here detected these decisions");
assert!(
matches!(
err,
AlignError::ForeignResolution(ref foreign)
if foreign.made_by() == empty.id() && foreign.asked() == set.id()
),
"{err:?}"
);
}
fn routes(policy: AlignmentFallback) -> [(&'static str, AlignmentSet); 3] {
[
(
"hit",
AlignmentSetBuilder::new()
.with_fallback(policy)
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build(),
),
(
"fallback",
AlignmentSetBuilder::new()
.with_fallback(policy)
.register(AlignerKey::Any, en_aligner())
.build(),
),
(
"miss",
AlignmentSetBuilder::new().with_fallback(policy).build(),
),
]
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_foreign_resolution_is_refused_on_every_route() {
let samples = load_jfk_wav();
let text = "ask not what your country can do for you, AT&T b4d";
let a = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let mut outcomes = Vec::new();
for policy in [AlignmentFallback::SkipChunk, AlignmentFallback::Error] {
for (route, b) in routes(policy) {
let resolution = a
.detect_oov(text, &Lang::En)
.expect("detect_oov")
.decide(fail_closed_policy);
assert!(
resolution
.resolved()
.is_some_and(|resolved| !resolved.is_empty()),
"A's aligner read the text and decided its events"
);
let outcome = b.align_chunk(&Lang::En, &samples, &[], text, clock, &abort, resolution);
outcomes.push((policy, route, b.id(), outcome));
}
}
let every: Vec<String> = outcomes
.iter()
.map(|(policy, route, _, outcome)| format!("{policy} {route}: {outcome:?}"))
.collect();
for (policy, route, b, outcome) in &outcomes {
assert!(
matches!(
outcome,
Err(AlignError::ForeignResolution(foreign))
if foreign.made_by() == a.id() && foreign.asked() == *b
),
"{policy} {route}: every outcome {every:#?}"
);
}
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_same_set_resolution_aligns_as_the_aligner_does() {
let samples = load_jfk_wav();
let speech = whole_chunk_is_speech(&samples);
let text = JFK_TRANSCRIPT.replace("Americans", "Américans");
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock");
let abort = AtomicBool::new(false);
let set = AlignmentSetBuilder::new()
.register(AlignerKey::Lang(Lang::En), en_aligner())
.build();
let through_the_set = set
.align_chunk(
&Lang::En,
&samples,
&speech,
&text,
clock,
&abort,
set
.detect_oov(&text, &Lang::En)
.expect("detect_oov")
.decide(default_policy),
)
.expect("the set aligns");
let aligner = en_aligner();
let direct = aligner
.align_chunk(
&samples,
&speech,
&text,
clock,
&abort,
aligner
.detect_oov(&text)
.expect("detect_oov")
.decide(asry::emissions::default_oov_policy),
)
.expect("the aligner aligns");
assert!(!direct.words().is_empty());
let words = |alignment: &UnitAlignment| {
alignment
.words()
.iter()
.map(|word| {
(
word.text().to_owned(),
word.range().start_pts(),
word.range().end_pts(),
word.score().to_bits(),
)
})
.collect::<Vec<_>>()
};
assert_eq!(words(&through_the_set), words(&direct));
}
const JFK_TRANSCRIPT: &str = "And so my fellow Americans ask not what your country can do for you, \
ask what you can do for your country.";
fn whole_chunk_is_speech(samples: &[f32]) -> [TimeRange; 1] {
[TimeRange::new(0, samples.len() as i64, ANALYSIS_TIMEBASE)]
}
fn load_jfk_wav() -> Vec<f32> {
let path = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests/whisper/fixtures/audio/jfk.wav");
let mut reader = hound::WavReader::open(&path)
.unwrap_or_else(|e| panic!("open the jfk.wav fixture at {path:?}: {e}"));
assert_eq!(reader.spec().sample_rate, 16_000, "fixture must be 16 kHz");
reader
.samples::<i16>()
.map(|s| f32::from(s.expect("valid sample")) / 32_768.0)
.collect()
}