use super::*;
use crate::ComputeUnits;
#[test]
fn model_is_send() {
fn assert_send<T: Send>() {}
assert_send::<Model>();
}
#[test]
fn load_missing_path_is_not_found() {
let err = Model::load("/nonexistent/Foo.mlmodelc", ComputeUnits::CpuOnly).unwrap_err();
assert!(matches!(err, crate::LoadError::NotFound(_)));
}
#[test]
fn compile_missing_source_is_not_found() {
let err = Model::compile("/nonexistent/foo.mlpackage").unwrap_err();
assert!(matches!(err, crate::CompileError::NotFound(_)));
}
fn shapes(list: &[&[usize]]) -> Vec<Vec<usize>> {
list.iter().map(|s| s.to_vec()).collect()
}
fn pinned(shape: &[usize]) -> Vec<AxisRange> {
shape.iter().map(|d| AxisRange::new(*d, 1)).collect()
}
struct Reading {
what: &'static str,
raw_type: isize,
declared: &'static [usize],
enumerated: &'static [&'static [usize]],
ranges: &'static [(usize, usize)],
dtype: isize,
verdict: ShapeConstraint,
}
const PROBE_READINGS: &[Reading] = &[
Reading {
what: "fixed.mlmodelc input mel",
raw_type: 2,
declared: &[1, 72, 401],
enumerated: &[&[1, 72, 401]],
ranges: &[(1, 1), (72, 1), (401, 1)],
dtype: 65552,
verdict: ShapeConstraint::Fixed,
},
Reading {
what: "fixed.mlmodelc output y",
raw_type: 2,
declared: &[1, 72, 401],
enumerated: &[&[1, 72, 401]],
ranges: &[(1, 1), (72, 1), (401, 1)],
dtype: 65552,
verdict: ShapeConstraint::Fixed,
},
Reading {
what: "enum3.mlmodelc input mel",
raw_type: 2,
declared: &[1, 72, 401],
enumerated: &[&[1, 72, 401], &[1, 72, 201], &[1, 72, 801]],
ranges: &[(1, 1), (72, 1), (401, 1)],
dtype: 65552,
verdict: ShapeConstraint::Enumerated,
},
Reading {
what: "range_equal.mlmodelc input mel",
raw_type: 3,
declared: &[1, 72, 401],
enumerated: &[],
ranges: &[(1, 1), (72, 1), (401, 1)],
dtype: 65552,
verdict: ShapeConstraint::Range,
},
Reading {
what: "range_open.mlmodelc input mel",
raw_type: 3,
declared: &[1, 72, 401],
enumerated: &[],
ranges: &[(1, 1), (72, 1), (10, 2992)],
dtype: 65552,
verdict: ShapeConstraint::Range,
},
Reading {
what: "nn_fixed.mlmodelc input mel",
raw_type: 2,
declared: &[1, 72, 401],
enumerated: &[&[1, 72, 401]],
ranges: &[(1, 1), (72, 1), (401, 1)],
dtype: 65568,
verdict: ShapeConstraint::Fixed,
},
Reading {
what: "nn_fixed.mlmodelc output y",
raw_type: 1,
declared: &[],
enumerated: &[],
ranges: &[],
dtype: 65568,
verdict: ShapeConstraint::Unspecified,
},
Reading {
what: "nn_range_unbounded.mlmodelc input mel",
raw_type: 3,
declared: &[1, 72, 401],
enumerated: &[],
ranges: &[(1, 1), (72, 1), (1, 9_223_372_036_854_775_807)],
dtype: 65568,
verdict: ShapeConstraint::Range,
},
];
#[test]
fn every_measured_probe_reading_classifies_to_its_verdict() {
for reading in PROBE_READINGS {
let ranges: Vec<AxisRange> = reading
.ranges
.iter()
.map(|(location, length)| AxisRange::new(*location, *length))
.collect();
assert_eq!(
classify_shape_constraint(
reading.raw_type,
reading.declared,
&shapes(reading.enumerated),
&ranges
),
reading.verdict,
"{}",
reading.what
);
}
}
#[test]
fn an_fp16_mlprogram_conversion_reports_float16_io() {
for reading in PROBE_READINGS {
let declared = crate::DataType::from_raw(reading.dtype);
let expected = if reading.what.starts_with("nn_") {
crate::DataType::F32
} else {
crate::DataType::F16
};
assert_eq!(declared, expected, "{}", reading.what);
}
}
#[test]
fn a_fixed_shape_graph_is_classified_without_a_dedicated_fixed_type_code() {
assert_eq!(
classify_shape_constraint(2, &[1, 4160], &shapes(&[&[1, 4160]]), &pinned(&[1, 4160])),
ShapeConstraint::Fixed
);
}
#[test]
fn a_range_axis_is_classified_range() {
assert_eq!(
classify_shape_constraint(
3,
&[1, 401, 60],
&[],
&[
AxisRange::new(1, 1),
AxisRange::new(10, 2992),
AxisRange::new(60, 1)
]
),
ShapeConstraint::Range
);
}
#[test]
fn several_enumerated_shapes_are_classified_enumerated() {
assert_eq!(
classify_shape_constraint(
2,
&[1, 72, 401],
&shapes(&[&[1, 72, 401], &[1, 72, 201]]),
&pinned(&[1, 72, 401])
),
ShapeConstraint::Enumerated
);
}
#[test]
fn an_unspecified_constraint_is_named_and_is_never_fixed() {
assert_eq!(
classify_shape_constraint(1, &[], &[], &[]),
ShapeConstraint::Unspecified
);
assert_eq!(
classify_shape_constraint(1, &[1, 1], &shapes(&[&[1, 1]]), &pinned(&[1, 1])),
ShapeConstraint::Unspecified,
"`…TypeUnspecified` establishes nothing, unit ranges or not"
);
}
#[test]
fn an_unmeasured_type_code_is_unknown_and_keeps_itself() {
assert_eq!(
classify_shape_constraint(7, &[1], &[], &[]),
ShapeConstraint::Unknown(7)
);
assert_eq!(
classify_shape_constraint(99, &[1], &shapes(&[&[1]]), &pinned(&[1])),
ShapeConstraint::Unknown(99),
"a code this door has never seen must carry itself into the diagnosis"
);
}
#[test]
fn a_range_constraint_stays_range_even_when_every_range_is_one() {
assert_eq!(
classify_shape_constraint(3, &[1, 72, 401], &[], &pinned(&[1, 72, 401])),
ShapeConstraint::Range,
"an equal-bound `RangeDim` reports unit ranges and is still symbolic"
);
assert_eq!(
classify_shape_constraint(
3,
&[1, 72, 401],
&shapes(&[&[1, 72, 401]]),
&pinned(&[1, 72, 401])
),
ShapeConstraint::Range,
"a `shapeRange` that also lists one enumerated shape is still a range"
);
}
#[test]
fn the_raw_code_conjunct_is_load_bearing() {
assert_ne!(
classify_shape_constraint(3, &[1, 72, 401], &[], &pinned(&[1, 72, 401])),
ShapeConstraint::Fixed
);
}
#[test]
fn raw_two_listing_no_enumerated_shape_fails_closed() {
assert_eq!(
classify_shape_constraint(2, &[1, 72, 401], &[], &pinned(&[1, 72, 401])),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::NoShapes)
);
}
#[test]
fn a_sole_enumerated_shape_that_is_not_the_declared_one_fails_closed() {
assert_eq!(
classify_shape_constraint(
2,
&[1, 72, 401],
&shapes(&[&[1, 72, 201]]),
&pinned(&[1, 72, 401])
),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SoleShapeIsNotDeclared)
);
}
#[test]
fn a_range_list_that_does_not_cover_every_axis_fails_closed() {
assert_eq!(
classify_shape_constraint(
2,
&[1, 72, 401],
&shapes(&[&[1, 72, 401]]),
&pinned(&[1, 72])
),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape)
);
}
#[test]
fn a_range_that_does_not_pin_the_declared_size_fails_closed() {
assert_eq!(
classify_shape_constraint(
2,
&[1, 72, 401],
&shapes(&[&[1, 72, 401]]),
&[
AxisRange::new(1, 1),
AxisRange::new(72, 1),
AxisRange::new(401, 2)
]
),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape),
"an axis admitting a second size is not pinned"
);
assert_eq!(
classify_shape_constraint(
2,
&[1, 72, 401],
&shapes(&[&[1, 72, 401]]),
&[
AxisRange::new(1, 1),
AxisRange::new(72, 1),
AxisRange::new(400, 1)
]
),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape),
"an axis pinned at a size other than the declared one is not this shape"
);
}
#[test]
fn a_zero_length_range_is_not_fixed() {
assert_eq!(
classify_shape_constraint(
2,
&[1, 72],
&shapes(&[&[1, 72]]),
&[AxisRange::new(1, 1), AxisRange::new(72, 0)]
),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape)
);
}
#[test]
fn a_declared_shape_with_no_axes_pins_nothing() {
assert_eq!(
classify_shape_constraint(2, &[], &shapes(&[&[]]), &[]),
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape)
);
}
#[test]
fn shape_constraint_renders_for_a_contract_mismatch_message() {
assert_eq!(ShapeConstraint::Fixed.to_string(), "fixed");
assert_eq!(ShapeConstraint::Enumerated.to_string(), "enumerated");
assert_eq!(ShapeConstraint::Range.to_string(), "range");
assert_eq!(ShapeConstraint::Unspecified.to_string(), "unspecified");
assert_eq!(ShapeConstraint::Unknown(7).to_string(), "unknown(7)");
assert_eq!(
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::NoShapes).to_string(),
"unmeasured(no enumerated shape)"
);
assert_eq!(
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SoleShapeIsNotDeclared).to_string(),
"unmeasured(sole enumerated shape is not the declared shape)"
);
assert_eq!(
ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape).to_string(),
"unmeasured(per-axis ranges do not pin the declared shape)"
);
}
#[test]
fn an_inclusive_axis_range_is_the_measured_span() {
assert_eq!(AxisRange::inclusive(10, 3001), AxisRange::new(10, 2992));
assert_eq!(AxisRange::inclusive(401, 401), AxisRange::new(401, 1));
assert_eq!(AxisRange::inclusive(10, 3001).min(), 10);
assert_eq!(AxisRange::inclusive(10, 3001).count(), 2992);
}
#[test]
fn an_inverted_inclusive_range_saturates_rather_than_wrapping() {
assert_eq!(AxisRange::inclusive(400, 10), AxisRange::new(400, 1));
}
#[test]
fn an_axis_range_renders_for_a_contract_mismatch_message() {
assert_eq!(AxisRange::new(401, 1).to_string(), "401");
assert_eq!(AxisRange::new(10, 2992).to_string(), "10..=3001");
assert_eq!(AxisRange::new(7, 0).to_string(), "(no size)");
}