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::{DynamicImage, ImageReader};
use std::io::Cursor;
const QWEN35_IMAGE_MEAN: f32 = 0.5;
const QWEN35_IMAGE_STD: f32 = 0.5;
#[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> {
let reader = ImageReader::new(Cursor::new(image_bytes))
.with_guessed_format()
.map_err(|e| VisionError::ImageDecode(format!("format detection failed: {e}")))?;
let img: DynamicImage = reader
.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 * merge;
if patch_size == 0 || merge == 0 || factor == 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 patch_len = in_channels * temporal * patch_size * patch_size;
let num_patches = grid.num_patches();
let mut out = vec![0.0f32; num_patches * patch_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))
}
#[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()));
}
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"
);
}
}