use super::*;
#[test]
fn rgb8_image_accepts_valid_geometry_and_exposes_dims() {
let data = vec![7u8; 4 * 3 * 3]; let img = Rgb8Image::new(&data, 4, 3).expect("valid");
assert_eq!(img.width(), 4);
assert_eq!(img.height(), 3);
assert_eq!(img.data().len(), 4 * 3 * 3);
assert_eq!(img.data(), data.as_slice());
}
#[test]
fn rgb8_image_rejects_zero_width() {
let data: Vec<u8> = Vec::new();
match Rgb8Image::new(&data, 0, 3) {
Err(Error::ImageDimensions(ref e)) if e.width() == 0 && e.height() == 3 => {}
other => panic!("expected ImageDimensions, got {other:?}"),
}
}
#[test]
fn rgb8_image_rejects_zero_height() {
let data: Vec<u8> = Vec::new();
match Rgb8Image::new(&data, 4, 0) {
Err(Error::ImageDimensions(ref e)) if e.width() == 4 && e.height() == 0 => {}
other => panic!("expected ImageDimensions, got {other:?}"),
}
}
#[test]
fn rgb8_image_rejects_length_mismatch() {
let data = vec![0u8; 4 * 3 * 3 - 1]; match Rgb8Image::new(&data, 4, 3) {
Err(Error::ImageDataLength(e)) => {
assert_eq!(e.got(), 4 * 3 * 3 - 1);
assert_eq!(e.expected(), 4 * 3 * 3);
}
other => panic!("expected ImageDataLength, got {other:?}"),
}
}
#[test]
fn rgb8_image_rejects_size_overflow() {
let data = [0u8; 1];
match Rgb8Image::new(&data, usize::MAX, 2) {
Err(Error::ImageDimensions(_)) => {}
other => panic!("expected ImageDimensions on overflow, got {other:?}"),
}
}
#[test]
fn options_default_equals_new_and_is_cpu_and_gpu() {
assert_eq!(ImageEmbedderOptions::default(), ImageEmbedderOptions::new());
assert_eq!(ImageEmbedderOptions::new().compute(), DEFAULT_IMAGE_COMPUTE);
assert_eq!(DEFAULT_IMAGE_COMPUTE, ComputeUnits::CpuAndGpu);
}
#[test]
fn options_with_and_set_compute() {
let opts = ImageEmbedderOptions::new().with_compute(ComputeUnits::All);
assert_eq!(opts.compute(), ComputeUnits::All);
let mut opts = ImageEmbedderOptions::new();
opts.set_compute(ComputeUnits::CpuOnly);
assert_eq!(opts.compute(), ComputeUnits::CpuOnly);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_roundtrip() {
let opts = ImageEmbedderOptions::new().with_compute(ComputeUnits::CpuAndNeuralEngine);
let json = serde_json::to_string(&opts).unwrap();
assert!(json.contains("cpu_and_neural_engine"), "serialized: {json}");
let back: ImageEmbedderOptions = serde_json::from_str(&json).unwrap();
assert_eq!(back, opts);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_defaults_missing_compute_to_the_module_default() {
let back: ImageEmbedderOptions = serde_json::from_str("{}").unwrap();
assert_eq!(back, ImageEmbedderOptions::new());
}
fn bundle(p: usize, n_real: usize) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let mut pixel_values = vec![0.0f32; p * PATCH_DIM];
let mut position_embeddings = vec![0.0f32; p * EMBEDDING_DIM];
let mut attention_mask = vec![0.0f32; p];
pixel_values[..n_real * PATCH_DIM].fill(0.5);
position_embeddings[..n_real * EMBEDDING_DIM].fill(0.5);
attention_mask[..n_real].fill(1.0);
(pixel_values, position_embeddings, attention_mask)
}
#[test]
fn preprocessed_image_accepts_well_formed_bundle() {
let (px, pos, mask) = bundle(4, 3);
let pre = PreprocessedImage::try_new(px, pos, mask, 4).expect("well-formed bundle");
assert_eq!(pre.max_num_patches(), 4);
assert_eq!(pre.pixel_values().len(), 4 * PATCH_DIM);
assert_eq!(pre.position_embeddings().len(), 4 * EMBEDDING_DIM);
assert_eq!(pre.attention_mask(), &[1.0, 1.0, 1.0, 0.0]);
}
#[test]
fn preprocessed_image_accepts_full_budget_bundle() {
let (px, pos, mask) = bundle(4, 4);
PreprocessedImage::try_new(px, pos, mask, 4).expect("full-budget bundle");
}
#[test]
fn preprocessed_image_accepts_negative_zero_mask_pad() {
let (px, pos, mut mask) = bundle(4, 3);
mask[3] = -0.0; PreprocessedImage::try_new(px, pos, mask, 4).expect("negative-zero pad accepted");
}
#[test]
fn preprocessed_image_rejects_zero_budget() {
match PreprocessedImage::try_new(vec![], vec![], vec![], 0) {
Err(Error::PreprocessedPatchBudget(0)) => {}
other => panic!("expected PreprocessedPatchBudget, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_overflowing_budget() {
match PreprocessedImage::try_new(vec![], vec![], vec![], usize::MAX) {
Err(Error::PreprocessedPatchBudget(_)) => {}
other => panic!("expected PreprocessedPatchBudget, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_wrong_pixel_values_length() {
let (mut px, pos, mask) = bundle(4, 3);
px.pop();
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedLength(e)) if e.feature() == "pixel_values" => {
assert_eq!(e.got(), 4 * PATCH_DIM - 1);
assert_eq!(e.expected(), 4 * PATCH_DIM);
}
other => panic!("expected PreprocessedLength, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_wrong_position_embeddings_length() {
let (px, mut pos, mask) = bundle(4, 3);
pos.push(0.0);
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedLength(e)) if e.feature() == "position_embeddings" => {
assert_eq!(e.got(), 4 * EMBEDDING_DIM + 1);
assert_eq!(e.expected(), 4 * EMBEDDING_DIM);
}
other => panic!("expected PreprocessedLength, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_wrong_mask_length() {
let (px, pos, _mask) = bundle(4, 3);
let mask = vec![0.0f32; 5]; match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedLength(ref e))
if e.feature() == "attention_mask" && e.got() == 5 && e.expected() == 4 => {}
other => panic!("expected PreprocessedLength, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_non_finite_pixel_values() {
let (mut px, pos, mask) = bundle(4, 3);
px[10] = f32::NAN;
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedNonFinite(ref e))
if e.feature() == "pixel_values" && e.index() == 10 => {}
other => panic!("expected PreprocessedNonFinite, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_non_finite_position_embeddings() {
let (px, mut pos, mask) = bundle(4, 3);
pos[0] = f32::NEG_INFINITY;
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedNonFinite(ref e))
if e.feature() == "position_embeddings" && e.index() == 0 => {}
other => panic!("expected PreprocessedNonFinite, got {other:?}"),
}
}
#[test]
fn preprocessed_image_classifies_nan_mask_as_mask_value() {
let (px, pos, mut mask) = bundle(4, 3);
mask[1] = f32::NAN;
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedMaskValue(e)) if e.index() == 1 => assert!(e.value().is_nan()),
other => panic!("expected PreprocessedMaskValue, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_mask_value_outside_domain() {
let (px, pos, mut mask) = bundle(4, 3);
mask[1] = 0.5;
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedMaskValue(e)) if e.index() == 1 => assert_eq!(e.value(), 0.5),
other => panic!("expected PreprocessedMaskValue, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_mask_one_after_zero() {
let (px, pos, _mask) = bundle(4, 3);
let mask = vec![1.0, 0.0, 1.0, 0.0];
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedMaskOrder(2)) => {}
other => panic!("expected PreprocessedMaskOrder, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_all_pad_mask() {
let (px, pos, mask) = bundle(4, 0);
match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedMaskEmpty) => {}
other => panic!("expected PreprocessedMaskEmpty, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_nonzero_pixel_pad_row() {
let (mut px, pos, mask) = bundle(4, 3);
px[3 * PATCH_DIM + 5] = 0.25; match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedPadNonZero(e)) if e.feature() == "pixel_values" => {
assert_eq!(e.index(), 3 * PATCH_DIM + 5)
}
other => panic!("expected PreprocessedPadNonZero, got {other:?}"),
}
}
#[test]
fn preprocessed_image_rejects_nonzero_position_embedding_pad_row() {
let (px, mut pos, mask) = bundle(4, 3);
pos[3 * EMBEDDING_DIM] = 1e-3; match PreprocessedImage::try_new(px, pos, mask, 4) {
Err(Error::PreprocessedPadNonZero(e)) if e.feature() == "position_embeddings" => {
assert_eq!(e.index(), 3 * EMBEDDING_DIM)
}
other => panic!("expected PreprocessedPadNonZero, got {other:?}"),
}
}
#[test]
fn check_patch_budget_accepts_equal_and_rejects_mismatch() {
check_patch_budget(512, 512).expect("equal budgets accepted");
match check_patch_budget(256, 512) {
Err(Error::PatchBudgetMismatch(ref e)) if e.input() == 256 && e.model() == 512 => {}
other => panic!("expected PatchBudgetMismatch, got {other:?}"),
}
}
#[test]
fn internal_pipeline_output_passes_public_validation() {
use super::preprocess::{POS_EMBED_ELEMS, preprocess_image};
let v = preprocess_image(
&[128u8; 8 * 8 * 3],
8,
8,
&vec![0.0f32; POS_EMBED_ELEMS],
512,
)
.expect("preprocess");
let real = v.grid.0 * v.grid.1;
let ones = v.attention_mask.iter().filter(|&&m| m == 1.0).count();
assert_eq!(ones, real, "mask real-count equals the resolved grid");
PreprocessedImage::try_new(v.pixel_values, v.position_embeddings, v.attention_mask, 512)
.expect("pipeline output passes public validation");
}
#[test]
fn preprocessed_image_debug_is_compact() {
let (px, pos, mask) = bundle(4, 3);
let pre = PreprocessedImage::try_new(px, pos, mask, 4).expect("well-formed");
let debug = format!("{pre:?}");
assert!(debug.contains("max_num_patches"), "{debug}");
assert!(debug.contains("num_real_patches: 3"), "{debug}");
assert!(!debug.contains("pixel_values"), "{debug}");
}
use crate::{
AxisRange, ComputeUnits, FeatureInfo, Model, ModelDescription,
embeddings::siglip::error::contract_violation, model::RawShapeConstraint,
};
const STAGED_PATCH_BUDGET: usize = 512;
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 vision_description(p: usize) -> ModelDescription {
ModelDescription::from_parts(
vec![
fixed(names::PIXEL_VALUES, &[1, p, PATCH_DIM], DataType::F32),
fixed(
names::POSITION_EMBEDDINGS,
&[1, p, EMBEDDING_DIM],
DataType::F32,
),
fixed(names::ATTENTION_MASK, &[1, p], DataType::F32),
],
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
)
}
fn check(description: &ModelDescription) -> Result<()> {
let declared = declared_patch_budget(description)?;
crate::model::contract::check_load_contract(description, &image_contract(declared))
.map_err(contract_violation)
}
#[test]
fn the_contract_accepts_the_staged_geometry() {
assert!(check(&vision_description(STAGED_PATCH_BUDGET)).is_ok());
}
#[test]
fn the_contract_reads_back_whatever_budget_the_graph_pins() {
for p in [1usize, 64, 256, STAGED_PATCH_BUDGET, 1024] {
let description = vision_description(p);
assert!(check(&description).is_ok(), "budget {p}");
assert_eq!(
declared_patch_budget(&description).expect("declared"),
p,
"the budget read back must be the one the graph pins"
);
}
}
#[test]
fn the_contract_refuses_inputs_that_disagree_about_the_budget() {
let description = ModelDescription::from_parts(
vec![
fixed(
names::PIXEL_VALUES,
&[1, STAGED_PATCH_BUDGET, PATCH_DIM],
DataType::F32,
),
fixed(
names::POSITION_EMBEDDINGS,
&[1, 256, EMBEDDING_DIM],
DataType::F32,
),
fixed(
names::ATTENTION_MASK,
&[1, STAGED_PATCH_BUDGET],
DataType::F32,
),
],
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::POSITION_EMBEDDINGS),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_flexible_patch_budget() {
let description = ModelDescription::from_parts(
vec![
multi_array(
names::PIXEL_VALUES,
&[1, STAGED_PATCH_BUDGET, PATCH_DIM],
DataType::F32,
false,
3,
Vec::new(),
&[1, STAGED_PATCH_BUDGET, PATCH_DIM],
),
fixed(
names::POSITION_EMBEDDINGS,
&[1, STAGED_PATCH_BUDGET, EMBEDDING_DIM],
DataType::F32,
),
fixed(
names::ATTENTION_MASK,
&[1, STAGED_PATCH_BUDGET],
DataType::F32,
),
],
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::PIXEL_VALUES),
"{err}"
);
}
#[test]
fn a_zero_patch_budget_is_refused_before_a_contract_exists() {
let description = vision_description(0);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::PIXEL_VALUES),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_wrong_patch_dim_or_mask_dtype() {
let wrong_patch_dim = ModelDescription::from_parts(
vec![
fixed(
names::PIXEL_VALUES,
&[1, STAGED_PATCH_BUDGET, 3 * 14 * 14],
DataType::F32,
),
fixed(
names::POSITION_EMBEDDINGS,
&[1, STAGED_PATCH_BUDGET, EMBEDDING_DIM],
DataType::F32,
),
fixed(
names::ATTENTION_MASK,
&[1, STAGED_PATCH_BUDGET],
DataType::F32,
),
],
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert!(matches!(
check(&wrong_patch_dim),
Err(Error::ContractMismatch(_))
));
let mut int_mask = vision_description(STAGED_PATCH_BUDGET);
int_mask = ModelDescription::from_parts(
vec![
fixed(
names::PIXEL_VALUES,
&[1, STAGED_PATCH_BUDGET, PATCH_DIM],
DataType::F32,
),
fixed(
names::POSITION_EMBEDDINGS,
&[1, STAGED_PATCH_BUDGET, EMBEDDING_DIM],
DataType::F32,
),
fixed(
names::ATTENTION_MASK,
&[1, STAGED_PATCH_BUDGET],
DataType::I32,
),
],
int_mask.outputs().to_vec(),
Vec::new(),
);
let err = check(&int_mask).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::ATTENTION_MASK),
"{err}"
);
}
#[test]
fn the_contract_refuses_an_extra_required_input() {
let mut inputs = vision_description(STAGED_PATCH_BUDGET).inputs().to_vec();
inputs.push(fixed("spatial_shapes", &[1, 2], DataType::I32));
let description = ModelDescription::from_parts(
inputs,
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert!(
matches!(check(&description), Err(Error::UnsatisfiableInput(name)) if name == "spatial_shapes"),
"{:?}",
check(&description)
);
}
#[test]
fn the_contract_accepts_an_extra_optional_input() {
let mut inputs = vision_description(STAGED_PATCH_BUDGET).inputs().to_vec();
inputs.push(multi_array(
"spatial_shapes",
&[1, 2],
DataType::I32,
true,
2,
vec![vec![1, 2]],
&[1, 2],
));
let description = ModelDescription::from_parts(
inputs,
vec![fixed(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert!(check(&description).is_ok());
}
#[test]
fn the_contract_refuses_an_optional_features_output() {
let description = ModelDescription::from_parts(
vision_description(STAGED_PATCH_BUDGET).inputs().to_vec(),
vec![multi_array(
names::IMAGE_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
true,
2,
vec![vec![1, EMBEDDING_DIM]],
&[1, EMBEDDING_DIM],
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::IMAGE_FEATURES),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_graph_that_declares_state() {
let base = vision_description(STAGED_PATCH_BUDGET);
let description = ModelDescription::from_parts(
base.inputs().to_vec(),
base.outputs().to_vec(),
vec![fixed("kv_cache", &[1, 8], DataType::F32)],
);
assert!(
matches!(check(&description), Err(Error::UnsatisfiableState(name)) if name == "kv_cache")
);
}
#[test]
fn the_image_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::PIXEL_VALUES).is_none(),
"silero declares no `pixel_values`, which is what makes it this gate's model"
);
let err = declared_patch_budget(model.description()).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m)
if m.feature() == names::PIXEL_VALUES && m.actual() == "missing"),
"{err}"
);
let violation = Checked::new(model, &image_contract(STAGED_PATCH_BUDGET))
.expect_err("silero does not satisfy the siglip vision contract");
assert!(
matches!(&violation, crate::model::contract::ContractViolation::Missing(m)
if m.feature() == names::PIXEL_VALUES),
"expected `pixel_values` missing, got {violation}"
);
}