use super::VisionError;
use super::checkpoint::Qwen35VisionWeights;
use super::vit::{batch_matvec, gelu, layer_norm, softmax_inplace};
use crate::model::qwen35_config::VisionModelConfig;
use image::ImageReader;
use std::io::Cursor;
const QWEN35_IMAGE_MEAN: f32 = 0.5;
const QWEN35_IMAGE_STD: f32 = 0.5;
const MAX_IMAGE_DIMENSION_PIXELS: u32 = 2048;
const MAX_SERVE_VISION_PATCHES: usize = 256;
const MAX_SERVE_PREPROCESSED_BYTES: usize = 16 * 1024 * 1024;
const SERVE_VISION_MAX_PATCHES_ENV: &str = "LATTICE_VISION_MAX_PATCHES";
fn resolve_serve_max_patches(raw_override: Option<&str>) -> usize {
raw_override
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(MAX_SERVE_VISION_PATCHES)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct GridThw {
pub t: usize,
pub h: usize,
pub w: usize,
}
impl GridThw {
pub fn num_patches(&self) -> usize {
self.t * self.h * self.w
}
}
pub fn preprocess_qwen35_image(
image_bytes: &[u8],
cfg: &VisionModelConfig,
) -> Result<(Vec<f32>, GridThw), VisionError> {
preprocess_qwen35_image_inner(image_bytes, cfg, None)
}
pub fn preprocess_qwen35_image_for_serve(
image_bytes: &[u8],
cfg: &VisionModelConfig,
) -> Result<(Vec<f32>, GridThw), VisionError> {
let max_patches =
resolve_serve_max_patches(std::env::var(SERVE_VISION_MAX_PATCHES_ENV).ok().as_deref());
preprocess_qwen35_image_inner(image_bytes, cfg, Some(max_patches))
}
fn preprocess_qwen35_image_inner(
image_bytes: &[u8],
cfg: &VisionModelConfig,
max_patches: Option<usize>,
) -> Result<(Vec<f32>, GridThw), VisionError> {
let img = if let Some(max_patches) = max_patches {
if cfg.patch_size == 0 {
return Err(VisionError::InvalidConfig(
"patch_size must be > 0".to_string(),
));
}
let mut limits = image::Limits::default();
limits.max_image_width = Some(MAX_IMAGE_DIMENSION_PIXELS);
limits.max_image_height = Some(MAX_IMAGE_DIMENSION_PIXELS);
let mut header_reader = ImageReader::new(Cursor::new(image_bytes))
.with_guessed_format()
.map_err(|e| VisionError::ImageDecode(format!("format detection failed: {e}")))?;
header_reader.limits(limits.clone());
let (header_width, header_height) = header_reader.into_dimensions().map_err(|e| {
if matches!(e, image::ImageError::Limits(_)) {
VisionError::DimensionsExceeded(format!(
"image dimensions exceed the serving maximum of \
{MAX_IMAGE_DIMENSION_PIXELS}px per side: {e}"
))
} else {
VisionError::ImageDecode(format!("dimension read failed: {e}"))
}
})?;
let patches = (header_width as usize)
.div_ceil(cfg.patch_size)
.checked_mul((header_height as usize).div_ceil(cfg.patch_size))
.ok_or_else(|| {
VisionError::InvalidConfig("serving image patch count overflowed".to_string())
})?;
if patches > max_patches {
return Err(VisionError::DimensionsExceeded(format!(
"image {header_width}x{header_height} produces {patches} patches; serving \
maximum is {max_patches}"
)));
}
let preprocessed_bytes = preprocessed_f32_len(patches, cfg)?
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
VisionError::InvalidConfig(
"serving image preprocessing byte count overflowed".to_string(),
)
})?;
if preprocessed_bytes > MAX_SERVE_PREPROCESSED_BYTES {
return Err(VisionError::DimensionsExceeded(format!(
"serving image preprocessing requires {preprocessed_bytes} bytes; maximum is \
{MAX_SERVE_PREPROCESSED_BYTES}"
)));
}
let mut reader = ImageReader::new(Cursor::new(image_bytes))
.with_guessed_format()
.map_err(|e| VisionError::ImageDecode(format!("format detection failed: {e}")))?;
reader.limits(limits);
reader
.decode()
.map_err(|e| VisionError::ImageDecode(format!("decode failed: {e}")))?
} else {
ImageReader::new(Cursor::new(image_bytes))
.with_guessed_format()
.map_err(|e| VisionError::ImageDecode(format!("format detection failed: {e}")))?
.decode()
.map_err(|e| VisionError::ImageDecode(format!("decode failed: {e}")))?
};
let rgb = img.into_rgb8();
let (width, height) = (rgb.width() as usize, rgb.height() as usize);
let patch_size = cfg.patch_size;
let merge = cfg.spatial_merge_size;
let factor = patch_size.checked_mul(merge).ok_or_else(|| {
VisionError::InvalidConfig("patch_size * spatial_merge_size overflowed".to_string())
})?;
if patch_size == 0 || merge == 0 {
return Err(VisionError::InvalidConfig(
"patch_size and spatial_merge_size must be > 0".into(),
));
}
if !height.is_multiple_of(factor) || !width.is_multiple_of(factor) {
return Err(VisionError::InvalidConfig(format!(
"image {width}x{height} is not a multiple of patch_size*spatial_merge_size={factor} \
(dynamic resize is out of scope for ADR-069 S3a)"
)));
}
let grid_h = height / patch_size;
let grid_w = width / patch_size;
let grid = GridThw {
t: 1,
h: grid_h,
w: grid_w,
};
let in_channels = cfg.in_channels;
let temporal = cfg.temporal_patch_size;
let num_patches = grid.num_patches();
let patch_len = preprocessed_f32_len(1, cfg)?;
let output_len = preprocessed_f32_len(num_patches, cfg)?;
let mut out = vec![0.0f32; output_len];
let blocks_h = grid_h / merge;
let blocks_w = grid_w / merge;
let mut patch_idx = 0usize;
for block_row in 0..blocks_h {
for block_col in 0..blocks_w {
for sub_row in 0..merge {
for sub_col in 0..merge {
let h = block_row * merge + sub_row;
let w = block_col * merge + sub_col;
let py = h * patch_size;
let px = w * patch_size;
let row = &mut out[patch_idx * patch_len..(patch_idx + 1) * patch_len];
let mut k = 0usize;
for c in 0..in_channels {
for _t in 0..temporal {
for dy in 0..patch_size {
for dx in 0..patch_size {
let pixel = rgb.get_pixel((px + dx) as u32, (py + dy) as u32);
let raw = pixel[c] as f32 / 255.0;
row[k] = (raw - QWEN35_IMAGE_MEAN) / QWEN35_IMAGE_STD;
k += 1;
}
}
}
}
patch_idx += 1;
}
}
}
}
Ok((out, grid))
}
fn preprocessed_f32_len(num_patches: usize, cfg: &VisionModelConfig) -> Result<usize, VisionError> {
cfg.in_channels
.checked_mul(cfg.temporal_patch_size)
.and_then(|value| value.checked_mul(cfg.patch_size))
.and_then(|value| value.checked_mul(cfg.patch_size))
.and_then(|patch_len| num_patches.checked_mul(patch_len))
.ok_or_else(|| {
VisionError::InvalidConfig("image preprocessing element count overflowed".to_string())
})
}
#[allow(clippy::too_many_arguments)]
fn bilinear_pos_embed(
pos_embed: &[f32],
hidden: usize,
side: usize,
grid_h: usize,
grid_w: usize,
h_idx: usize,
w_idx: usize,
out: &mut [f32],
) {
let h_frac_pos = if grid_h > 1 {
h_idx as f32 * (side - 1) as f32 / (grid_h - 1) as f32
} else {
0.0
};
let w_frac_pos = if grid_w > 1 {
w_idx as f32 * (side - 1) as f32 / (grid_w - 1) as f32
} else {
0.0
};
let h_floor = h_frac_pos.floor() as usize;
let w_floor = w_frac_pos.floor() as usize;
let h_ceil = (h_floor + 1).min(side - 1);
let w_ceil = (w_floor + 1).min(side - 1);
let h_frac = h_frac_pos - h_floor as f32;
let w_frac = w_frac_pos - w_floor as f32;
let corners = [
(h_floor, w_floor, (1.0 - h_frac) * (1.0 - w_frac)),
(h_floor, w_ceil, (1.0 - h_frac) * w_frac),
(h_ceil, w_floor, h_frac * (1.0 - w_frac)),
(h_ceil, w_ceil, h_frac * w_frac),
];
for (ch, cw, weight) in corners {
if weight == 0.0 {
continue;
}
let row_idx = ch * side + cw;
let row = &pos_embed[row_idx * hidden..(row_idx + 1) * hidden];
for (o, &v) in out.iter_mut().zip(row.iter()) {
*o += v * weight;
}
}
}
pub(crate) fn apply_rope_inplace(x: &mut [f32], cos: &[f32], sin: &[f32]) {
let half = x.len() / 2;
let mut rotated = vec![0.0f32; x.len()];
for i in 0..half {
rotated[i] = -x[half + i];
rotated[half + i] = x[i];
}
for i in 0..x.len() {
x[i] = x[i] * cos[i] + rotated[i] * sin[i];
}
}
pub(crate) fn build_pos_embed_and_rope_tables(
weights: &Qwen35VisionWeights,
cfg: &VisionModelConfig,
grid: GridThw,
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let hidden = cfg.hidden_size;
let n = grid.num_patches();
let side = (cfg.num_position_embeddings as f64).sqrt().round() as usize;
let merge = cfg.spatial_merge_size;
let head_dim = hidden / cfg.num_heads;
let rope_dim = head_dim / 2; let rope_half = rope_dim / 2; let theta = 10_000.0_f32;
let inv_freq: Vec<f32> = (0..rope_half)
.map(|i| 1.0 / theta.powf((2 * i) as f32 / rope_dim as f32))
.collect();
let mut pos_embed_contrib = vec![0.0f32; n * hidden];
let mut cos_table = vec![0.0f32; n * head_dim];
let mut sin_table = vec![0.0f32; n * head_dim];
let blocks_h = grid.h / merge;
let blocks_w = grid.w / merge;
let mut patch_idx = 0usize;
for block_row in 0..blocks_h {
for block_col in 0..blocks_w {
for sub_row in 0..merge {
for sub_col in 0..merge {
let h_idx = block_row * merge + sub_row;
let w_idx = block_col * merge + sub_col;
let pos_slice =
&mut pos_embed_contrib[patch_idx * hidden..(patch_idx + 1) * hidden];
bilinear_pos_embed(
&weights.pos_embed,
hidden,
side,
grid.h,
grid.w,
h_idx,
w_idx,
pos_slice,
);
let mut rotary = vec![0.0f32; rope_dim];
for i in 0..rope_half {
rotary[i] = h_idx as f32 * inv_freq[i];
rotary[rope_half + i] = w_idx as f32 * inv_freq[i];
}
let cos_row = &mut cos_table[patch_idx * head_dim..(patch_idx + 1) * head_dim];
let sin_row = &mut sin_table[patch_idx * head_dim..(patch_idx + 1) * head_dim];
for i in 0..rope_dim {
let (s, c) = rotary[i].sin_cos();
cos_row[i] = c;
cos_row[rope_dim + i] = c;
sin_row[i] = s;
sin_row[rope_dim + i] = s;
}
patch_idx += 1;
}
}
}
}
debug_assert_eq!(patch_idx, n);
(pos_embed_contrib, cos_table, sin_table)
}
pub fn qwen35_vit_forward(
weights: &Qwen35VisionWeights,
cfg: &VisionModelConfig,
pixel_values: &[f32],
grid: GridThw,
) -> Result<Vec<f32>, VisionError> {
let hidden = cfg.hidden_size;
let n = grid.num_patches();
let patch_len = cfg.in_channels * cfg.temporal_patch_size * cfg.patch_size * cfg.patch_size;
if pixel_values.len() != n * patch_len {
return Err(VisionError::ShapeMismatch {
expected: n * patch_len,
actual: pixel_values.len(),
context: "qwen35_vit_forward: pixel_values length".into(),
});
}
let mut hidden_states = batch_matvec(
&weights.patch_embed_weight,
pixel_values,
n,
hidden,
patch_len,
);
for i in 0..n {
for j in 0..hidden {
hidden_states[i * hidden + j] += weights.patch_embed_bias[j];
}
}
let head_dim = hidden / cfg.num_heads;
let (pos_embed_contrib, cos_table, sin_table) =
build_pos_embed_and_rope_tables(weights, cfg, grid);
for i in 0..n * hidden {
hidden_states[i] += pos_embed_contrib[i];
}
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let n_heads = cfg.num_heads;
for block in &weights.blocks {
let residual = hidden_states.clone();
let mut normed = hidden_states.clone();
for i in 0..n {
layer_norm(
&mut normed[i * hidden..(i + 1) * hidden],
&block.norm1_weight,
&block.norm1_bias,
1e-6,
);
}
let mut qkv = batch_matvec(&block.qkv_weight, &normed, n, 3 * hidden, hidden);
for i in 0..n {
for j in 0..3 * hidden {
qkv[i * 3 * hidden + j] += block.qkv_bias[j];
}
}
for i in 0..n {
let base = i * 3 * hidden;
let cos_row = &cos_table[i * head_dim..(i + 1) * head_dim];
let sin_row = &sin_table[i * head_dim..(i + 1) * head_dim];
for h in 0..n_heads {
let q = &mut qkv[base + h * head_dim..base + (h + 1) * head_dim];
apply_rope_inplace(q, cos_row, sin_row);
let k_base = base + hidden;
let k = &mut qkv[k_base + h * head_dim..k_base + (h + 1) * head_dim];
apply_rope_inplace(k, cos_row, sin_row);
}
}
let attn_out = multihead_attention_full(&qkv, n, hidden, n_heads, head_dim, scale);
let proj_out = batch_matvec(&block.proj_weight, &attn_out, n, hidden, hidden);
for i in 0..n * hidden {
hidden_states[i] = residual[i] + proj_out[i] + block.proj_bias[i % hidden];
}
let residual = hidden_states.clone();
let mut normed = hidden_states.clone();
for i in 0..n {
layer_norm(
&mut normed[i * hidden..(i + 1) * hidden],
&block.norm2_weight,
&block.norm2_bias,
1e-6,
);
}
let mlp_dim = block.fc1_bias.len();
let mut fc1_out = batch_matvec(&block.fc1_weight, &normed, n, mlp_dim, hidden);
for i in 0..n {
for j in 0..mlp_dim {
let idx = i * mlp_dim + j;
fc1_out[idx] = gelu(fc1_out[idx] + block.fc1_bias[j]);
}
}
let fc2_out = batch_matvec(&block.fc2_weight, &fc1_out, n, hidden, mlp_dim);
for i in 0..n * hidden {
hidden_states[i] = residual[i] + fc2_out[i] + block.fc2_bias[i % hidden];
}
}
Ok(hidden_states)
}
fn multihead_attention_full(
qkv: &[f32],
n: usize,
hidden: usize,
n_heads: usize,
head_dim: usize,
scale: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; n * hidden];
for h in 0..n_heads {
let mut q_h = vec![0.0f32; n * head_dim];
let mut k_h = vec![0.0f32; n * head_dim];
let mut v_h = vec![0.0f32; n * head_dim];
for i in 0..n {
let base = i * 3 * hidden;
q_h[i * head_dim..(i + 1) * head_dim]
.copy_from_slice(&qkv[base + h * head_dim..base + (h + 1) * head_dim]);
k_h[i * head_dim..(i + 1) * head_dim].copy_from_slice(
&qkv[base + hidden + h * head_dim..base + hidden + (h + 1) * head_dim],
);
v_h[i * head_dim..(i + 1) * head_dim].copy_from_slice(
&qkv[base + 2 * hidden + h * head_dim..base + 2 * hidden + (h + 1) * head_dim],
);
}
let mut scores = vec![0.0f32; n * n];
for i in 0..n {
let qi = &q_h[i * head_dim..(i + 1) * head_dim];
for j in 0..n {
let kj = &k_h[j * head_dim..(j + 1) * head_dim];
let dot: f32 = qi.iter().zip(kj.iter()).map(|(a, b)| a * b).sum();
scores[i * n + j] = dot * scale;
}
}
for i in 0..n {
softmax_inplace(&mut scores[i * n..(i + 1) * n]);
}
for i in 0..n {
let attn_row = &scores[i * n..(i + 1) * n];
for j in 0..head_dim {
let mut acc = 0.0f32;
for k in 0..n {
acc += attn_row[k] * v_h[k * head_dim + j];
}
out[i * hidden + h * head_dim + j] = acc;
}
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn tiny_cfg() -> VisionModelConfig {
VisionModelConfig {
depth: 1,
hidden_size: 8,
num_heads: 2,
patch_size: 2,
spatial_merge_size: 2,
out_hidden_size: 8,
temporal_patch_size: 1,
num_position_embeddings: 16, in_channels: 3,
deepstack_visual_indexes: vec![],
intermediate_size: None,
}
}
fn make_test_png(w: u32, h: u32) -> Vec<u8> {
use image::RgbImage;
let mut img = RgbImage::new(w, h);
for y in 0..h {
for x in 0..w {
let v = ((x + y) % 256) as u8;
img.put_pixel(x, y, image::Rgb([v, v, v]));
}
}
let mut buf = Vec::new();
img.write_to(&mut std::io::Cursor::new(&mut buf), image::ImageFormat::Png)
.unwrap();
buf
}
#[test]
fn preprocess_rejects_misaligned_image() {
let cfg = tiny_cfg(); let png = make_test_png(6, 4); let err = preprocess_qwen35_image(&png, &cfg).unwrap_err();
assert!(matches!(err, VisionError::InvalidConfig(_)));
}
#[test]
fn preprocess_produces_expected_shape_and_grid() {
let cfg = tiny_cfg();
let png = make_test_png(8, 8); let (patches, grid) = preprocess_qwen35_image(&png, &cfg).expect("preprocess");
assert_eq!(grid, GridThw { t: 1, h: 4, w: 4 });
let patch_len = cfg.in_channels * cfg.temporal_patch_size * cfg.patch_size * cfg.patch_size;
assert_eq!(patches.len(), grid.num_patches() * patch_len);
assert!(patches.iter().all(|v| v.is_finite()));
}
#[test]
fn serving_preprocess_enforces_patch_budget_without_narrowing_embedding_path() {
let cfg = tiny_cfg();
let boundary = make_test_png(32, 32); assert!(
preprocess_qwen35_image_inner(&boundary, &cfg, Some(MAX_SERVE_VISION_PATCHES)).is_ok()
);
let over = make_test_png(36, 32); let err =
preprocess_qwen35_image_inner(&over, &cfg, Some(MAX_SERVE_VISION_PATCHES)).unwrap_err();
assert!(
matches!(err, VisionError::DimensionsExceeded(message) if message.contains("serving maximum is 256"))
);
assert!(
preprocess_qwen35_image(&over, &cfg).is_ok(),
"the serving-only latency budget must not narrow the embedding API"
);
}
#[test]
fn serving_preprocess_rejects_declared_pixel_dimensions_over_the_hard_cap() {
let cfg = tiny_cfg();
let over_width = make_test_png(MAX_IMAGE_DIMENSION_PIXELS + 1, 1);
let err = preprocess_qwen35_image_for_serve(&over_width, &cfg).unwrap_err();
assert!(
matches!(&err, VisionError::DimensionsExceeded(message) if message.contains("2048px")),
"expected DimensionsExceeded naming the 2048px cap, got: {err}"
);
}
#[test]
fn resolve_serve_max_patches_falls_back_to_default_for_every_invalid_shape() {
assert_eq!(resolve_serve_max_patches(None), MAX_SERVE_VISION_PATCHES);
assert_eq!(
resolve_serve_max_patches(Some("")),
MAX_SERVE_VISION_PATCHES
);
assert_eq!(
resolve_serve_max_patches(Some("not-a-number")),
MAX_SERVE_VISION_PATCHES
);
assert_eq!(
resolve_serve_max_patches(Some("0")),
MAX_SERVE_VISION_PATCHES
);
assert_eq!(
resolve_serve_max_patches(Some("-4")),
MAX_SERVE_VISION_PATCHES
);
}
#[test]
fn resolve_serve_max_patches_honors_a_positive_override() {
assert_eq!(resolve_serve_max_patches(Some("64")), 64);
assert_eq!(resolve_serve_max_patches(Some(" 64 ")), 64);
}
#[test]
fn serving_preprocess_inner_honors_a_lower_override_below_the_default_boundary() {
let cfg = tiny_cfg();
let boundary = make_test_png(32, 32); assert!(preprocess_qwen35_image_inner(&boundary, &cfg, Some(256)).is_ok());
let err = preprocess_qwen35_image_inner(&boundary, &cfg, Some(200)).unwrap_err();
assert!(
matches!(err, VisionError::DimensionsExceeded(message) if message.contains("serving maximum is 200"))
);
}
#[test]
fn serving_preprocess_caps_output_allocation_for_pathological_config() {
let mut cfg = tiny_cfg();
cfg.patch_size = 8;
cfg.spatial_merge_size = 1;
cfg.temporal_patch_size = 128;
let image = make_test_png(128, 128);
let err = preprocess_qwen35_image_for_serve(&image, &cfg).unwrap_err();
assert!(
matches!(err, VisionError::DimensionsExceeded(message) if message.contains("preprocessing requires"))
);
}
fn make_test_weights(cfg: &VisionModelConfig) -> Qwen35VisionWeights {
use crate::vision::checkpoint::{VisualBlockWeights, VisualMergerWeights};
let hidden = cfg.hidden_size;
let patch_len = cfg.in_channels * cfg.temporal_patch_size * cfg.patch_size * cfg.patch_size;
let mlp_dim = 2 * hidden;
let merge_in = cfg.spatial_merge_size * cfg.spatial_merge_size * hidden;
let mut state = 0x1234_5678_u32;
let mut next = move || {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
(state as f32 / u32::MAX as f32) * 0.2 - 0.1
};
let mut v = |n: usize| (0..n).map(|_| next()).collect::<Vec<f32>>();
let block = VisualBlockWeights {
qkv_weight: v(3 * hidden * hidden),
qkv_bias: v(3 * hidden),
proj_weight: v(hidden * hidden),
proj_bias: v(hidden),
fc1_weight: v(mlp_dim * hidden),
fc1_bias: v(mlp_dim),
fc2_weight: v(hidden * mlp_dim),
fc2_bias: v(hidden),
norm1_weight: vec![1.0; hidden],
norm1_bias: vec![0.0; hidden],
norm2_weight: vec![1.0; hidden],
norm2_bias: vec![0.0; hidden],
};
Qwen35VisionWeights {
patch_embed_weight: v(hidden * patch_len),
patch_embed_weight_shape: vec![
hidden,
cfg.in_channels,
cfg.temporal_patch_size,
cfg.patch_size,
cfg.patch_size,
],
patch_embed_bias: v(hidden),
pos_embed: v(cfg.num_position_embeddings * hidden),
blocks: vec![block],
merger: VisualMergerWeights {
fc1_weight: v(merge_in * merge_in),
fc1_bias: v(merge_in),
fc2_weight: v(cfg.out_hidden_size * merge_in),
fc2_bias: v(cfg.out_hidden_size),
norm_weight: vec![1.0; hidden],
norm_bias: vec![0.0; hidden],
},
}
}
#[test]
fn vit_forward_output_shape_and_finite() {
let cfg = tiny_cfg();
let weights = make_test_weights(&cfg);
let png = make_test_png(8, 8);
let (pixel_values, grid) = preprocess_qwen35_image(&png, &cfg).expect("preprocess");
let out = qwen35_vit_forward(&weights, &cfg, &pixel_values, grid).expect("forward");
assert_eq!(out.len(), grid.num_patches() * cfg.hidden_size);
assert!(out.iter().all(|v| v.is_finite()));
}
#[test]
fn vit_forward_rejects_pixel_length_mismatch() {
let cfg = tiny_cfg();
let weights = make_test_weights(&cfg);
let grid = GridThw { t: 1, h: 4, w: 4 };
let bad_pixels = vec![0.0f32; 3]; let err = qwen35_vit_forward(&weights, &cfg, &bad_pixels, grid).unwrap_err();
assert!(matches!(err, VisionError::ShapeMismatch { .. }));
}
#[test]
fn vit_forward_is_deterministic() {
let cfg = tiny_cfg();
let weights = make_test_weights(&cfg);
let png = make_test_png(8, 8);
let (pixel_values, grid) = preprocess_qwen35_image(&png, &cfg).expect("preprocess");
let out1 = qwen35_vit_forward(&weights, &cfg, &pixel_values, grid).expect("forward 1");
let out2 = qwen35_vit_forward(&weights, &cfg, &pixel_values, grid).expect("forward 2");
assert_eq!(out1, out2);
}
#[test]
fn vit_forward_is_sensitive_to_weight_mutation() {
let cfg = tiny_cfg();
let mut weights = make_test_weights(&cfg);
let png = make_test_png(8, 8);
let (pixel_values, grid) = preprocess_qwen35_image(&png, &cfg).expect("preprocess");
let baseline = qwen35_vit_forward(&weights, &cfg, &pixel_values, grid).expect("forward");
weights.blocks[0].qkv_weight[0] += 5.0;
let mutated = qwen35_vit_forward(&weights, &cfg, &pixel_values, grid).expect("forward");
assert_ne!(
baseline, mutated,
"weight mutation had no effect on ViT output"
);
}
}