use super::*;
use core::num::NonZeroU32;
fn staged(real_samples: usize, available_frames: usize) -> usize {
truncated_frame_count(AcousticGeometry::WAV2VEC2, real_samples, available_frames)
}
fn contract(blank: u32, geometry: AcousticGeometry) -> AcousticContract {
AcousticContract::new(
blank,
geometry,
crate::audio::align::acoustic::Tokenization::new(
crate::audio::align::acoustic::WordDelimiter::Pipe,
crate::audio::align::acoustic::LetterCase::Upper,
crate::audio::align::acoustic::Granularity::Character,
&[],
),
OutputKind::LogProbabilities,
)
}
fn geometry(receptive_field: u32, stride: u32) -> AcousticGeometry {
AcousticGeometry::new(
16_000,
NonZeroU32::new(receptive_field).expect("nonzero"),
NonZeroU32::new(stride).expect("nonzero"),
)
.expect("a geometry asry's seam times")
}
fn bundled_seam() -> asry::emissions::EmissionsAligner {
asry::emissions::EmissionsAligner::builder(
crate::audio::align::Lang::En,
crate::audio::align::vocab::tokenizer_json_bytes(),
)
.normalizer(Box::new(crate::audio::align::EnglishNormalizer::new()))
.blank_token_id(crate::audio::align::vocab::BLANK_ID)
.build()
.expect("build the En seam from the bundled tokenizer")
}
fn clock() -> asry::emissions::OutputClock {
asry::emissions::OutputClock::new(0, asry::time::ANALYSIS_TIMEBASE, 0).expect("clock")
}
fn through_asry(output: EncoderOutput) -> Result<asry::emissions::Emissions, AlignError> {
let seam = bundled_seam();
let resolution = seam
.detect_oov("A")
.expect("detect_oov")
.decide(asry::emissions::wildcard_all_policy);
let prepared = seam
.prepare(
&[0.1; 400],
&asry::emissions::SpeechSpans::all_speech(),
"A",
resolution,
clock(),
&core::sync::atomic::AtomicBool::new(false),
)
.expect("prepare");
assert!(!prepared.is_trivial(), "`A` is alignable");
prepared.encode_with(|_| Ok::<_, AlignError>(output))
}
#[test]
fn truncated_frame_count_zero_samples_is_zero() {
assert_eq!(staged(0, 2999), 0);
}
#[test]
fn truncated_frame_count_sub_receptive_field_is_one_frame() {
for real_samples in [1, 200, 320, 399, 400] {
assert_eq!(
staged(real_samples, 2999),
1,
"real_samples={real_samples} is within the receptive field: one frame"
);
}
}
#[test]
fn truncated_frame_count_no_phantom_frame_from_receptive_field_slack() {
assert_eq!(staged(321, 2999), 1); assert_eq!(staged(641, 2999), 1); }
#[test]
fn truncated_frame_count_adds_one_frame_per_hop_past_the_receptive_field() {
assert_eq!(staged(401, 2999), 1);
assert_eq!(staged(720, 2999), 2);
assert_eq!(staged(1040, 2999), 3);
}
#[test]
fn truncated_frame_count_receptive_field_boundary_is_pinned_by_hand() {
assert_eq!(staged(399, 2999), 1);
assert_eq!(staged(400, 2999), 1);
assert_eq!(staged(719, 2999), 1);
assert_eq!(staged(720, 2999), 2);
}
#[test]
fn the_staged_geometry_is_pinned() {
let geometry = AcousticGeometry::WAV2VEC2;
assert_eq!(geometry.sample_rate().get(), 16_000);
assert_eq!(geometry.receptive_field().get(), 400);
assert_eq!(geometry.stride().get(), 320);
assert_eq!(HOP_SAMPLES, 320);
assert_eq!(AcousticContract::BASE960H.geometry(), geometry);
}
#[test]
fn truncated_frame_count_reference_short_clip() {
assert_eq!(staged(48_000, 2_999), 149);
}
#[test]
fn truncated_frame_count_full_window_is_the_model_frame_count() {
assert_eq!(staged(ENCODER_WINDOW_SAMPLES, 2_999), 2_999);
}
#[test]
fn truncated_frame_count_approaches_the_full_window_without_overshoot() {
assert_eq!(staged(2_999 * HOP_SAMPLES, 2_999), 2_998);
assert_eq!(staged(959_760, 2_999), 2_999);
assert_eq!(staged(ENCODER_WINDOW_SAMPLES - 1, 2_999), 2_999);
}
#[test]
fn truncated_frame_count_clamp_engages_only_below_the_formula() {
assert_eq!(staged(48_000, 100), 100);
assert_eq!(staged(48_000, 149), 149); }
#[test]
fn truncated_frame_count_never_exceeds_available_frames_near_full_window() {
let available_frames = 2_999;
for real_samples in [
ENCODER_WINDOW_SAMPLES - 1,
ENCODER_WINDOW_SAMPLES,
available_frames * HOP_SAMPLES,
available_frames * HOP_SAMPLES + 1,
ENCODER_WINDOW_SAMPLES + HOP_SAMPLES,
ENCODER_WINDOW_SAMPLES + 10 * HOP_SAMPLES,
] {
let t = staged(real_samples, available_frames);
assert!(
t <= available_frames,
"staged({real_samples}, {available_frames}) = {t} exceeds available_frames"
);
}
assert_eq!(
staged(ENCODER_WINDOW_SAMPLES + HOP_SAMPLES, available_frames),
2_999
);
assert_eq!(
staged(ENCODER_WINDOW_SAMPLES + 10 * HOP_SAMPLES, available_frames),
2_999
);
}
#[test]
fn a_640_320_contract_truncates_720_samples_to_one_frame() {
let wide = geometry(640, 320);
assert_eq!(
wide.frames(ENCODER_WINDOW_SAMPLES),
AcousticGeometry::WAV2VEC2.frames(ENCODER_WINDOW_SAMPLES),
"both geometries make the staged window's 2999 frames: the declaration cannot tell them apart"
);
assert_eq!(truncated_frame_count(wide, 720, 2_999), 1);
assert_eq!(staged(720, 2_999), 2);
assert_eq!(truncated_frame_count(wide, 1, 2_999), 1);
assert_eq!(truncated_frame_count(wide, 959, 2_999), 1);
assert_eq!(truncated_frame_count(wide, 960, 2_999), 2);
assert_eq!(truncated_frame_count(wide, 0, 2_999), 0);
}
#[test]
fn check_frame_count_refuses_a_geometry_the_declaration_contradicts() {
let window = NonZeroUsize::new(ENCODER_WINDOW_SAMPLES).expect("nonzero");
let frames = NonZeroUsize::new(EXPECTED_OUTPUT_FRAMES).expect("nonzero");
assert_eq!(
check_frame_count(AcousticGeometry::WAV2VEC2, window, frames),
Ok(())
);
assert_eq!(
check_frame_count(geometry(640, 320), window, frames),
Ok(())
);
for (geometry, window, derived) in [
(geometry(400, 160), ENCODER_WINDOW_SAMPLES, 5_998),
(geometry(400, 321), ENCODER_WINDOW_SAMPLES, 2_990),
(geometry(400, 320), 399, 0),
] {
let window = NonZeroUsize::new(window).expect("nonzero");
let Err(AlignerError::FrameCountMismatch(mismatch)) =
check_frame_count(geometry, window, frames)
else {
panic!("{geometry:?} must be refused against a {window}-sample window of 2999 frames");
};
assert_eq!(mismatch.geometry(), geometry);
assert_eq!(mismatch.window(), window.get());
assert_eq!(mismatch.declared(), EXPECTED_OUTPUT_FRAMES);
assert_eq!(mismatch.derived(), derived);
}
}
#[test]
fn encoder_input_from_samples_binds_real_length_to_the_slice() {
let chunk = vec![0.0f32; 176_000];
let input = EncoderInput::from_samples(&chunk);
assert_eq!(input.real_samples, 176_000);
assert_eq!(input.encoder_input.len(), 176_000);
assert_eq!(staged(input.real_samples, 2_999), 549);
assert_eq!(staged(175_360, 2_999), 547);
assert_ne!(staged(input.real_samples, 2_999), staged(175_360, 2_999));
}
#[test]
fn encoder_input_gate_binds_real_length_independent_of_the_padded_buffer() {
let real_len = 200usize;
let padded_buffer = vec![0.0f32; 400];
let input = EncoderInput::new(&padded_buffer, real_len);
assert_eq!(input.real_samples, 200); assert_eq!(input.encoder_input.len(), 400);
assert_eq!(staged(input.real_samples, 2_999), 1);
assert_eq!(staged(padded_buffer.len(), 2_999), 1);
}
#[test]
fn from_prepared_records_the_true_pre_pad_provenance_not_the_padded_length() {
use core::sync::atomic::AtomicBool;
let aligner = bundled_seam();
let samples: Vec<f32> = (0..200).map(|i| (i as f32 * 0.05).sin() * 0.2).collect();
let abort = AtomicBool::new(false);
let resolution = aligner
.detect_oov("test")
.expect("detect_oov")
.decide(asry::emissions::wildcard_all_policy);
let prepared = aligner
.prepare(
&samples,
&crate::audio::align::SpeechSpans::all_speech(),
"test",
resolution,
clock(),
&abort,
)
.expect("prepare 200 real samples with alignable text");
assert!(
!prepared.is_trivial(),
"`test` must tokenize to alignable tokens, or there is no prepared buffer to test"
);
assert_eq!(
prepared.encoder_input().len(),
400,
"asry pads 200 real samples up to the 400-sample receptive field"
);
let via_prepared = EncoderInput::from_prepared(&prepared);
assert_eq!(
via_prepared.real_samples, 200,
"from_prepared must record the true pre-pad real_samples (200), never the padded 400"
);
let via_raw = EncoderInput::from_samples(prepared.encoder_input());
assert_eq!(
via_raw.real_samples, 400,
"from_samples records the buffer length it is handed (400) — the distinguisher"
);
assert_ne!(
via_prepared.real_samples, via_raw.real_samples,
"the two doors record DIFFERENT provenance for the one buffer; only the frame \
count coincides, which is why a count-only test cannot bind from_prepared"
);
}
#[test]
fn check_window_refuses_a_buffer_longer_than_the_models_window() {
let window = NonZeroUsize::new(ENCODER_WINDOW_SAMPLES).expect("nonzero");
let Err(AlignError::InputTooLong(too_long)) = check_window(ENCODER_WINDOW_SAMPLES + 1, window)
else {
panic!("one sample past the window must be refused");
};
assert_eq!(
(too_long.got(), too_long.max()),
(ENCODER_WINDOW_SAMPLES + 1, ENCODER_WINDOW_SAMPLES)
);
assert!(check_window(ENCODER_WINDOW_SAMPLES, window).is_ok());
assert!(check_window(0, window).is_ok());
let short = NonZeroUsize::new(480_000).expect("nonzero");
assert!(check_window(480_001, short).is_err());
assert!(check_window(480_000, short).is_ok());
}
#[test]
fn encoder_input_binds_a_full_window_as_real() {
let full = vec![0.0f32; ENCODER_WINDOW_SAMPLES];
let input = EncoderInput::from_samples(&full);
assert_eq!(input.real_samples, ENCODER_WINDOW_SAMPLES);
assert_eq!(input.encoder_input.len(), ENCODER_WINDOW_SAMPLES);
}
const STAGED_BAND: Option<SentinelBand> = AcousticContract::BASE960H.sentinel_band();
#[test]
fn check_sentinel_band_accepts_real_log_probs() {
let data = [0.0, -0.06, -19.0, -21.75, -30.02, -30.81];
assert!(check_sentinel_band(&data, STAGED_BAND, ComputeUnits::CpuOnly).is_ok());
}
#[test]
fn check_sentinel_band_accepts_an_empty_matrix() {
assert!(check_sentinel_band(&[], STAGED_BAND, ComputeUnits::CpuOnly).is_ok());
}
#[test]
fn check_sentinel_band_rejects_the_fp16_log_zero_sentinel() {
let data = [0.0, -1.5, -45_440.0, -20.0];
let Err(err) = check_sentinel_band(&data, STAGED_BAND, ComputeUnits::All) else {
panic!("the -45440 fp16 log(0) sentinel must be rejected");
};
let AlignError::CorruptEmissions(ref e) = err else {
panic!("expected AlignError::CorruptEmissions, got {err:?}");
};
assert_eq!(e.compute(), ComputeUnits::All);
assert_eq!(e.band(), SentinelBand::Fp16Saturation);
assert_eq!(e.min(), -45_440.0);
assert_eq!(e.cells(), 1);
assert_eq!(e.total(), 4);
}
#[test]
fn the_band_holds_its_ceiling_and_nothing_above_it() {
let ceiling = SentinelBand::Fp16Saturation.ceiling();
assert_eq!(ceiling, -32_768.0);
assert!(check_sentinel_band(&[ceiling], STAGED_BAND, ComputeUnits::CpuOnly).is_err());
assert!(
check_sentinel_band(&[-32_767.0], STAGED_BAND, ComputeUnits::CpuOnly).is_ok(),
"a value above the binade is no saturated fp16 log(0)"
);
}
#[test]
fn check_sentinel_band_leaves_non_finite_values_to_asrys_scan() {
assert!(check_sentinel_band(&[f32::NAN], STAGED_BAND, ComputeUnits::CpuOnly).is_ok());
assert!(check_sentinel_band(&[f32::INFINITY], STAGED_BAND, ComputeUnits::CpuOnly).is_ok());
assert!(check_sentinel_band(&[f32::NEG_INFINITY], STAGED_BAND, ComputeUnits::CpuOnly).is_err());
}
fn guard_row(
tail: f32,
band: Option<SentinelBand>,
) -> Result<asry::emissions::Emissions, AlignError> {
let two = NonZeroUsize::new(2).expect("nonzero");
let output = RawEmissions {
frames: 1,
vocab_size: two,
data: vec![0.0, tail],
}
.check_value_domain(band, ComputeUnits::All)?
.into_output();
through_asry(output)
}
#[test]
fn a_generic_contract_refuses_no_finite_log_probability() {
let generic = contract(0, AcousticGeometry::WAV2VEC2);
assert_eq!(generic.sentinel_band(), None);
for tail in [-101.0f32, -500.0, -32_000.0, -40_000.0, -45_440.0] {
let emissions = guard_row(tail, generic.sentinel_band())
.unwrap_or_else(|err| panic!("[0, {tail}] is a row of log-probabilities: {err:?}"));
assert_eq!(emissions.vocab().get(), 2);
}
for tail in [-40_000.0f32, -45_440.0] {
assert!(
matches!(
guard_row(tail, STAGED_BAND),
Err(AlignError::CorruptEmissions(_))
),
"[0, {tail}] is in the staged model's band"
);
}
for tail in [-101.0f32, -500.0, -32_000.0] {
assert!(
guard_row(tail, STAGED_BAND).is_ok(),
"[0, {tail}] is above the staged model's band"
);
}
}
#[test]
fn read_emissions_refuses_every_shape_but_the_declared_one() {
const FRAMES: usize = 3;
let width = WIDTH.get();
let cells = |count: usize| (0..count).map(|i| -(i as f32)).collect::<Vec<f32>>();
let declared = cells(FRAMES * width);
let tensor = MultiArray::from_slice(&[1, FRAMES, width], &declared).expect("build a tensor");
assert_eq!(
read_emissions(&tensor, FRAMES, WIDTH).expect("the declared shape is read"),
declared
);
for shape in [
vec![1, width, FRAMES],
vec![FRAMES, width, 1],
vec![FRAMES, width],
vec![1, FRAMES * width],
vec![1, FRAMES, 32],
vec![1, FRAMES - 1, width],
] {
let tensor =
MultiArray::from_slice(&shape, &cells(shape.iter().product())).expect("build a tensor");
let Err(AlignError::OutputShape(mismatch)) = read_emissions(&tensor, FRAMES, WIDTH) else {
panic!("{shape:?} must be refused as an output-shape mismatch");
};
assert_eq!(mismatch.got(), shape.as_slice());
assert_eq!(mismatch.expected(), [1, FRAMES, width]);
}
}
const WIDTH: NonZeroUsize = match NonZeroUsize::new(crate::audio::align::vocab::VOCAB_SIZE) {
Some(width) => width,
None => unreachable!(),
};
fn uniform_frame(value: f32) -> [f32; crate::audio::align::vocab::VOCAB_SIZE] {
[value; crate::audio::align::vocab::VOCAB_SIZE]
}
#[test]
fn check_log_prob_normalization_accepts_normalized_log_probs() {
let ln29 = f64::from(crate::audio::align::vocab::VOCAB_SIZE as u32).ln();
let uniform = uniform_frame(-ln29 as f32);
let logits: [f32; crate::audio::align::vocab::VOCAB_SIZE] =
core::array::from_fn(|j| (j as f32) * 0.5 - 3.0);
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let z_lse = f64::from(max)
+ logits
.iter()
.map(|&z| (f64::from(z) - f64::from(max)).exp())
.sum::<f64>()
.ln();
let peaked: Vec<f32> = logits
.iter()
.map(|&z| (f64::from(z) - z_lse) as f32)
.collect();
let mut data = Vec::new();
data.extend_from_slice(&uniform);
data.extend_from_slice(&peaked);
assert!(
check_log_prob_normalization(&data, WIDTH, ComputeUnits::CpuOnly).is_ok(),
"normalized log-prob frames (logsumexp ≈ 0) must pass"
);
}
#[test]
fn check_log_prob_normalization_accepts_an_empty_matrix() {
assert!(check_log_prob_normalization(&[], WIDTH, ComputeUnits::CpuOnly).is_ok());
}
#[test]
fn check_log_prob_normalization_rejects_shifted_raw_logits() {
let mut data = Vec::with_capacity(2999 * crate::audio::align::vocab::VOCAB_SIZE);
for _ in 0..2999 {
for j in 0..crate::audio::align::vocab::VOCAB_SIZE {
data
.push(-10.0 - (j as f32) * (10.0 / (crate::audio::align::vocab::VOCAB_SIZE as f32 - 1.0)));
}
}
assert!(
check_sentinel_band(&data, STAGED_BAND, ComputeUnits::CpuOnly).is_ok(),
"shifted raw logits in [-20, -10] are all above the staged band — it cannot catch them"
);
assert!(
data.iter().all(|v| v.is_finite() && *v <= 0.0),
"shifted raw logits are finite and <= 0 — asry's scan cannot catch them"
);
let Err(err) = check_log_prob_normalization(&data, WIDTH, ComputeUnits::CpuOnly) else {
panic!("raw logits shifted into [-20, -10] must be rejected as un-normalized");
};
let AlignError::UnnormalizedEmissions(ref e) = err else {
panic!("expected AlignError::UnnormalizedEmissions, got {err:?}");
};
let logsumexp = e.logsumexp();
assert!(
logsumexp.abs() > 6.6,
"a [-20, -10] shifted frame's |logsumexp| is >= 6.63, got {logsumexp}"
);
assert_eq!(e.tolerance(), log_prob_sum_tolerance(WIDTH));
}
#[test]
fn check_log_prob_normalization_rejects_an_all_zero_frame() {
let data = uniform_frame(0.0);
assert!(
check_sentinel_band(&data, STAGED_BAND, ComputeUnits::CpuOnly).is_ok(),
"an all-zeros frame is above the band"
);
assert!(data.iter().all(|v| v.is_finite() && *v <= 0.0));
let Err(AlignError::UnnormalizedEmissions(e)) =
check_log_prob_normalization(&data, WIDTH, ComputeUnits::All)
else {
panic!("an all-zeros frame (logsumexp = ln 29) must be rejected");
};
let logsumexp = e.logsumexp();
assert_eq!(e.row(), 0);
assert_eq!(e.compute(), ComputeUnits::All); let ln29 = f64::from(crate::audio::align::vocab::VOCAB_SIZE as u32).ln();
assert!(
(logsumexp - ln29).abs() < 1e-5,
"all-zeros logsumexp must be ln(29) ≈ {ln29}, got {logsumexp}"
);
}
#[test]
fn check_log_prob_normalization_names_the_worst_frame() {
let ln29 = f64::from(crate::audio::align::vocab::VOCAB_SIZE as u32).ln();
let normalized = uniform_frame(-ln29 as f32);
let bad_index = 2usize;
let mut data = Vec::new();
for i in 0..5 {
if i == bad_index {
data.extend_from_slice(&uniform_frame(0.0)); } else {
data.extend_from_slice(&normalized); }
}
let Err(AlignError::UnnormalizedEmissions(e)) =
check_log_prob_normalization(&data, WIDTH, ComputeUnits::CpuOnly)
else {
panic!("the un-normalized frame must be rejected");
};
assert_eq!(e.row(), bad_index, "the error must name the worst frame");
}
#[test]
fn check_log_prob_normalization_thresholds_on_the_tolerance() {
let ln29 = f64::from(crate::audio::align::vocab::VOCAB_SIZE as u32).ln();
let tol = log_prob_sum_tolerance(WIDTH);
let inside = uniform_frame((tol / 2.0 - ln29) as f32);
let outside = uniform_frame((2.0 * tol - ln29) as f32);
assert!(
check_log_prob_normalization(&inside, WIDTH, ComputeUnits::CpuOnly).is_ok(),
"logsumexp = TOL/2 is within tolerance"
);
assert!(
check_log_prob_normalization(&outside, WIDTH, ComputeUnits::CpuOnly).is_err(),
"logsumexp = 2·TOL exceeds tolerance"
);
}
#[test]
fn the_allowance_grows_with_the_head_and_still_refuses_an_all_zeros_frame() {
let width = |v: usize| NonZeroUsize::new(v).expect("nonzero");
assert_eq!(log_prob_sum_tolerance(width(29)), 60.0 / 2048.0);
assert_eq!(log_prob_sum_tolerance(width(3_500)), 7_002.0 / 2048.0);
for v in [2usize, 29, 32, 100, 1_000, 3_500, 8_000] {
let tolerance = log_prob_sum_tolerance(width(v));
let ln_v = (v as f64).ln();
assert!(
(v as f64 + 2.0 + 2.0 * ln_v) / 2048.0 <= tolerance,
"{v} classes: the fp16 rounding bound must fit the allowance"
);
let zeros = vec![0.0f32; v];
assert!(
matches!(
check_log_prob_normalization(&zeros, width(v), ComputeUnits::CpuOnly),
Err(AlignError::UnnormalizedEmissions(_))
),
"{v} classes: an all-zeros frame (logsumexp ln {v} = {ln_v}) must be refused"
);
}
let v = 3_500usize;
let wide = vec![(0.5 - (v as f64).ln()) as f32; v];
assert!(
check_log_prob_normalization(&wide, width(v), ComputeUnits::CpuOnly).is_ok(),
"a 3,500-class frame within its own allowance must pass"
);
assert!(0.5 > log_prob_sum_tolerance(WIDTH));
}
#[test]
fn check_log_prob_normalization_frames_rows_by_the_width_it_is_given() {
let data = [-(4.0f32.ln()); 8];
let four = NonZeroUsize::new(4).expect("nonzero");
let two = NonZeroUsize::new(2).expect("nonzero");
assert!(
check_log_prob_normalization(&data, four, ComputeUnits::CpuOnly).is_ok(),
"two 4-class frames of -ln 4 are normalized"
);
let Err(AlignError::UnnormalizedEmissions(e)) =
check_log_prob_normalization(&data, two, ComputeUnits::CpuOnly)
else {
panic!("the same cells read as 2-class frames are not distributions");
};
assert!(
(e.logsumexp() + 2.0f64.ln()).abs() < 1e-6,
"a [-ln 4, -ln 4] row has logsumexp -ln 2, got {}",
e.logsumexp()
);
}
#[test]
fn the_wrap_carries_the_width_the_encoder_read() {
let four = NonZeroUsize::new(4).expect("nonzero");
let output = RawEmissions {
frames: 2,
vocab_size: four,
data: vec![-(4.0f32.ln()); 8],
}
.check_value_domain(None, ComputeUnits::CpuOnly)
.expect("two normalized 4-class frames clear the guard")
.into_output();
assert_eq!(output_shape(&output), (2, four));
let emissions = through_asry(output).expect("and asry takes them as log-probabilities");
assert_eq!(emissions.frames(), 2);
assert_eq!(emissions.vocab(), four);
}
#[test]
fn raw_emissions_check_value_domain_binds_the_guard_and_the_minted_buffer() {
let mut shifted = Vec::with_capacity(4 * crate::audio::align::vocab::VOCAB_SIZE);
for _ in 0..4 {
for j in 0..crate::audio::align::vocab::VOCAB_SIZE {
shifted
.push(-10.0 - (j as f32) * (10.0 / (crate::audio::align::vocab::VOCAB_SIZE as f32 - 1.0)));
}
}
let raw = RawEmissions {
frames: 4,
vocab_size: WIDTH,
data: shifted,
};
assert!(
matches!(
raw.check_value_domain(STAGED_BAND, ComputeUnits::CpuOnly),
Err(AlignError::UnnormalizedEmissions(_))
),
"the minter must reject a shifted-raw-logit tensor as un-normalized"
);
let raw = RawEmissions {
frames: 1,
vocab_size: WIDTH,
data: uniform_frame(0.0).to_vec(),
};
assert!(
matches!(
raw.check_value_domain(STAGED_BAND, ComputeUnits::CpuOnly),
Err(AlignError::UnnormalizedEmissions(_))
),
"the minter must reject an all-zeros frame as un-normalized"
);
let ln29 = f64::from(crate::audio::align::vocab::VOCAB_SIZE as u32).ln();
let normalized = uniform_frame(-ln29 as f32).to_vec();
let token = RawEmissions {
frames: 1,
vocab_size: WIDTH,
data: normalized.clone(),
}
.check_value_domain(STAGED_BAND, ComputeUnits::CpuOnly)
.expect("the minter must accept a normalized log-prob frame");
assert_eq!(token.frames, 1);
assert_eq!(
token.data, normalized,
"the minted token must own the exact buffer the guard validated"
);
}
#[test]
fn into_output_takes_the_log_prob_door_not_the_logit_door() {
use asry::emissions::{EmissionsError, LogProbsValueClass};
let mut data = vec![-20.0f32; crate::audio::align::vocab::VOCAB_SIZE];
data[0] = 0.001;
assert!(
check_sentinel_band(&data, STAGED_BAND, ComputeUnits::CpuOnly).is_ok(),
"min cell -20.0 is far above the staged band: the band cannot catch a positive cell"
);
assert!(
check_log_prob_normalization(&data, WIDTH, ComputeUnits::CpuOnly).is_ok(),
"logsumexp ≈ 0.001 is within the 29-class allowance: normalization cannot catch it"
);
let token = RawEmissions {
frames: 1,
vocab_size: WIDTH,
data,
}
.check_value_domain(STAGED_BAND, ComputeUnits::CpuOnly)
.expect("a frame that clears the band and the normalization guard must mint a token");
let Err(err) = through_asry(token.into_output()) else {
panic!(
"asry accepted a frame with a positive cell (0.001) handed on by into_output: the \
log-prob door must refuse it. Only the logit door — the wrong one — would renormalize \
and accept."
);
};
let AlignError::Alignment(EmissionsError::Value(value)) = err else {
panic!("expected AlignError::Alignment(EmissionsError::Value), got {err:?}");
};
assert_eq!(
value.class(),
LogProbsValueClass::Positive,
"cell 0 (0.001) is finite and > 0 — the positive log-prob-domain class"
);
assert_eq!(value.frame(), 0, "the positive cell is in frame 0");
assert_eq!(
value.vocab_index(),
0,
"the positive cell is at vocab index 0"
);
}
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 encoder_path() -> std::path::PathBuf {
models_dir().join("base960h_aligner.mlmodelc")
}
fn load_encoder() -> Encoder {
staged_encoder(DEFAULT_ENCODER_COMPUTE)
}
fn staged_encoder(compute: ComputeUnits) -> Encoder {
Encoder::load(encoder_path(), &AcousticContract::BASE960H, compute)
.expect("load base960h_aligner.mlmodelc (set ALIGNKIT_TEST_MODELS to the model directory)")
}
fn window_input(samples: &[f32]) -> EncoderInput<'_> {
EncoderInput::from_samples(samples)
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn from_file_loads_and_reports_frame_count() {
let encoder = load_encoder();
assert_eq!(encoder.window_samples(), ENCODER_WINDOW_SAMPLES);
assert_eq!(encoder.frames(), 2_999);
assert_eq!(
encoder.vocab_size().get(),
crate::audio::align::vocab::VOCAB_SIZE
);
assert_eq!(encoder.contract(), &AcousticContract::BASE960H);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn a_geometry_the_model_contradicts_is_refused_at_load() {
for (stride, derived) in [(160u32, 5_998usize), (321, 2_990)] {
let contract = contract(0, geometry(400, stride));
let Err(AlignerError::FrameCountMismatch(mismatch)) =
Encoder::load(encoder_path(), &contract, DEFAULT_ENCODER_COMPUTE)
else {
panic!("a {stride}-sample stride must be refused against the staged model's 2999 frames");
};
assert_eq!(
(mismatch.window(), mismatch.declared(), mismatch.derived()),
(ENCODER_WINDOW_SAMPLES, 2_999, derived)
);
}
let wide = contract(0, geometry(640, 320));
let encoder = Encoder::load(encoder_path(), &wide, DEFAULT_ENCODER_COMPUTE)
.expect("a 640/320 geometry fits the staged declaration");
assert_eq!(encoder.contract(), &wide);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_refuse_a_buffer_longer_than_the_window_before_predicting() {
let encoder = load_encoder();
let too_long = vec![0.0f32; ENCODER_WINDOW_SAMPLES + 1];
let Err(AlignError::InputTooLong(err)) = encoder.emissions(window_input(&too_long)) else {
panic!("a buffer past the window must be refused");
};
assert_eq!(
(err.got(), err.max()),
(ENCODER_WINDOW_SAMPLES + 1, ENCODER_WINDOW_SAMPLES)
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_on_full_window_produces_correctly_shaped_finite_log_probs() {
let encoder = load_encoder();
let samples = vec![0.0f32; ENCODER_WINDOW_SAMPLES];
let raw = encoder
.emissions_raw(window_input(&samples))
.expect("emissions on silence");
assert_eq!(raw.frames, encoder.frames());
assert_eq!(
raw.data.len(),
raw.frames * crate::audio::align::vocab::VOCAB_SIZE
);
assert!(
raw.data.iter().all(|v| v.is_finite()),
"all log-probs finite"
);
assert!(
raw.data.iter().all(|&v| v <= 0.0),
"log-probs must satisfy log(p) <= 0"
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_have_no_fp16_log_zero_sentinel() {
let encoder = load_encoder();
let samples = load_jfk_wav();
let raw = encoder
.emissions_raw(window_input(&samples))
.expect("emissions on jfk.wav");
let min = raw.data.iter().copied().fold(f32::INFINITY, f32::min);
let band = SentinelBand::Fp16Saturation;
let sentinels = raw.data.iter().filter(|v| band.holds(**v)).count();
assert_eq!(
sentinels,
0,
"{sentinels} of {} emission cells are at or below {} (min = {min}) — the fp16 `log(0)` \
sentinel. The encoder is on {:?}; an ANE placement corrupts this model's emissions and \
cannot be used. See DEFAULT_ENCODER_COMPUTE.",
raw.data.len(),
band.ceiling(),
DEFAULT_ENCODER_COMPUTE,
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_reject_an_ane_corrupted_matrix() {
let encoder = Encoder::load(
encoder_path(),
&AcousticContract::BASE960H,
ComputeUnits::All,
)
.expect("load base960h_aligner.mlmodelc on ComputeUnits::All");
let samples = load_jfk_wav();
let Err(err) = encoder.emissions(window_input(&samples)) else {
panic!(
"an ANE-corrupted emission matrix was accepted. asry's log-probability scan cannot catch \
this — -45440 is finite and <= 0 — so the caller now has plausible, silently wrong word \
timings. The staged contract's sentinel band is the only thing standing here."
);
};
let AlignError::CorruptEmissions(ref e) = err else {
panic!("expected AlignError::CorruptEmissions, got {err:?}");
};
let (compute, min, cells, total) = (e.compute(), e.min(), e.cells(), e.total());
assert_eq!(compute, ComputeUnits::All);
assert_eq!(total, 549 * crate::audio::align::vocab::VOCAB_SIZE);
assert!(
cells > 0 && cells <= total,
"corrupt cells: {cells}/{total}"
);
assert_eq!(e.band(), SentinelBand::Fp16Saturation);
assert!(
e.band().holds(min),
"reported min {min} must be in the band it tripped"
);
let rendered = AlignError::CorruptEmissions(e.clone()).to_string();
assert!(
rendered.contains("All"),
"error must name the placement: {rendered}"
);
println!("rejected with: {rendered}");
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_accept_the_cpu_and_gpu_placement() {
let encoder = Encoder::load(
encoder_path(),
&AcousticContract::BASE960H,
ComputeUnits::CpuAndGpu,
)
.expect("load base960h_aligner.mlmodelc on ComputeUnits::CpuAndGpu");
let samples = load_jfk_wav();
let output = encoder
.emissions(window_input(&samples))
.expect("CpuAndGpu emissions are clean log-probs and must pass the band");
let (frames, vocab) = output_shape(&output);
assert_eq!(frames, 549);
assert_eq!(vocab.get(), crate::audio::align::vocab::VOCAB_SIZE);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_accept_the_default_placement_on_real_speech() {
let encoder = load_encoder();
let samples = load_jfk_wav();
let output = encoder
.emissions(window_input(&samples))
.unwrap_or_else(|e| panic!("the SHIPPING placement must produce clean log-probs: {e}"));
let (frames, vocab) = output_shape(&output);
assert_eq!(frames, 549);
assert_eq!(vocab.get(), crate::audio::align::vocab::VOCAB_SIZE);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_pass_the_normalization_guard_on_real_speech() {
for compute in [ComputeUnits::CpuOnly, ComputeUnits::CpuAndGpu] {
let encoder = Encoder::load(encoder_path(), &AcousticContract::BASE960H, compute)
.unwrap_or_else(|e| panic!("load base960h_aligner.mlmodelc on {compute:?}: {e}"));
for (name, samples) in [("jfk", load_jfk_wav()), ("ted_60", load_ted_60_wav())] {
let raw = encoder
.emissions_raw(window_input(&samples))
.unwrap_or_else(|e| panic!("{compute:?} {name}: emissions_raw: {e}"));
check_sentinel_band(&raw.data, STAGED_BAND, compute)
.unwrap_or_else(|e| panic!("{compute:?} {name}: real emissions tripped the band: {e}"));
check_log_prob_normalization(&raw.data, raw.vocab_size, compute).unwrap_or_else(|e| {
panic!("{compute:?} {name}: real emissions tripped the normalization guard: {e}")
});
let worst = raw
.data
.as_chunks::<{ crate::audio::align::vocab::VOCAB_SIZE }>()
.0
.iter()
.map(|frame| {
let max = f64::from(frame.iter().copied().fold(f32::NEG_INFINITY, f32::max));
let sum: f64 = frame.iter().map(|&x| (f64::from(x) - max).exp()).sum();
(max + sum.ln()).abs()
})
.fold(0.0f64, f64::max);
let tolerance = log_prob_sum_tolerance(raw.vocab_size);
println!(
"{compute:?} {name}: {} frames, worst |logsumexp| = {worst:.6e} (allowance {tolerance:e})",
raw.frames,
);
assert!(
worst < tolerance,
"{compute:?} {name}: worst |logsumexp| {worst} is not under the allowance {tolerance} — \
the allowance's measured headroom has been lost"
);
}
}
}
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}"));
let spec = reader.spec();
assert_eq!(spec.channels, 1, "fixture must be mono");
assert_eq!(spec.sample_rate, 16_000, "fixture must be 16 kHz");
assert_eq!(spec.sample_format, hound::SampleFormat::Int);
reader
.samples::<i16>()
.map(|s| f32::from(s.expect("valid sample")) / 32_768.0)
.collect()
}
fn load_ted_60_wav() -> Vec<f32> {
let path = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests/whisper/fixtures/audio/ted_60.wav");
let mut reader = hound::WavReader::open(&path)
.unwrap_or_else(|e| panic!("open the ted_60.wav fixture at {path:?}: {e}"));
let spec = reader.spec();
assert_eq!(spec.channels, 1, "fixture must be mono");
assert_eq!(spec.sample_rate, 16_000, "fixture must be 16 kHz");
assert_eq!(spec.sample_format, hound::SampleFormat::Int);
let samples: Vec<f32> = reader
.samples::<i16>()
.map(|s| f32::from(s.expect("valid sample")) / 32_768.0)
.collect();
assert_eq!(
samples.len(),
ENCODER_WINDOW_SAMPLES,
"ted_60.wav must fill the encoder window exactly (the zero-padding-free path)"
);
samples
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_wraps_into_validated_emissions() {
let encoder = load_encoder();
let samples = vec![0.0f32; 48_000];
let output = encoder
.emissions(window_input(&samples))
.expect("emissions clear the guards");
assert!(matches!(output, EncoderOutput::LogProbs { .. }));
let emissions = through_asry(output).expect("asry makes validated Emissions of them");
assert_eq!(emissions.frames(), 149);
assert_eq!(
emissions.vocab().get(),
crate::audio::align::vocab::VOCAB_SIZE
);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_on_short_input_truncates_to_hermetic_formula() {
let encoder = load_encoder();
let samples = vec![0.0f32; 48_000];
let raw = encoder
.emissions_raw(window_input(&samples))
.expect("emissions on short input");
assert_eq!(raw.frames, staged(48_000, encoder.frames()));
assert_eq!(raw.frames, 149);
assert_eq!(raw.data.len(), 149 * crate::audio::align::vocab::VOCAB_SIZE);
}
#[test]
#[ignore = "requires local alignkit models (ALIGNKIT_TEST_MODELS)"]
fn emissions_is_deterministic_across_repeated_calls() {
let encoder = load_encoder();
let samples: Vec<f32> = (0..ENCODER_WINDOW_SAMPLES)
.map(|i| 0.01 * (i as f32 * 0.001).sin())
.collect();
let first = encoder
.emissions_raw(window_input(&samples))
.expect("first emissions call");
let second = encoder
.emissions_raw(window_input(&samples))
.expect("second emissions call");
assert_eq!(first.frames, second.frames);
assert_eq!(
first.data, second.data,
"repeated emissions_raw() must be bit-identical"
);
}
use crate::{AxisRange, FeatureInfo, ModelDescription, model::RawShapeConstraint};
fn fixed(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
multi_array(name, shape, dtype, false, 2, vec![shape.to_vec()], shape)
}
fn multi_array(
name: &str,
shape: &[usize],
dtype: DataType,
optional: bool,
raw_type: isize,
enumerated: Vec<Vec<usize>>,
pinned: &[usize],
) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
optional,
Some(RawShapeConstraint::new(
raw_type,
enumerated,
pinned.iter().map(|d| AxisRange::new(*d, 1)).collect(),
)),
)
}
fn aligner_description() -> ModelDescription {
ModelDescription::from_parts(
vec![fixed(
names::WAVEFORM,
&[1, ENCODER_WINDOW_SAMPLES],
DataType::F32,
)],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
Vec::new(),
)
}
fn check(description: &ModelDescription) -> Result<(), AlignerError> {
check_with(AcousticGeometry::WAV2VEC2, description)
}
fn check_with(
geometry: AcousticGeometry,
description: &ModelDescription,
) -> Result<(), AlignerError> {
crate::model::contract::check_load_contract(
description,
&align_contract(geometry.receptive_field()),
)
.map_err(contract_violation)
}
fn check_staged(description: &ModelDescription) -> Result<Declared, AlignerError> {
check(description)?;
let declared = declared(description);
check_frame_count(
AcousticContract::BASE960H.geometry(),
declared.window,
declared.frames,
)?;
Ok(declared)
}
#[test]
fn the_contract_accepts_the_staged_geometry() {
let declared = check_staged(&aligner_description()).expect("the staged declaration loads");
assert_eq!(declared.window.get(), ENCODER_WINDOW_SAMPLES);
assert_eq!(declared.frames.get(), EXPECTED_OUTPUT_FRAMES);
assert_eq!(
declared.vocab_size.get(),
crate::audio::align::vocab::VOCAB_SIZE
);
}
#[test]
fn the_contract_refuses_a_flexible_waveform_declaring_its_exact_numbers() {
let description = ModelDescription::from_parts(
vec![multi_array(
names::WAVEFORM,
&[1, ENCODER_WINDOW_SAMPLES],
DataType::F32,
false,
3,
Vec::new(),
&[1, ENCODER_WINDOW_SAMPLES],
)],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m) if m.feature() == names::WAVEFORM),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_flexible_emissions_declaring_its_exact_numbers() {
let description = ModelDescription::from_parts(
vec![fixed(
names::WAVEFORM,
&[1, ENCODER_WINDOW_SAMPLES],
DataType::F32,
)],
vec![multi_array(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
false,
3,
Vec::new(),
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m) if m.feature() == names::EMISSIONS),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_missing_waveform() {
let description = ModelDescription::from_parts(
vec![fixed("audio", &[1, ENCODER_WINDOW_SAMPLES], DataType::F32)],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m)
if m.feature() == names::WAVEFORM && m.actual() == "missing"),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_wrong_dtype_or_an_empty_axis() {
const VOCAB: usize = crate::audio::align::vocab::VOCAB_SIZE;
let staged_emissions = || {
fixed(
names::EMISSIONS,
&[1, EXPECTED_OUTPUT_FRAMES, VOCAB],
DataType::F32,
)
};
let staged_waveform = || fixed(names::WAVEFORM, &[1, ENCODER_WINDOW_SAMPLES], DataType::F32);
let cases: [(FeatureInfo, FeatureInfo, &str); 5] = [
(
fixed(names::WAVEFORM, &[1, ENCODER_WINDOW_SAMPLES], DataType::F16),
staged_emissions(),
names::WAVEFORM,
),
(
staged_waveform(),
fixed(
names::EMISSIONS,
&[1, EXPECTED_OUTPUT_FRAMES, VOCAB],
DataType::F16,
),
names::EMISSIONS,
),
(
fixed(names::WAVEFORM, &[1, 0], DataType::F32),
staged_emissions(),
names::WAVEFORM,
),
(
staged_waveform(),
fixed(names::EMISSIONS, &[1, 0, VOCAB], DataType::F32),
names::EMISSIONS,
),
(
staged_waveform(),
fixed(
names::EMISSIONS,
&[1, EXPECTED_OUTPUT_FRAMES, 0],
DataType::F32,
),
names::EMISSIONS,
),
];
for (waveform, emissions, feature) in cases {
let description = ModelDescription::from_parts(vec![waveform], vec![emissions], Vec::new());
let err = check(&description).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m) if m.feature() == feature),
"expected a {feature} mismatch, got {err}"
);
}
}
#[test]
fn the_staged_geometry_refuses_a_window_or_frame_count_it_does_not_make() {
const VOCAB: usize = crate::audio::align::vocab::VOCAB_SIZE;
for (window, frames, derived) in [
(ENCODER_WINDOW_SAMPLES, 2_998, 2_999),
(ENCODER_WINDOW_SAMPLES, 3_000, 2_999),
(480_000, EXPECTED_OUTPUT_FRAMES, 1_499),
] {
let description = ModelDescription::from_parts(
vec![fixed(names::WAVEFORM, &[1, window], DataType::F32)],
vec![fixed(names::EMISSIONS, &[1, frames, VOCAB], DataType::F32)],
Vec::new(),
);
assert!(check(&description).is_ok(), "the contract fixes each axis");
let Err(AlignerError::FrameCountMismatch(mismatch)) = check_staged(&description) else {
panic!("{frames} frames of a {window}-sample window must be refused");
};
assert_eq!(
(mismatch.window(), mismatch.declared(), mismatch.derived()),
(window, frames, derived)
);
}
}
#[test]
fn the_contract_reads_the_window_frames_and_head_width_back() {
for (window, frames, width) in [
(ENCODER_WINDOW_SAMPLES, EXPECTED_OUTPUT_FRAMES, 29),
(ENCODER_WINDOW_SAMPLES, EXPECTED_OUTPUT_FRAMES, 32),
(ENCODER_WINDOW_SAMPLES, EXPECTED_OUTPUT_FRAMES, 1),
(ENCODER_WINDOW_SAMPLES, EXPECTED_OUTPUT_FRAMES, 64),
(480_000, 1_499, 29),
] {
let description = ModelDescription::from_parts(
vec![fixed(names::WAVEFORM, &[1, window], DataType::F32)],
vec![fixed(names::EMISSIONS, &[1, frames, width], DataType::F32)],
Vec::new(),
);
let declared = check_staged(&description).unwrap_or_else(|err| {
panic!("a {window}-sample window of {frames} frames and {width} classes loads: {err}")
});
assert_eq!(
(
declared.window.get(),
declared.frames.get(),
declared.vocab_size.get()
),
(window, frames, width)
);
}
}
#[test]
fn a_window_under_the_receptive_field_is_refused_at_load() {
const VOCAB: usize = crate::audio::align::vocab::VOCAB_SIZE;
let with_window = |window: usize| {
ModelDescription::from_parts(
vec![fixed(names::WAVEFORM, &[1, window], DataType::F32)],
vec![fixed(names::EMISSIONS, &[1, 1, VOCAB], DataType::F32)],
Vec::new(),
)
};
let err = check(&with_window(100)).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m) if m.feature() == names::WAVEFORM),
"{err}"
);
assert!(check(&with_window(399)).is_err());
assert!(check(&with_window(400)).is_ok());
let wide = geometry(640, 320);
assert!(check_with(wide, &with_window(400)).is_err());
assert!(check_with(wide, &with_window(639)).is_err());
assert!(check_with(wide, &with_window(640)).is_ok());
let narrow = geometry(200, 100);
assert!(check_with(narrow, &with_window(199)).is_err());
assert!(check_with(narrow, &with_window(300)).is_ok());
}
#[test]
fn a_log_probability_head_too_wide_to_check_is_refused_at_load() {
let width = |v: usize| NonZeroUsize::new(v).expect("nonzero");
let v = 5_000usize;
let raw_row = vec![-4.0f32; v];
let lse = -4.0 + (v as f64).ln();
assert!(
lse < log_prob_sum_tolerance(width(v)),
"the reason: at 5,000 classes the allowance ({}) admits the raw -4 row (logsumexp {lse})",
log_prob_sum_tolerance(width(v))
);
let Err(AlignerError::UnprovableNormalization(refused)) =
check_output_width(OutputKind::LogProbabilities, width(v))
else {
panic!("a 5,000-class log-probability head must be refused at load");
};
assert_eq!(
(refused.vocab_size(), refused.widest()),
(v, MAX_LOG_PROB_WIDTH)
);
assert_eq!(check_output_width(OutputKind::Logits, width(v)), Ok(()));
assert_eq!(MAX_LOG_PROB_WIDTH, 353);
assert_eq!(
check_output_width(OutputKind::LogProbabilities, width(353)),
Ok(())
);
assert!(check_output_width(OutputKind::LogProbabilities, width(354)).is_err());
let emissions = through_asry(
RawEmissions {
frames: 1,
vocab_size: width(v),
data: raw_row,
}
.into_logit_output(),
)
.expect("logits are normalized, not checked");
assert_eq!(emissions.vocab(), width(v));
}
#[test]
fn the_widest_checkable_head_still_refuses_a_frame_off_by_a_factor_of_two() {
let v = MAX_LOG_PROB_WIDTH;
let width = NonZeroUsize::new(v).expect("nonzero");
let ln_v = (v as f64).ln();
let doubled = vec![(UNNORMALIZED_LOGSUMEXP - ln_v) as f32; v];
assert!(matches!(
check_log_prob_normalization(&doubled, width, ComputeUnits::CpuOnly),
Err(AlignError::UnnormalizedEmissions(_))
));
let within = vec![((log_prob_sum_tolerance(width) / 2.0) - ln_v) as f32; v];
assert!(check_log_prob_normalization(&within, width, ComputeUnits::CpuOnly).is_ok());
}
#[test]
fn the_contract_refuses_an_extra_required_input() {
let description = ModelDescription::from_parts(
vec![
fixed(names::WAVEFORM, &[1, ENCODER_WINDOW_SAMPLES], DataType::F32),
fixed(
"attention_mask",
&[1, ENCODER_WINDOW_SAMPLES],
DataType::I32,
),
],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
Vec::new(),
);
assert!(
matches!(check(&description), Err(AlignerError::UnsatisfiableInput(ref name))
if name == "attention_mask"),
"{:?}",
check(&description)
);
}
#[test]
fn the_contract_accepts_an_extra_optional_input() {
let description = ModelDescription::from_parts(
vec![
fixed(names::WAVEFORM, &[1, ENCODER_WINDOW_SAMPLES], DataType::F32),
multi_array(
"attention_mask",
&[1, ENCODER_WINDOW_SAMPLES],
DataType::I32,
true,
2,
vec![vec![1, ENCODER_WINDOW_SAMPLES]],
&[1, ENCODER_WINDOW_SAMPLES],
),
],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
Vec::new(),
);
assert!(check(&description).is_ok());
}
#[test]
fn the_contract_refuses_an_optional_emissions_output() {
let description = ModelDescription::from_parts(
vec![fixed(
names::WAVEFORM,
&[1, ENCODER_WINDOW_SAMPLES],
DataType::F32,
)],
vec![multi_array(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
true,
2,
vec![vec![
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
]],
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, AlignerError::ContractMismatch(m) if m.feature() == names::EMISSIONS),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_graph_that_declares_state() {
let description = ModelDescription::from_parts(
vec![fixed(
names::WAVEFORM,
&[1, ENCODER_WINDOW_SAMPLES],
DataType::F32,
)],
vec![fixed(
names::EMISSIONS,
&[
1,
EXPECTED_OUTPUT_FRAMES,
crate::audio::align::vocab::VOCAB_SIZE,
],
DataType::F32,
)],
vec![fixed("kv_cache", &[1, 8], DataType::F32)],
);
assert!(
matches!(check(&description), Err(AlignerError::UnsatisfiableState(ref name))
if name == "kv_cache")
);
}
#[test]
fn the_align_contract_refuses_the_vendored_silero_bundle() {
let bundle = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../Models/vadkit/silero-vad-unified-256ms-v6.2.1.mlmodelc");
assert!(
bundle.is_dir(),
"the vendored silero bundle is committed, so this gate is NOT model-gated; \
looked for {}",
bundle.display()
);
let model = Model::load(&bundle, ComputeUnits::CpuOnly).expect("the committed bundle loads");
assert!(
model.description().input(names::WAVEFORM).is_none(),
"silero declares no `waveform`, which is what makes it this gate's model"
);
let violation = Checked::new(
model,
&align_contract(AcousticGeometry::WAV2VEC2.receptive_field()),
)
.expect_err("silero does not satisfy the aligner contract");
assert!(
matches!(&violation, ContractViolation::Missing(m) if m.feature() == names::WAVEFORM),
"expected `waveform` missing, got {violation}"
);
}