use super::*;
use crate::model::RawShapeConstraint;
fn pinned(shape: &[usize]) -> Vec<AxisRange> {
shape.iter().map(|d| AxisRange::new(*d, 1)).collect()
}
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()],
pinned(shape),
)),
)
}
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())),
)
}
fn optional(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
true,
Some(RawShapeConstraint::new(
2,
vec![shape.to_vec()],
pinned(shape),
)),
)
}
const MEL: &str = "mel";
const EMBEDDING: &str = "embedding";
const MEL_SHAPE: &[usize] = &[1, 72, 401];
const EMBEDDING_SHAPE: &[usize] = &[1, 192];
fn identity_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
MEL,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(72), Dim::Exactly(401)],
)],
vec![FeatureContract::new(
EMBEDDING,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(192)],
)],
StateContract::None,
)
}
fn redimnet_description() -> ModelDescription {
ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
)
}
#[test]
fn the_identity_contract_accepts_the_published_redimnet_description() {
assert_eq!(
check_load_contract(&redimnet_description(), &identity_contract()),
Ok(())
);
}
#[test]
fn a_named_feature_the_model_does_not_declare_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed("audio", MEL_SHAPE, DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Missing(m) if m.feature() == MEL),
"{error}"
);
assert!(error.to_string().contains("declares no feature `mel`"));
}
#[test]
fn a_missing_output_is_refused_too() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![fixed("logits", EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Missing(m) if m.feature() == EMBEDDING),
"{error}"
);
}
#[test]
fn a_feature_of_the_wrong_element_type_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F16)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::DataType(d) if d.feature() == MEL),
"{error}"
);
assert!(error.to_string().contains("float16"), "{error}");
}
#[test]
fn a_feature_with_a_different_number_of_axes_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[72, 401], DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Rank(r) if r.feature() == MEL),
"{error}"
);
assert!(error.to_string().contains("rank 2"), "{error}");
}
#[test]
fn an_all_fixed_contract_refuses_a_flexible_feature_declaring_its_numbers() {
let description = ModelDescription::from_parts(
vec![ranged(MEL, MEL_SHAPE, DataType::F32, &pinned(MEL_SHAPE))],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == MEL),
"{error}"
);
assert!(error.to_string().contains("is range"), "{error}");
}
#[test]
fn an_all_fixed_contract_refuses_an_enumerated_feature_whose_ranges_look_pinned() {
let mel = FeatureInfo::from_parts(
MEL.to_string(),
MEL_SHAPE.to_vec(),
Some(DataType::F32),
false,
Some(RawShapeConstraint::new(
2,
vec![MEL_SHAPE.to_vec(), vec![1, 72, 201]],
pinned(MEL_SHAPE),
)),
);
let description = ModelDescription::from_parts(
vec![mel],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == MEL),
"{error}"
);
assert!(error.to_string().contains("is enumerated"), "{error}");
}
#[test]
fn an_all_fixed_contract_refuses_an_unspecified_feature() {
let embedding = FeatureInfo::from_parts(
EMBEDDING.to_string(),
EMBEDDING_SHAPE.to_vec(),
Some(DataType::F32),
false,
Some(RawShapeConstraint::new(1, Vec::new(), Vec::new())),
);
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![embedding],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == EMBEDDING),
"{error}"
);
assert!(error.to_string().contains("is unspecified"), "{error}");
}
#[test]
fn an_exactly_axis_pinned_at_another_size_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72, 400], DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Axis(a) if a.feature() == MEL),
"{error}"
);
let rendered = error.to_string();
assert!(rendered.contains("axis 2 400"), "{rendered}");
assert!(rendered.contains("axis 2 401"), "{rendered}");
}
#[test]
fn a_transposed_shape_of_the_same_size_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[1, 401, 72], DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
assert!(matches!(
check_load_contract(&description, &identity_contract()),
Err(ContractViolation::Axis(_))
));
}
fn any_fixed_batch_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
MEL,
DataType::F32,
vec![Dim::AnyFixed, Dim::Exactly(72), Dim::Exactly(401)],
)],
Vec::new(),
StateContract::None,
)
}
#[test]
fn an_any_fixed_axis_accepts_whatever_one_size_is_pinned_and_is_read_back() {
for batch in [1_usize, 3, 32] {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[batch, 72, 401], DataType::F32)],
Vec::new(),
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &any_fixed_batch_contract()),
Ok(())
);
assert_eq!(description.input(MEL).expect("mel").shape()[0], batch);
}
}
#[test]
fn an_any_fixed_axis_refuses_an_axis_admitting_more_than_one_size() {
let description = ModelDescription::from_parts(
vec![ranged(
MEL,
&[3, 72, 401],
DataType::F32,
&[
AxisRange::inclusive(1, 8),
AxisRange::new(72, 1),
AxisRange::new(401, 1),
],
)],
Vec::new(),
Vec::new(),
);
let error = check_load_contract(&description, &any_fixed_batch_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == MEL),
"{error}"
);
}
#[test]
fn an_any_fixed_axis_the_model_pins_at_zero_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[0, 72, 401], DataType::F32)],
Vec::new(),
Vec::new(),
);
assert_eq!(
description.input(MEL).expect("mel").axis_ranges()[0],
AxisRange::new(0, 1)
);
let error = check_load_contract(&description, &any_fixed_batch_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::ZeroSizedAxis(z) if z.feature() == MEL),
"{error}"
);
let rendered = error.to_string();
assert!(rendered.contains("axis 0 0"), "{rendered}");
assert!(
rendered.contains("axis 0 any one non-zero fixed size"),
"{rendered}"
);
}
const INPUT_IDS: &str = "input_ids";
fn any_fixed_window_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
INPUT_IDS,
DataType::I32,
vec![Dim::Exactly(1), Dim::AnyFixed],
)],
Vec::new(),
StateContract::None,
)
}
#[test]
fn an_any_fixed_axis_pinned_at_zero_is_refused_on_a_trailing_axis() {
let description = ModelDescription::from_parts(
vec![fixed(INPUT_IDS, &[1, 0], DataType::I32)],
Vec::new(),
Vec::new(),
);
let declared = description.input(INPUT_IDS).expect("input_ids");
assert_eq!(declared.axis_ranges()[0], AxisRange::new(1, 1));
assert_eq!(declared.axis_ranges()[1], AxisRange::new(0, 1));
let error = check_load_contract(&description, &any_fixed_window_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::ZeroSizedAxis(z) if z.feature() == INPUT_IDS),
"{error}"
);
let rendered = error.to_string();
assert!(rendered.contains("axis 1 0"), "{rendered}");
assert!(
rendered.contains("axis 1 any one non-zero fixed size"),
"{rendered}"
);
}
#[test]
fn an_any_fixed_trailing_axis_accepts_every_non_zero_window() {
for window in [1_usize, 16, 64, 512] {
let description = ModelDescription::from_parts(
vec![fixed(INPUT_IDS, &[1, window], DataType::I32)],
Vec::new(),
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &any_fixed_window_contract()),
Ok(()),
"window {window}"
);
}
}
#[test]
fn a_zero_on_an_exactly_axis_stays_an_ordinary_axis_mismatch() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[3, 0, 401], DataType::F32)],
Vec::new(),
Vec::new(),
);
assert!(matches!(
check_load_contract(&description, &any_fixed_batch_contract()),
Err(ContractViolation::Axis(_))
));
}
#[test]
fn an_any_fixed_axis_accepts_every_non_zero_size() {
for batch in [1_usize, 2, 3, 32, 4096] {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[batch, 72, 401], DataType::F32)],
Vec::new(),
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &any_fixed_batch_contract()),
Ok(()),
"batch {batch}"
);
}
}
fn lid_shaped_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
"mel_features",
DataType::F32,
vec![
Dim::Exactly(1),
Dim::Range(AxisRange::inclusive(10, 3001)),
Dim::Exactly(60),
],
)],
Vec::new(),
StateContract::None,
)
}
fn lid_shaped_description(ranges: &[AxisRange]) -> ModelDescription {
ModelDescription::from_parts(
vec![ranged("mel_features", &[1, 301, 60], DataType::F32, ranges)],
Vec::new(),
Vec::new(),
)
}
#[test]
fn a_lid_shaped_flexible_contract_is_expressible() {
let description = lid_shaped_description(&[
AxisRange::new(1, 1),
AxisRange::inclusive(10, 3001),
AxisRange::new(60, 1),
]);
assert_eq!(
check_load_contract(&description, &lid_shaped_contract()),
Ok(())
);
}
#[test]
fn a_flexible_axis_whose_bounds_differ_is_refused() {
let description = lid_shaped_description(&[
AxisRange::new(1, 1),
AxisRange::inclusive(10, 1500),
AxisRange::new(60, 1),
]);
let error = check_load_contract(&description, &lid_shaped_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Axis(a) if a.feature() == "mel_features"),
"{error}"
);
let rendered = error.to_string();
assert!(rendered.contains("axis 1 10..=1500"), "{rendered}");
assert!(rendered.contains("axis 1 10..=3001"), "{rendered}");
}
#[test]
fn a_fixed_axis_inside_a_flexible_feature_is_still_checked() {
let description = lid_shaped_description(&[
AxisRange::new(1, 1),
AxisRange::inclusive(10, 3001),
AxisRange::new(80, 1),
]);
assert!(matches!(
check_load_contract(&description, &lid_shaped_contract()),
Err(ContractViolation::Axis(_))
));
}
#[test]
fn a_flexible_contract_refuses_a_fully_fixed_feature() {
let description = ModelDescription::from_parts(
vec![fixed("mel_features", &[1, 301, 60], DataType::F32)],
Vec::new(),
Vec::new(),
);
let error = check_load_contract(&description, &lid_shaped_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == "mel_features"),
"{error}"
);
}
#[test]
fn an_axis_the_constraint_lists_no_range_for_is_refused() {
let description = lid_shaped_description(&[AxisRange::new(1, 1), AxisRange::inclusive(10, 3001)]);
let error = check_load_contract(&description, &lid_shaped_contract()).unwrap_err();
assert!(matches!(&error, ContractViolation::Axis(_)), "{error}");
assert!(error.to_string().contains("axis 2 none"), "{error}");
}
#[test]
fn a_required_input_the_contract_does_not_name_is_refused() {
let description = ModelDescription::from_parts(
vec![
fixed(MEL, MEL_SHAPE, DataType::F32),
fixed("speaker_mask", &[1, 401], DataType::F32),
],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::UnsatisfiableInput(i) if i.name() == "speaker_mask"),
"{error}"
);
}
#[test]
fn an_optional_extra_input_is_accepted() {
let description = ModelDescription::from_parts(
vec![
fixed(MEL, MEL_SHAPE, DataType::F32),
optional("mask", &[1, 401], DataType::F32),
],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &identity_contract()),
Ok(())
);
}
#[test]
fn an_extra_output_is_accepted() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![
fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32),
fixed("logits", &[1, 5994], DataType::F32),
],
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &identity_contract()),
Ok(())
);
}
#[test]
fn a_named_output_the_model_declares_optional_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![optional(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::OptionalOutput(o) if o.feature() == EMBEDDING),
"{error}"
);
assert!(error.to_string().contains("`embedding`"), "{error}");
assert_eq!(
check_load_contract(&redimnet_description(), &identity_contract()),
Ok(())
);
}
#[test]
fn a_named_input_the_model_declares_optional_is_accepted() {
let description = ModelDescription::from_parts(
vec![optional(MEL, MEL_SHAPE, DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &identity_contract()),
Ok(())
);
}
#[test]
fn the_reported_unsatisfiable_input_is_stable() {
let description = ModelDescription::from_parts(
vec![
fixed("aaa", &[1], DataType::F32),
fixed(MEL, MEL_SHAPE, DataType::F32),
fixed("zzz", &[1], DataType::F32),
],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::UnsatisfiableInput(i) if i.name() == "aaa"),
"{error}"
);
}
#[test]
fn a_declared_state_buffer_is_refused() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
vec![fixed("kv_cache", &[1, 8], DataType::F32)],
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::UnsatisfiableState(s) if s.name() == "kv_cache"),
"{error}"
);
assert!(
error.to_string().contains("state buffer `kv_cache`"),
"{error}"
);
}
#[test]
fn the_reported_state_buffer_is_stable() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F32)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
vec![
fixed("aaa", &[1], DataType::F32),
fixed("zzz", &[1], DataType::F32),
],
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::UnsatisfiableState(s) if s.name() == "aaa"),
"{error}"
);
}
fn at_least_frames_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
MEL,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(72), Dim::AtLeast(224)],
)],
Vec::new(),
StateContract::None,
)
}
#[test]
fn an_at_least_axis_accepts_any_pinned_size_from_the_floor_up_and_is_read_back() {
for frames in [224_usize, 225, 448, 4096] {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72, frames], DataType::F32)],
Vec::new(),
Vec::new(),
);
assert_eq!(
check_load_contract(&description, &at_least_frames_contract()),
Ok(()),
"{frames} frames"
);
assert_eq!(description.input(MEL).expect("mel").shape()[2], frames);
}
}
#[test]
fn an_at_least_axis_refuses_every_size_below_its_floor() {
for frames in [0_usize, 1, 100, 223] {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72, frames], DataType::F32)],
Vec::new(),
Vec::new(),
);
let error = check_load_contract(&description, &at_least_frames_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Axis(a) if a.feature() == MEL),
"{frames} frames: {error}"
);
let rendered = error.to_string();
assert!(rendered.contains("at least 224"), "{rendered}");
assert!(rendered.contains(&format!("axis 2 {frames}")), "{rendered}");
}
}
#[test]
fn a_zero_on_an_at_least_axis_is_an_ordinary_axis_mismatch() {
let description = ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72, 0], DataType::F32)],
Vec::new(),
Vec::new(),
);
let error = check_load_contract(&description, &at_least_frames_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Axis(a) if a.feature() == MEL),
"{error}"
);
}
#[test]
fn an_at_least_axis_still_requires_the_axis_to_be_pinned() {
let description = ModelDescription::from_parts(
vec![ranged(
MEL,
&[1, 72, 448],
DataType::F32,
&[
AxisRange::new(1, 1),
AxisRange::new(72, 1),
AxisRange::inclusive(224, 448),
],
)],
Vec::new(),
Vec::new(),
);
let error = check_load_contract(&description, &at_least_frames_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::Flexibility(f) if f.feature() == MEL),
"{error}"
);
}
fn one_violation_per_clause() -> Vec<ContractViolation> {
let refuse = |description: ModelDescription| {
check_load_contract(&description, &identity_contract())
.expect_err("each description below fails exactly one clause")
};
let embedding = || fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32);
let mel = || fixed(MEL, MEL_SHAPE, DataType::F32);
vec![
refuse(ModelDescription::from_parts(
Vec::new(),
vec![embedding()],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![fixed(MEL, MEL_SHAPE, DataType::F16)],
vec![embedding()],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72], DataType::F32)],
vec![embedding()],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![ranged(MEL, MEL_SHAPE, DataType::F32, &pinned(MEL_SHAPE))],
vec![embedding()],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![fixed(MEL, &[1, 72, 400], DataType::F32)],
vec![embedding()],
Vec::new(),
)),
check_load_contract(
&ModelDescription::from_parts(
vec![fixed(MEL, &[0, 72, 401], DataType::F32)],
Vec::new(),
Vec::new(),
),
&any_fixed_batch_contract(),
)
.expect_err("a read-back axis pinned at zero"),
refuse(ModelDescription::from_parts(
vec![mel()],
vec![optional(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![mel(), fixed("prompt", &[1], DataType::I32)],
vec![embedding()],
Vec::new(),
)),
refuse(ModelDescription::from_parts(
vec![mel()],
vec![embedding()],
vec![fixed("kv", &[1], DataType::F16)],
)),
]
}
#[test]
fn every_clause_reduces_to_the_three_cases_a_door_distinguishes() {
let expected = [
Rendered::Feature(FeatureRendering::new(
MEL,
"a declared feature".to_string(),
"missing".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
MEL,
"float32".to_string(),
"float16".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
MEL,
"rank 3".to_string(),
"rank 2".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
MEL,
"fixed".to_string(),
"range".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
MEL,
"axis 2 401".to_string(),
"axis 2 400".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
MEL,
"axis 0 any one non-zero fixed size".to_string(),
"axis 0 0".to_string(),
)),
Rendered::Feature(FeatureRendering::new(
EMBEDDING,
"a required output".to_string(),
"optional".to_string(),
)),
Rendered::UnsatisfiableInput("prompt".to_string()),
Rendered::UnsatisfiableState("kv".to_string()),
];
assert!(
expected
.iter()
.any(|case| matches!(case, Rendered::UnsatisfiableInput(_)))
&& expected
.iter()
.any(|case| matches!(case, Rendered::UnsatisfiableState(_)))
);
let violations = one_violation_per_clause();
assert_eq!(violations.len(), expected.len());
for (violation, want) in violations.into_iter().zip(expected) {
let message = violation.to_string();
let rendered = violation.rendered();
assert_eq!(rendered, want, "{message}");
if let Rendered::Feature(rendering) = rendered {
assert!(!rendering.feature().is_empty(), "{message}");
let states = rendering.clone().expected();
let declares = rendering.actual();
assert_ne!(states, declares, "{message}");
}
}
}
#[test]
fn a_dim_renders_for_a_violation_message() {
assert_eq!(Dim::Exactly(401).to_string(), "401");
assert_eq!(Dim::AnyFixed.to_string(), "any one non-zero fixed size");
assert_eq!(
Dim::AtLeast(224).to_string(),
"any one fixed size, at least 224"
);
assert_eq!(
Dim::Range(AxisRange::inclusive(10, 3001)).to_string(),
"10..=3001"
);
}
#[test]
fn a_feature_with_no_multi_array_constraint_renders_as_none() {
let description = ModelDescription::from_parts(
vec![FeatureInfo::from_parts(
MEL.to_string(),
Vec::new(),
None,
false,
None,
)],
vec![fixed(EMBEDDING, EMBEDDING_SHAPE, DataType::F32)],
Vec::new(),
);
let error = check_load_contract(&description, &identity_contract()).unwrap_err();
assert!(
matches!(&error, ContractViolation::DataType(d) if d.observed() == "none"),
"{error}"
);
assert!(error.to_string().contains("is none"), "{error}");
}
#[test]
fn checked_materialises_only_the_outputs_its_contract_names() {
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, crate::ComputeUnits::CpuOnly).expect("the committed bundle loads");
assert_eq!(
model
.description()
.outputs()
.iter()
.map(FeatureInfo::name)
.collect::<Vec<_>>(),
vec!["new_cell_state", "new_hidden_state", "vad_output"],
"this gate is about a graph with MORE outputs than the contract names"
);
let contract = LoadContract::new(
vec![
FeatureContract::new(
"audio_input",
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(4160)],
),
FeatureContract::new(
"cell_state",
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(128)],
),
FeatureContract::new(
"hidden_state",
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(128)],
),
],
vec![FeatureContract::new(
"vad_output",
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(1), Dim::Exactly(1)],
)],
StateContract::None,
);
let checked = Checked::new(model, &contract).expect("silero satisfies this contract");
let audio = MultiArray::zeros(&[1, 4160], DataType::F32).expect("one 256 ms window");
let hidden = MultiArray::zeros(&[1, 128], DataType::F32).expect("the LSTM's hidden state");
let cell = MultiArray::zeros(&[1, 128], DataType::F32).expect("the LSTM's cell state");
let outputs = checked
.predict_with(&[
("audio_input", &audio),
("hidden_state", &hidden),
("cell_state", &cell),
])
.expect("a real prediction through the door's own entry point");
assert_eq!(
outputs.names().collect::<Vec<_>>(),
vec!["vad_output"],
"the door asked for one output; the other two must not have been materialised"
);
assert_eq!(outputs.len(), 1);
assert_eq!(outputs.get("vad_output").unwrap().shape(), &[1, 1, 1]);
}