use super::*;
#[test]
fn names_match_recorded_ground_truth() {
assert_eq!(names::LOGITS, "logits");
assert_eq!(names::KEY_UPDATES, "key_cache_updates");
assert_eq!(names::VALUE_UPDATES, "value_cache_updates");
assert_eq!(names::ALIGNMENT, "alignment_heads_weights");
assert_eq!(names::KV_UPDATE_MASK, "kv_cache_update_mask");
}
use crate::{
AxisRange, FeatureInfo, audio::whisper::constants::MAX_TOKEN_CONTEXT, model::RawShapeConstraint,
};
fn fixed(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(
2,
vec![shape.to_vec()],
shape.iter().map(|d| AxisRange::new(*d, 1)).collect(),
)),
)
}
fn ranged(name: &str, shape: &[usize], dtype: DataType, ranges: &[AxisRange]) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
false,
Some(RawShapeConstraint::new(3, Vec::new(), ranges.to_vec())),
)
}
const TINY_WINDOW: usize = 480_000;
const TINY_MELS: usize = 80;
const TINY_MEL_FRAMES: usize = 3_000;
const TINY_EMBED: usize = 384;
const TINY_AUDIO_CTX: usize = 1_500;
const TINY_KV: usize = 1_536;
const TINY_CTX: usize = 224;
const TINY_VOCAB: usize = 51_865;
fn tiny_mel_description() -> ModelDescription {
ModelDescription::from_parts(
vec![fixed(names::AUDIO, &[TINY_WINDOW], DataType::F16)],
vec![fixed(
names::MEL,
&[1, TINY_MELS, 1, TINY_MEL_FRAMES],
DataType::F16,
)],
Vec::new(),
)
}
fn tiny_encoder_description() -> ModelDescription {
ModelDescription::from_parts(
vec![fixed(
names::MEL,
&[1, TINY_MELS, 1, TINY_MEL_FRAMES],
DataType::F16,
)],
vec![fixed(
names::ENCODER,
&[1, TINY_EMBED, 1, TINY_AUDIO_CTX],
DataType::F16,
)],
Vec::new(),
)
}
fn tiny_decoder_description_with(
mutate: impl FnOnce(&mut Vec<FeatureInfo>, &mut Vec<FeatureInfo>),
) -> ModelDescription {
let mut inputs = vec![
fixed(names::INPUT_IDS, &[1], DataType::I32),
fixed(names::CACHE_LENGTH, &[1], DataType::I32),
fixed(names::KEY_CACHE, &[1, TINY_KV, 1, TINY_CTX], DataType::F16),
fixed(
names::VALUE_CACHE,
&[1, TINY_KV, 1, TINY_CTX],
DataType::F16,
),
fixed(names::KV_UPDATE_MASK, &[1, TINY_CTX], DataType::F16),
fixed(
names::ENCODER,
&[1, TINY_EMBED, 1, TINY_AUDIO_CTX],
DataType::F16,
),
fixed(names::PADDING_MASK, &[1, TINY_CTX], DataType::F16),
];
let mut outputs = vec![
fixed(names::LOGITS, &[1, 1, TINY_VOCAB], DataType::F16),
fixed(names::KEY_UPDATES, &[1, TINY_KV, 1, 1], DataType::F16),
fixed(names::VALUE_UPDATES, &[1, TINY_KV, 1, 1], DataType::F16),
fixed(names::ALIGNMENT, &[1, TINY_AUDIO_CTX], DataType::F16),
];
mutate(&mut inputs, &mut outputs);
ModelDescription::from_parts(inputs, outputs, Vec::new())
}
fn tiny_decoder_description() -> ModelDescription {
tiny_decoder_description_with(|_, _| {})
}
fn tiny_decoder_contract(supports_alignment: bool) -> LoadContract {
decoder_contract(
TINY_EMBED,
TINY_AUDIO_CTX,
TINY_KV,
TINY_CTX,
TINY_VOCAB,
supports_alignment,
)
}
fn check(description: &ModelDescription, contract: &LoadContract) -> Result<(), BackendError> {
crate::model::contract::check_load_contract(description, contract)
.map_err(contract_violation("decoder"))
}
#[test]
fn the_three_contracts_accept_the_staged_tiny_descriptions() {
assert!(
crate::model::contract::check_load_contract(&tiny_mel_description(), &mel_contract()).is_ok()
);
assert!(
crate::model::contract::check_load_contract(
&tiny_encoder_description(),
&encoder_contract(TINY_MELS, TINY_MEL_FRAMES),
)
.is_ok()
);
assert_eq!(
check(&tiny_decoder_description(), &tiny_decoder_contract(true)),
Ok(())
);
}
#[test]
fn the_decoder_contract_refuses_a_mistyped_kv_cache_update_mask() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[4] = fixed(names::KV_UPDATE_MASK, &[1, TINY_CTX], DataType::I32);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.model() == "decoder"
&& c.feature() == names::KV_UPDATE_MASK
&& c.expected() == "float16"
&& c.actual() == "int32"),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_mistyped_padding_mask() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[6] = fixed(names::PADDING_MASK, &[1, TINY_CTX], DataType::F32);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == names::PADDING_MASK),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_decoder_that_declares_state() {
let base = tiny_decoder_description();
let description = ModelDescription::from_parts(
base.inputs().to_vec(),
base.outputs().to_vec(),
vec![fixed("kv_cache", &[1, TINY_KV], DataType::F16)],
);
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.feature() == "kv_cache" && c.actual() == "a declared state buffer"),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_value_cache_that_disagrees_with_the_key_cache() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[3] = fixed(
names::VALUE_CACHE,
&[1, TINY_KV * 2, 1, TINY_CTX],
DataType::F16,
);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == names::VALUE_CACHE),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_mask_narrower_than_the_kv_cache() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[4] = fixed(names::KV_UPDATE_MASK, &[1, TINY_CTX / 2], DataType::F16);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == names::KV_UPDATE_MASK),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_an_encoder_output_from_another_model_size() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[5] = fixed(names::ENCODER, &[1, 768, 1, TINY_AUDIO_CTX], DataType::F16);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == names::ENCODER),
"{err}"
);
}
#[test]
fn the_encoder_contract_refuses_an_input_that_is_not_the_mels_output() {
let description = ModelDescription::from_parts(
vec![fixed(
names::MEL,
&[1, 128, 1, TINY_MEL_FRAMES],
DataType::F16,
)],
vec![fixed(
names::ENCODER,
&[1, 1280, 1, TINY_AUDIO_CTX],
DataType::F16,
)],
Vec::new(),
);
let violation = crate::model::contract::check_load_contract(
&description,
&encoder_contract(TINY_MELS, TINY_MEL_FRAMES),
)
.unwrap_err();
assert!(
matches!(&violation, ContractViolation::Axis(a) if a.feature() == names::MEL),
"{violation}"
);
}
#[test]
fn the_decoder_contract_refuses_a_transposed_logits_head() {
let description = tiny_decoder_description_with(|_, outputs| {
outputs[0] = fixed(names::LOGITS, &[1, TINY_VOCAB, 1, 1], DataType::F16);
});
assert_eq!(
[1, TINY_VOCAB, 1, 1].iter().product::<usize>(),
[1, 1, TINY_VOCAB].iter().product::<usize>()
);
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == names::LOGITS),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_flexible_key_cache() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs[2] = ranged(
names::KEY_CACHE,
&[1, TINY_KV, 1, TINY_CTX],
DataType::F16,
&[
AxisRange::new(1, 1),
AxisRange::new(TINY_KV, 1),
AxisRange::new(1, 1),
AxisRange::inclusive(1, 448),
],
);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.feature() == names::KEY_CACHE && c.expected() == "fixed" && c.actual() == "range"),
"{err}"
);
}
#[test]
fn the_decoder_contract_refuses_a_required_input_it_does_not_name() {
let description = tiny_decoder_description_with(|inputs, _| {
inputs.push(fixed("prompt_ids", &[1, 16], DataType::I32));
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c) if c.feature() == "prompt_ids"),
"{err}"
);
}
#[test]
fn a_decoder_without_the_alignment_head_is_accepted_when_the_contract_omits_it() {
let description = tiny_decoder_description_with(|_, outputs| {
outputs.pop();
});
assert_eq!(check(&description, &tiny_decoder_contract(false)), Ok(()));
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.feature() == names::ALIGNMENT && c.actual() == "missing"),
"{err}"
);
}
#[test]
fn the_backend_contracts_refuse_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 load = || Model::load(&bundle, crate::ComputeUnits::CpuOnly).expect("committed bundle loads");
let err = CoreMlBackend::new(load(), load(), load(), TINY_VOCAB)
.expect_err("silero is not a whisper model");
assert!(
matches!(&err, BackendError::Contract(c)
if c.model() == "mel" && c.feature() == names::AUDIO && c.actual() == "missing"),
"{err}"
);
}
fn with_axis_bumped(base: &ModelDescription, feature: &str, axis: usize) -> ModelDescription {
let bump = |declared: &FeatureInfo| -> FeatureInfo {
if declared.name() != feature {
return declared.clone();
}
let mut shape = declared.shape().to_vec();
shape[axis] += 1;
fixed(
declared.name(),
&shape,
declared.data_type().expect("a multi-array feature"),
)
};
ModelDescription::from_parts(
base.inputs().iter().map(bump).collect(),
base.outputs().iter().map(bump).collect(),
base.states().to_vec(),
)
}
fn with_dtype_changed(base: &ModelDescription, feature: &str, dtype: DataType) -> ModelDescription {
let swap = |declared: &FeatureInfo| -> FeatureInfo {
if declared.name() != feature {
return declared.clone();
}
fixed(declared.name(), declared.shape(), dtype)
};
ModelDescription::from_parts(
base.inputs().iter().map(swap).collect(),
base.outputs().iter().map(swap).collect(),
base.states().to_vec(),
)
}
struct Stage {
name: &'static str,
contract: LoadContract,
description: ModelDescription,
free_axes: &'static [(&'static str, usize)],
}
fn contract_stages() -> [Stage; 3] {
[
Stage {
name: "mel",
contract: mel_contract(),
description: tiny_mel_description(),
free_axes: &[(names::AUDIO, 0), (names::MEL, 1), (names::MEL, 3)],
},
Stage {
name: "encoder",
contract: encoder_contract(TINY_MELS, TINY_MEL_FRAMES),
description: tiny_encoder_description(),
free_axes: &[(names::ENCODER, 1), (names::ENCODER, 3)],
},
Stage {
name: "decoder",
contract: tiny_decoder_contract(true),
description: tiny_decoder_description(),
free_axes: &[(names::KEY_CACHE, 1), (names::KEY_CACHE, 3)],
},
]
}
#[test]
fn every_axis_is_pinned_except_the_dimensions_each_stage_reads_back() {
let mut perturbations = 0_usize;
for stage in contract_stages() {
let base = &stage.description;
for declared in base.inputs().iter().chain(base.outputs()) {
for axis in 0..declared.shape().len() {
let perturbed = with_axis_bumped(base, declared.name(), axis);
let free = stage.free_axes.contains(&(declared.name(), axis));
let accepted =
crate::model::contract::check_load_contract(&perturbed, &stage.contract).is_ok();
assert_eq!(
accepted,
free,
"{}: `{}` axis {axis}: the contract {} it",
stage.name,
declared.name(),
if free { "must accept" } else { "must refuse" }
);
perturbations += 1;
}
}
}
assert_eq!(perturbations, 44);
}
#[test]
fn every_named_features_element_type_is_pinned() {
let mut checked = 0_usize;
for stage in contract_stages() {
let base = &stage.description;
let names: Vec<String> = base
.inputs()
.iter()
.chain(base.outputs())
.map(|f| f.name().to_string())
.collect();
for name in names {
let declared = base
.input(&name)
.or_else(|| base.output(&name))
.expect("just enumerated");
let other = if declared.data_type() == Some(DataType::I32) {
DataType::F16
} else {
DataType::I32
};
let perturbed = with_dtype_changed(base, &name, other);
assert!(
crate::model::contract::check_load_contract(&perturbed, &stage.contract).is_err(),
"{}: `{name}` re-declared {other:?} must be refused",
stage.name
);
checked += 1;
}
}
assert_eq!(checked, 15);
}
fn with_axis_set(
base: &ModelDescription,
feature: &str,
axis: usize,
size: usize,
) -> ModelDescription {
let set = |declared: &FeatureInfo| -> FeatureInfo {
if declared.name() != feature {
return declared.clone();
}
let mut shape = declared.shape().to_vec();
shape[axis] = size;
fixed(
declared.name(),
&shape,
declared.data_type().expect("a multi-array feature"),
)
};
ModelDescription::from_parts(
base.inputs().iter().map(set).collect(),
base.outputs().iter().map(set).collect(),
base.states().to_vec(),
)
}
#[test]
fn every_read_back_axis_refuses_a_zero_size() {
let mut floors = 0_usize;
for stage in contract_stages() {
for (feature, axis) in stage.free_axes {
let zeroed = with_axis_set(&stage.description, feature, *axis, 0);
assert!(
crate::model::contract::check_load_contract(&zeroed, &stage.contract).is_err(),
"{}: `{feature}` axis {axis} at size 0 must be refused",
stage.name
);
floors += 1;
}
}
assert_eq!(floors, 7);
}
#[test]
fn the_decoder_contract_refuses_a_vocabulary_the_tokenizer_overruns() {
const SHORT_VOCAB: usize = 50_363; let description = tiny_decoder_description_with(|_, outputs| {
outputs[0] = fixed(names::LOGITS, &[1, 1, SHORT_VOCAB], DataType::F16);
});
let err = check(&description, &tiny_decoder_contract(true)).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.feature() == names::LOGITS
&& c.expected().contains(&TINY_VOCAB.to_string())
&& c.actual().contains(&SHORT_VOCAB.to_string())),
"{err}"
);
}
#[test]
fn the_audio_boundary_refuses_a_window_no_f16_input_can_carry() {
const WINDOW: usize = 8;
let clean = vec![0.5_f32; WINDOW];
assert!(audio_input(&clean, WINDOW).is_ok());
assert!(matches!(
audio_input(&clean, WINDOW + 1),
Err(BackendError::AudioLength(ref length))
if length.got() == WINDOW && length.expected() == WINDOW + 1
));
for (index, bad) in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY]
.into_iter()
.enumerate()
{
let mut window = clean.clone();
window[index] = bad;
assert!(
matches!(
audio_input(&window, WINDOW),
Err(BackendError::NonFiniteAudio(at)) if at == index
),
"{bad} at {index}"
);
}
let mut window = clean.clone();
window[5] = 70_000.0;
assert!(matches!(
audio_input(&window, WINDOW),
Err(BackendError::F16OverflowAudio(5))
));
for value in [f32::from(f16::MAX), -f32::from(f16::MAX)] {
let mut window = clean.clone();
window[5] = value;
assert!(audio_input(&window, WINDOW).is_ok(), "{value}");
}
let mut window = clean.clone();
window[2] = 70_000.0;
window[6] = f32::NAN;
assert!(matches!(
audio_input(&window, WINDOW),
Err(BackendError::F16OverflowAudio(2))
));
}
fn decoder_description_at_context(c: usize) -> ModelDescription {
tiny_decoder_description_with(|inputs, _| {
inputs[2] = fixed(names::KEY_CACHE, &[1, TINY_KV, 1, c], DataType::F16);
inputs[3] = fixed(names::VALUE_CACHE, &[1, TINY_KV, 1, c], DataType::F16);
inputs[4] = fixed(names::KV_UPDATE_MASK, &[1, c], DataType::F16);
inputs[6] = fixed(names::PADDING_MASK, &[1, c], DataType::F16);
})
}
#[test]
fn the_decoder_contract_refuses_a_kv_context_the_decode_loop_overruns() {
const SHORT_CTX: usize = 100;
let description = decoder_description_at_context(SHORT_CTX);
let contract = decoder_contract(
TINY_EMBED,
TINY_AUDIO_CTX,
TINY_KV,
SHORT_CTX,
TINY_VOCAB,
true,
);
let err = check(&description, &contract).unwrap_err();
assert!(
matches!(&err, BackendError::Contract(c)
if c.feature() == names::KEY_CACHE
&& c.expected().contains(&MAX_TOKEN_CONTEXT.to_string())
&& c.actual().contains(&SHORT_CTX.to_string())),
"{err}"
);
}
#[test]
fn the_decoder_contract_accepts_a_kv_context_larger_than_the_loop_uses() {
const LARGE_V3_CTX: usize = 448;
const { assert!(LARGE_V3_CTX > MAX_TOKEN_CONTEXT) };
let description = decoder_description_at_context(LARGE_V3_CTX);
let contract = decoder_contract(
TINY_EMBED,
TINY_AUDIO_CTX,
TINY_KV,
LARGE_V3_CTX,
TINY_VOCAB,
true,
);
assert_eq!(check(&description, &contract), Ok(()));
}
#[test]
fn the_kv_context_floor_is_max_token_context_exactly() {
for (c, accepted) in [
(MAX_TOKEN_CONTEXT - 1, false),
(MAX_TOKEN_CONTEXT, true),
(MAX_TOKEN_CONTEXT + 1, true),
] {
let contract = decoder_contract(TINY_EMBED, TINY_AUDIO_CTX, TINY_KV, c, TINY_VOCAB, true);
assert_eq!(
check(&decoder_description_at_context(c), &contract).is_ok(),
accepted,
"a KV context of {c}"
);
}
}