File size: 7,193 Bytes
eae424a | 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 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 | //! 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<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
}
/// 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<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
}
/// 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<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();
// 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<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;
// 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<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();
// 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]);
}
}
|