use super::*;
const P: usize = 512;
#[test]
fn budget_solver_matches_known_320x240_oracle() {
assert_eq!(fit_to_patch_budget(240, 320, PATCH_SIZE, P), (19, 26));
}
#[test]
fn budget_solver_square_image_fills_the_square_grid() {
assert_eq!(fit_to_patch_budget(512, 512, PATCH_SIZE, P), (22, 22));
assert_eq!(fit_to_patch_budget(64, 64, PATCH_SIZE, P), (22, 22));
}
#[test]
fn budget_solver_reproduces_probe_aspect_table() {
let rows: &[(usize, usize, usize, usize)] = &[
(640, 425, 28, 18), (640, 480, 26, 19), (640, 586, 23, 22), (483, 640, 19, 26), (480, 640, 19, 26), (427, 640, 18, 27), (426, 640, 18, 28), ];
for &(h, w, eh, ew) in rows {
assert_eq!(
fit_to_patch_budget(h, w, PATCH_SIZE, P),
(eh, ew),
"budget solver diverged from the probe oracle for {h}×{w}"
);
}
}
#[test]
fn budget_solver_respects_budget_and_is_maximal() {
for &(h, w) in &[
(240, 320),
(640, 425),
(512, 512),
(100, 3000),
(1000, 1000),
(37, 91),
(640, 587),
] {
let (hp, wp) = fit_to_patch_budget(h, w, PATCH_SIZE, P);
assert!(hp >= 1 && wp >= 1, "{h}×{w} degenerate grid ({hp},{wp})");
assert!(hp * wp <= P, "{h}×{w} exceeds budget: {hp}·{wp}");
assert!(
(hp + 1) * (wp + 1) > P,
"{h}×{w} not maximal: ({}+1)·({}+1) ≤ {P}",
hp,
wp
);
}
}
#[test]
fn budget_solver_is_scale_invariant_for_a_fixed_aspect() {
let g = fit_to_patch_budget(240, 320, PATCH_SIZE, P);
assert_eq!(g, (19, 26));
assert_eq!(fit_to_patch_budget(480, 640, PATCH_SIZE, P), g);
assert_eq!(fit_to_patch_budget(120, 160, PATCH_SIZE, P), g);
}
#[test]
fn budget_solver_extreme_aspect_clamps_short_side_to_one_patch() {
let (hp, wp) = fit_to_patch_budget(16, 16_000, PATCH_SIZE, P);
assert_eq!(hp, 1, "the 1-patch short-side clamp must hold");
assert!(wp >= 1 && hp * wp <= P, "grid ({hp},{wp}) invalid");
}
#[test]
fn resize_identity_is_exact() {
let src: Vec<f32> = (0..2 * 2 * 3).map(|i| i as f32).collect();
let out = resize_bilinear_antialias(&src, 2, 2, 3, 2, 2).expect("identity resize");
assert_eq!(out, src);
}
#[test]
fn resize_of_constant_field_is_constant() {
let src = vec![0.375f32; 3 * 5 * 2]; let up = resize_bilinear_antialias(&src, 3, 5, 2, 7, 9).expect("upscale");
assert!(
up.iter().all(|&v| (v - 0.375).abs() <= 1e-6),
"upscale drifted"
);
let down = resize_bilinear_antialias(&src, 3, 5, 2, 2, 2).expect("downscale");
assert!(
down.iter().all(|&v| (v - 0.375).abs() <= 1e-6),
"downscale drifted"
);
}
#[test]
fn resize_upscale_matches_hand_computed_bilinear() {
let out = resize_bilinear_antialias(&[0.0, 1.0], 1, 2, 1, 1, 4).expect("upscale resize");
assert_eq!(out, vec![0.0, 0.25, 0.75, 1.0]);
}
#[test]
fn resize_checker_upscale_matches_hand_computed_surface() {
let src = [0.0f32, 1.0, 1.0, 0.0];
let out = resize_bilinear_antialias(&src, 2, 2, 1, 4, 4).expect("checker upscale");
let at = |r: usize, c: usize| out[r * 4 + c];
assert_eq!(
(at(0, 0), at(0, 3), at(3, 0), at(3, 3)),
(0.0, 1.0, 1.0, 0.0)
);
for (r, c, want) in [
(1, 1, 0.375f32),
(1, 2, 0.625),
(2, 1, 0.625),
(2, 2, 0.375),
] {
assert!(
(at(r, c) - want).abs() <= 1e-6,
"checker[{r}][{c}] = {}",
at(r, c)
);
}
}
#[test]
fn resize_antialias_downscale_is_symmetric_lowpass() {
let out = resize_bilinear_antialias(&[0.0, 0.0, 255.0, 255.0], 1, 4, 1, 1, 2)
.expect("antialias downscale");
assert!(out[0] < out[1], "must be monotonic increasing");
assert!(
(out[0] + out[1] - 255.0).abs() <= 1e-3,
"must be symmetric around 127.5"
);
assert!(
out[0] > 1.0 && out[1] < 254.0,
"antialias must low-pass the edge"
);
}
#[test]
fn resize_keeps_channels_independent() {
let src = [0.0f32, 10.0, 20.0, 40.0, 50.0, 60.0]; let out = resize_bilinear_antialias(&src, 1, 2, 3, 1, 4).expect("channel resize");
let up1d = |a: f32, b: f32| [a, 0.75 * a + 0.25 * b, 0.25 * a + 0.75 * b, b];
for c in 0..3 {
let want = up1d(src[c], src[3 + c]);
for x in 0..4 {
assert!(
(out[x * 3 + c] - want[x]).abs() <= 1e-6,
"channel {c} pixel {x}: {} != {}",
out[x * 3 + c],
want[x]
);
}
}
}
#[test]
fn resize_preserves_vertical_constancy() {
let src = [3.0f32, 7.0, 3.0, 7.0, 3.0, 7.0, 3.0, 7.0];
let out = resize_bilinear_antialias(&src, 4, 2, 1, 2, 2).expect("height resize"); assert!((out[0] - out[2]).abs() <= 1e-6, "column 0 not constant");
assert!((out[1] - out[3]).abs() <= 1e-6, "column 1 not constant");
}
fn mono_to_rgb(mono: &[u8]) -> Vec<u8> {
let mut rgb = Vec::with_capacity(mono.len() * CHANNELS);
for &v in mono {
rgb.extend_from_slice(&[v; CHANNELS]);
}
rgb
}
fn channel(rgb: &[u8], c: usize) -> Vec<u8> {
rgb.iter().skip(c).step_by(CHANNELS).copied().collect()
}
#[test]
fn resize_u8_identity_is_exact() {
#[rustfmt::skip]
let src: [u8; 27] = [
1, 2, 3, 4, 5, 6, 7, 8, 9,
10, 11, 12, 13, 14, 15, 16, 17, 18,
19, 20, 21, 22, 23, 24, 25, 26, 27,
]; let out = resize_bilinear_antialias_u8(&src, 3, 3, 3, 3).expect("resize");
assert_eq!(out, src);
}
#[test]
fn resize_u8_constant_field_is_constant() {
let src = vec![123u8; 5 * 7 * 3]; let up = resize_bilinear_antialias_u8(&src, 5, 7, 9, 4).expect("upscale");
assert!(
up.iter().all(|&v| v == 123),
"upscale drifted from constant"
);
let down = resize_bilinear_antialias_u8(&src, 5, 7, 2, 2).expect("downscale");
assert!(
down.iter().all(|&v| v == 123),
"downscale drifted from constant"
);
}
#[test]
fn resize_u8_matches_hand_computed_pil_grid_checker() {
let src = mono_to_rgb(&[0u8, 255, 255, 0]);
let out = resize_bilinear_antialias_u8(&src, 2, 2, 4, 4).expect("resize");
#[rustfmt::skip]
let expected: [u8; 16] = [
0, 64, 191, 255,
64, 96, 159, 191,
191, 159, 96, 64,
255, 191, 64, 0,
];
for c in 0..CHANNELS {
assert_eq!(
channel(&out, c),
expected,
"channel {c} diverged from the PIL grid"
);
}
}
#[test]
fn resize_u8_per_pass_rounding_discriminates_from_float() {
let src = mono_to_rgb(&[0u8, 1, 255, 255]);
let out = resize_bilinear_antialias_u8(&src, 2, 2, 4, 4).expect("resize");
#[rustfmt::skip]
let expected: [u8; 16] = [
0, 0, 1, 1,
64, 64, 65, 65,
191, 191, 192, 192,
255, 255, 255, 255,
];
for c in 0..CHANNELS {
let ch = channel(&out, c);
assert_eq!(ch, expected, "channel {c} diverged from the PIL grid");
assert_eq!(
(ch[6], ch[10]),
(65, 192),
"channel {c} per-pass discriminant"
);
}
}
fn rgb_from_channels(r: &[u8], g: &[u8], b: &[u8]) -> Vec<u8> {
assert_eq!(r.len(), g.len());
assert_eq!(r.len(), b.len());
let mut rgb = Vec::with_capacity(r.len() * CHANNELS);
for i in 0..r.len() {
rgb.extend_from_slice(&[r[i], g[i], b[i]]);
}
rgb
}
#[test]
fn resize_u8_distinct_channels_match_independent_pil_grids() {
let r_src = [0u8, 255, 255, 0]; let g_src = [255u8, 0, 0, 255]; let b_src = [0u8, 1, 255, 255]; let src = rgb_from_channels(&r_src, &g_src, &b_src);
let out = resize_bilinear_antialias_u8(&src, 2, 2, 4, 4).expect("resize");
#[rustfmt::skip]
let r_expected: [u8; 16] = [
0, 64, 191, 255,
64, 96, 159, 191,
191, 159, 96, 64,
255, 191, 64, 0,
];
#[rustfmt::skip]
let g_expected: [u8; 16] = [
255, 191, 64, 0,
191, 159, 96, 64,
64, 96, 159, 191,
0, 64, 191, 255,
];
#[rustfmt::skip]
let b_expected: [u8; 16] = [
0, 0, 1, 1,
64, 64, 65, 65,
191, 191, 192, 192,
255, 255, 255, 255,
];
assert_eq!(
channel(&out, 0),
r_expected,
"R channel diverged from its own PIL grid"
);
assert_eq!(
channel(&out, 1),
g_expected,
"G channel diverged from its own PIL grid"
);
assert_eq!(
channel(&out, 2),
b_expected,
"B channel diverged from its own PIL grid"
);
}
#[test]
fn resize_u8_rejects_overflowing_geometry() {
match resize_bilinear_antialias_u8(&[0u8; 12], usize::MAX / 2, 2, 2, 2) {
Err(Error::PreprocessAllocation(bytes)) => assert_eq!(bytes, usize::MAX),
other => panic!("expected PreprocessAllocation, got {other:?}"),
}
}
#[test]
fn resize_u8_rejects_over_wide_source_extent() {
match resize_bilinear_antialias_u8(&[0u8; 12], 2, usize::MAX / 2, 2, 16) {
Err(Error::PreprocessAllocation(bytes)) => assert_eq!(bytes, usize::MAX),
other => panic!("expected PreprocessAllocation, got {other:?}"),
}
}
#[test]
fn resize_u8_rejects_overflowing_destination_length() {
match resize_bilinear_antialias_u8(&[0u8; 12], 2, 2, 2, usize::MAX / 2) {
Err(Error::PreprocessAllocation(bytes)) => assert_eq!(bytes, usize::MAX),
other => panic!("expected PreprocessAllocation, got {other:?}"),
}
}
#[test]
fn preprocess_image_pixel_values_come_from_u8_resize() {
let pattern = [0u8, 1, 255, 255];
let mut rgb = vec![0u8; 2 * 2 * CHANNELS];
for (px, &v) in pattern.iter().enumerate() {
for c in 0..CHANNELS {
rgb[px * CHANNELS + c] = v;
}
}
let base = vec![0.0f32; POS_EMBED_ELEMS];
let out = preprocess_image(&rgb, 2, 2, &base, 1).expect("preprocess");
assert_eq!(out.grid, (1, 1), "budget 1 → single-patch grid");
let resized_u8 = resize_bilinear_antialias_u8(&rgb, 2, 2, 16, 16).expect("resize");
assert_eq!(resized_u8.len(), 16 * 16 * CHANNELS);
for (k, &v) in resized_u8.iter().enumerate() {
assert_eq!(
out.pixel_values[k],
normalize_u8(v),
"pixel_values[{k}] must be normalize_u8 of the u8-resized byte"
);
}
let rgb_f32: Vec<f32> = rgb.iter().map(|&b| f32::from(b)).collect();
let float_resized =
resize_bilinear_antialias(&rgb_f32, 2, 2, CHANNELS, 16, 16).expect("float resize");
let float_u8: Vec<u8> = float_resized
.iter()
.map(|&v| v.round().clamp(0.0, 255.0) as u8)
.collect();
assert!(
resized_u8.iter().zip(&float_u8).any(|(&u, &f)| u != f),
"uint8 per-pass resize must differ from a float-then-round resize here"
);
}
#[test]
fn normalize_u8_maps_range_to_unit_interval() {
assert_eq!(normalize_u8(0), -1.0);
assert_eq!(normalize_u8(255), 1.0);
assert!(normalize_u8(128).abs() < 0.01);
}
#[test]
fn patchify_places_pixels_at_exact_slots() {
let grid_h = 2;
let grid_w = 3;
let img_h = grid_h * PATCH_SIZE;
let img_w = grid_w * PATCH_SIZE;
let mut img = vec![0.0f32; img_h * img_w * CHANNELS];
for y in 0..img_h {
for x in 0..img_w {
for c in 0..CHANNELS {
img[(y * img_w + x) * CHANNELS + c] = (y * 10_000 + x * 10 + c) as f32;
}
}
}
let budget = 10;
let (pixel_values, mask) = patchify(&img, grid_h, grid_w, budget).expect("patchify");
assert_eq!(pixel_values.len(), budget * PATCH_DIM);
assert_eq!(mask.len(), budget);
for ph in 0..grid_h {
for pw in 0..grid_w {
let row = ph * grid_w + pw;
let mut k = 0;
for py in 0..PATCH_SIZE {
for px in 0..PATCH_SIZE {
for c in 0..CHANNELS {
let y = ph * PATCH_SIZE + py;
let x = pw * PATCH_SIZE + px;
let want = (y * 10_000 + x * 10 + c) as f32;
assert_eq!(
pixel_values[row * PATCH_DIM + k],
want,
"slot (ph={ph},pw={pw},py={py},px={px},c={c}) misplaced"
);
k += 1;
}
}
}
}
}
}
#[test]
fn patchify_mask_and_padding_are_correct() {
let grid_h = 2;
let grid_w = 3; let img = vec![0.5f32; (grid_h * PATCH_SIZE) * (grid_w * PATCH_SIZE) * CHANNELS];
let budget = 8;
let (pixel_values, mask) = patchify(&img, grid_h, grid_w, budget).expect("patchify");
let n_real = grid_h * grid_w;
assert_eq!(
mask.iter().sum::<f32>(),
n_real as f32,
"mask must count real patches"
);
assert!(
mask[..n_real].iter().all(|&m| m == 1.0),
"real patches masked 1.0"
);
assert!(
mask[n_real..].iter().all(|&m| m == 0.0),
"pad patches masked 0.0"
);
assert!(
pixel_values[n_real * PATCH_DIM..].iter().all(|&v| v == 0.0),
"padded pixel rows must be zero"
);
}
#[test]
fn patchify_rejects_grid_over_budget() {
let grid_h = 3;
let grid_w = 3; let img = vec![0.0f32; (grid_h * PATCH_SIZE) * (grid_w * PATCH_SIZE) * CHANNELS];
match patchify(&img, grid_h, grid_w, 8) {
Err(Error::PatchCount(ref e)) if e.got() == 9 && e.max() == 8 => {}
other => panic!("expected PatchCount, got {other:?}"),
}
}
#[test]
fn parse_base_pos_grid_validates_exact_byte_length() {
assert_eq!(POS_EMBED_BYTES, 16 * 16 * 768 * 4);
assert_eq!(POS_EMBED_BYTES, 786_432);
let good = vec![0u8; POS_EMBED_BYTES];
let grid = parse_base_pos_grid(&good).expect("exact length parses");
assert_eq!(grid.len(), POS_EMBED_ELEMS);
match parse_base_pos_grid(&[0u8; 16]) {
Err(Error::PosEmbedLength(e)) if e.got() == 16 => assert_eq!(e.expected(), POS_EMBED_BYTES),
other => panic!("expected PosEmbedLength, got {other:?}"),
}
let long = vec![0u8; POS_EMBED_BYTES + 4];
assert!(matches!(
parse_base_pos_grid(&long),
Err(Error::PosEmbedLength(_))
));
}
#[test]
fn parse_base_pos_grid_decodes_little_endian_f32() {
let mut bytes = vec![0u8; POS_EMBED_BYTES];
bytes[0..4].copy_from_slice(&1.5f32.to_le_bytes());
bytes[4..8].copy_from_slice(&(-2.25f32).to_le_bytes());
let grid = parse_base_pos_grid(&bytes).expect("parse");
assert_eq!(grid[0], 1.5);
assert_eq!(grid[1], -2.25);
}
#[test]
fn lift_position_embeddings_flattens_and_zero_pads() {
let base = vec![0.7f32; POS_EMBED_ELEMS]; let grid_h = 3;
let grid_w = 4; let budget = 20;
let lifted = lift_position_embeddings(&base, grid_h, grid_w, budget).expect("lift");
assert_eq!(lifted.len(), budget * EMBEDDING_DIM);
let n_real = grid_h * grid_w;
assert!(
lifted[..n_real * EMBEDDING_DIM]
.iter()
.all(|&v| (v - 0.7).abs() <= 1e-6),
"real position rows must carry the resized (constant) grid"
);
assert!(
lifted[n_real * EMBEDDING_DIM..].iter().all(|&v| v == 0.0),
"padded position rows must be zero"
);
}
#[test]
fn lift_position_embeddings_row_order_matches_patch_order() {
let mut base = vec![0.0f32; POS_EMBED_ELEMS];
for gy in 0..POS_GRID_SIDE {
for gx in 0..POS_GRID_SIDE {
base[(gy * POS_GRID_SIDE + gx) * EMBEDDING_DIM] = (gy * 100 + gx) as f32;
}
}
let budget = POS_GRID_SIDE * POS_GRID_SIDE + 5;
let lifted = lift_position_embeddings(&base, POS_GRID_SIDE, POS_GRID_SIDE, budget).expect("lift");
for gy in 0..POS_GRID_SIDE {
for gx in 0..POS_GRID_SIDE {
let row = gy * POS_GRID_SIDE + gx;
assert_eq!(
lifted[row * EMBEDDING_DIM],
(gy * 100 + gx) as f32,
"position row {row} misaligned with patch order"
);
}
}
}
fn synthetic_rgb(width: usize, height: usize) -> Vec<u8> {
let mut data = vec![0u8; width * height * CHANNELS];
for y in 0..height {
for x in 0..width {
let base = (y * width + x) * CHANNELS;
data[base] = ((x * 255) / width.max(1)) as u8;
data[base + 1] = ((y * 255) / height.max(1)) as u8;
data[base + 2] = ((x + y) % 256) as u8;
}
}
data
}
#[test]
fn preprocess_image_produces_budget_shaped_tensors() {
let (w, h) = (320usize, 240usize);
let base = vec![0.1f32; POS_EMBED_ELEMS];
let rgb = synthetic_rgb(w, h);
let out = preprocess_image(&rgb, w, h, &base, P).expect("preprocess");
assert_eq!(out.grid, (19, 26), "grid must match the 320×240 oracle");
assert_eq!(out.pixel_values.len(), P * PATCH_DIM);
assert_eq!(out.attention_mask.len(), P);
assert_eq!(out.position_embeddings.len(), P * EMBEDDING_DIM);
let n_real = out.grid.0 * out.grid.1;
assert_eq!(out.attention_mask.iter().sum::<f32>(), n_real as f32);
assert!(
out.pixel_values.iter().all(|&v| (-1.0..=1.0).contains(&v)),
"normalized pixels must lie in [-1, 1]"
);
assert!(
out.pixel_values[n_real * PATCH_DIM..]
.iter()
.all(|&v| v == 0.0)
);
assert!(
out.position_embeddings[n_real * EMBEDDING_DIM..]
.iter()
.all(|&v| v == 0.0)
);
}
#[test]
fn preprocess_image_is_deterministic() {
let (w, h) = (200usize, 150usize);
let base: Vec<f32> = (0..POS_EMBED_ELEMS)
.map(|i| (i % 97) as f32 * 0.01)
.collect();
let rgb = synthetic_rgb(w, h);
let a = preprocess_image(&rgb, w, h, &base, P).expect("a");
let b = preprocess_image(&rgb, w, h, &base, P).expect("b");
assert_eq!(a.grid, b.grid);
assert_eq!(a.pixel_values, b.pixel_values);
assert_eq!(a.attention_mask, b.attention_mask);
assert_eq!(a.position_embeddings, b.position_embeddings);
}
#[test]
fn preprocess_rejects_width_over_axis_bound() {
let w = MAX_IMAGE_AXIS + 1;
let rgb = vec![0u8; w * 3];
let base = vec![0.0f32; POS_EMBED_ELEMS];
let err = preprocess_image(&rgb, w, 1, &base, P).unwrap_err();
assert!(matches!(err, Error::ImageDimensions(ref e) if e.width() == w && e.height() == 1));
}
#[test]
fn preprocess_rejects_height_over_axis_bound() {
let h = MAX_IMAGE_AXIS + 1;
let rgb = vec![0u8; h * 3];
let base = vec![0.0f32; POS_EMBED_ELEMS];
let err = preprocess_image(&rgb, 1, h, &base, P).unwrap_err();
assert!(matches!(err, Error::ImageDimensions(ref e) if e.width() == 1 && e.height() == h));
}
#[test]
fn preprocess_rejects_pillow_f32_inexact_extent() {
let w = 16_777_219usize; let rgb = vec![0u8; w * 3];
let base = vec![0.0f32; POS_EMBED_ELEMS];
let err = preprocess_image(&rgb, w, 1, &base, P).unwrap_err();
assert!(matches!(err, Error::ImageDimensions(ref e) if e.width() == w && e.height() == 1));
}
#[test]
fn preprocess_accepts_axis_bound_wide_panorama() {
let rgb = vec![127u8; MAX_IMAGE_AXIS * 3];
let base = vec![0.0f32; POS_EMBED_ELEMS];
let out = preprocess_image(&rgb, MAX_IMAGE_AXIS, 1, &base, P).unwrap();
let (gh, gw) = out.grid;
assert_eq!(gh, 1);
assert!((1..=P).contains(&gw));
assert_eq!(out.pixel_values.len(), P * PATCH_DIM);
assert_eq!(out.attention_mask.len(), P);
let real = out.attention_mask.iter().filter(|&&m| m == 1.0).count();
assert_eq!(real, gh * gw);
}
#[test]
fn preprocess_accepts_axis_bound_tall_strip() {
let rgb = vec![127u8; MAX_IMAGE_AXIS * 3];
let base = vec![0.0f32; POS_EMBED_ELEMS];
let out = preprocess_image(&rgb, 1, MAX_IMAGE_AXIS, &base, P).unwrap();
let (gh, gw) = out.grid;
assert_eq!(gw, 1);
assert!((1..=P).contains(&gh));
assert_eq!(out.pixel_values.len(), P * PATCH_DIM);
assert_eq!(out.attention_mask.len(), P);
let real = out.attention_mask.iter().filter(|&&m| m == 1.0).count();
assert_eq!(real, gh * gw);
}