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]);
    }
}