|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| use crate::dinov3::Config;
|
|
|
|
|
| pub const IMAGE_MEAN: [f32; 3] = [0.485, 0.456, 0.406];
|
| pub const IMAGE_STD: [f32; 3] = [0.229, 0.224, 0.225];
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn patches_from_pixels_chw(pixels: &[f32], config: &Config) -> Vec<f32> {
|
| 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
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn patches_from_rgb8(rgb: &[u8], config: &Config) -> Vec<f32> {
|
| 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
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn conv_weight_to_matmul(weight: &[f32], out_channels: usize, patch_dim: usize) -> Vec<f32> {
|
| 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();
|
|
|
|
|
| let size = c.image_size;
|
| let pixels: Vec<f32> = (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;
|
|
|
| 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;
|
|
|
|
|
|
|
| let rgb: Vec<u8> = (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();
|
|
|
| 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() {
|
|
|
| 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]);
|
| }
|
| }
|
|
|