//! Turning an image into the `"patches"` tensor the encoder graph wants. //! //! The graph folds DINOv3's patch-embedding Conv2d into a single matmul, //! which means the flattening order here has to agree exactly with the //! order the convolution weight was flattened in. PyTorch stores that //! weight as `[out_channels, in_channels, kh, kw]`, so within one patch //! the element order is **channel-major**: //! //! ```text //! index = c * patch_size² + ky * patch_size + kx //! ``` //! //! Get this wrong and the model still runs, producing confident nonsense //! — so it is pinned down by tests below. use crate::dinov3::Config; /// ImageNet statistics from the model's `preprocessor_config.json`. pub const IMAGE_MEAN: [f32; 3] = [0.485, 0.456, 0.406]; pub const IMAGE_STD: [f32; 3] = [0.229, 0.224, 0.225]; /// Flatten an already-normalized CHW pixel tensor into patches. /// /// `pixels` is `[3, image_size, image_size]`, matching what HuggingFace's /// image processor hands to the model as `pixel_values`. Taking this form /// directly is what lets the desktop verifier feed the exact same tensor /// as the reference implementation, keeping resize and normalization /// differences out of a numerics comparison. /// /// Returns `[num_patches, patch_dim]` in row-major grid order. pub fn patches_from_pixels_chw(pixels: &[f32], config: &Config) -> Vec { let size = config.image_size; let ps = config.patch_size; let grid = config.grid(); assert_eq!( pixels.len(), 3 * size * size, "expected a [3, {size}, {size}] pixel tensor, got {} values", pixels.len() ); let plane = size * size; let patch_area = ps * ps; let mut out = vec![0.0f32; config.num_patches() * config.patch_dim()]; for gy in 0..grid { for gx in 0..grid { let patch = (gy * grid + gx) * config.patch_dim(); for c in 0..3 { for ky in 0..ps { let src_row = c * plane + (gy * ps + ky) * size + gx * ps; let dst_row = patch + c * patch_area + ky * ps; out[dst_row..dst_row + ps].copy_from_slice(&pixels[src_row..src_row + ps]); } } } } out } /// Flatten interleaved 8-bit RGB into patches, rescaling to `[0, 1]` and /// applying the ImageNet normalization on the way. /// /// `rgb` is `[image_size, image_size, 3]` — the layout image decoders and /// camera conversions naturally produce. No resizing happens here; the /// caller supplies an image already at `config.image_size`. pub fn patches_from_rgb8(rgb: &[u8], config: &Config) -> Vec { let size = config.image_size; let ps = config.patch_size; let grid = config.grid(); assert_eq!( rgb.len(), 3 * size * size, "expected a [{size}, {size}, 3] RGB image, got {} bytes", rgb.len() ); let patch_area = ps * ps; let mut out = vec![0.0f32; config.num_patches() * config.patch_dim()]; for gy in 0..grid { for gx in 0..grid { let patch = (gy * grid + gx) * config.patch_dim(); for ky in 0..ps { let y = gy * ps + ky; for kx in 0..ps { let x = gx * ps + kx; let src = (y * size + x) * 3; for c in 0..3 { let v = rgb[src + c] as f32 / 255.0; out[patch + c * patch_area + ky * ps + kx] = (v - IMAGE_MEAN[c]) / IMAGE_STD[c]; } } } } } out } /// Reshape a `[out, 3, patch, patch]` Conv2d weight into the /// `[patch_dim, out]` matrix the graph's patch-embedding matmul expects. /// /// The source is already contiguous in channel-major order per output /// channel, so this is purely a transpose of a `[out, patch_dim]` view. pub fn conv_weight_to_matmul(weight: &[f32], out_channels: usize, patch_dim: usize) -> Vec { assert_eq!( weight.len(), out_channels * patch_dim, "conv weight has {} values, expected {out_channels} * {patch_dim}", weight.len() ); let mut m = vec![0.0f32; patch_dim * out_channels]; for o in 0..out_channels { for i in 0..patch_dim { m[i * out_channels + o] = weight[o * patch_dim + i]; } } m } #[cfg(test)] mod tests { use super::*; #[test] fn chw_patch_layout_is_channel_major() { let c = Config::vits16(); // Encode each pixel's identity as its flat CHW index so the // mapping is checkable by arithmetic. let size = c.image_size; let pixels: Vec = (0..3 * size * size).map(|i| i as f32).collect(); let patches = patches_from_pixels_chw(&pixels, &c); let ps = c.patch_size; let plane = size * size; // Patch (gy=3, gx=5), channel 2, offset (ky=7, kx=11). let (gy, gx, ch, ky, kx) = (3, 5, 2, 7, 11); let got = patches[(gy * c.grid() + gx) * c.patch_dim() + ch * ps * ps + ky * ps + kx]; let want = (ch * plane + (gy * ps + ky) * size + gx * ps + kx) as f32; assert_eq!(got, want); } #[test] fn rgb8_and_chw_paths_agree() { let c = Config::vits16(); let size = c.image_size; // Build an arbitrary but reproducible RGB image, then the // equivalent normalized CHW tensor, and check both flatteners // land on the same patch tensor. let rgb: Vec = (0..3 * size * size).map(|i| (i % 251) as u8).collect(); let mut chw = vec![0.0f32; 3 * size * size]; for y in 0..size { for x in 0..size { for ch in 0..3 { let v = rgb[(y * size + x) * 3 + ch] as f32 / 255.0; chw[ch * size * size + y * size + x] = (v - IMAGE_MEAN[ch]) / IMAGE_STD[ch]; } } } let from_rgb = patches_from_rgb8(&rgb, &c); let from_chw = patches_from_pixels_chw(&chw, &c); assert_eq!(from_rgb.len(), from_chw.len()); let worst = from_rgb .iter() .zip(&from_chw) .map(|(a, b)| (a - b).abs()) .fold(0.0f32, f32::max); assert!(worst < 1e-6, "paths disagree by {worst}"); } #[test] fn normalization_maps_midgray_near_zero() { let c = Config::vits16(); // 0.485*255 ≈ 124 is the red-channel mean, so red lands near 0. let rgb = vec![124u8; 3 * c.image_size * c.image_size]; let patches = patches_from_rgb8(&rgb, &c); assert!(patches[0].abs() < 0.02, "red channel not centred: {}", patches[0]); } #[test] fn conv_weight_transpose_roundtrip() { // [out=2, patch_dim=3] stored row-major becomes [3, 2]. let w = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let m = conv_weight_to_matmul(&w, 2, 3); assert_eq!(m, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]); } }