File size: 5,790 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
//! Do the features actually mean anything?
//!
//! The point of DINO features is that patches of the same thing land near
//! each other in feature space regardless of position, and patches of
//! different things do not. That property is what the PCA colouring
//! displays and what any downstream use would rely on, and it is
//! surprisingly sensitive to exactly the mistakes that are easy to make
//! here: a transposed projection, a mis-paired RoPE half, or a patch
//! flattening in the wrong channel order all leave the magnitudes looking
//! healthy while destroying the structure.
//!
//! So rather than compare against a reference dump — which would need
//! torch on the machine — this checks the property directly, using a scene
//! with known regions.
//!
//! Requires the checkpoint. Set `DINOVISION_WEIGHTS` to a
//! `model.safetensors`; without it the test reports itself skipped, since
//! failing for a missing asset would be noise.

use dinovision::dinov3::Config;

/// Cosine similarity, the metric DINO features are usually compared under.
fn cosine(a: &[f32], b: &[f32]) -> f32 {
    let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
    let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
    let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
    dot / (na * nb).max(f32::MIN_POSITIVE)
}

/// A scene with four unambiguous regions, drawn as flat colour blocks so
/// the expected grouping is not a matter of opinion. Quadrants, at 224²:
/// top-left red, top-right green, bottom-left blue, bottom-right a fine
/// checkerboard (same mean colour as the background, different texture).
fn quadrant_scene(size: usize) -> Vec<u8> {
    let mut rgb = vec![0u8; size * size * 3];
    let half = size / 2;
    for y in 0..size {
        for x in 0..size {
            let px = match (x < half, y < half) {
                (true, true) => [200, 40, 40],
                (false, true) => [40, 180, 60],
                (true, false) => [50, 70, 210],
                (false, false) => {
                    let on = ((x / 4) + (y / 4)) % 2 == 0;
                    if on { [230, 230, 230] } else { [30, 30, 30] }
                }
            };
            let i = (y * size + x) * 3;
            rgb[i..i + 3].copy_from_slice(&px);
        }
    }
    rgb
}

/// Index of the patch at grid position (gx, gy), skipping prefix tokens.
fn patch(features: &[f32], config: &Config, gx: usize, gy: usize) -> Vec<f32> {
    let t = config.num_prefix_tokens() + gy * config.grid() + gx;
    let h = config.hidden_size;
    features[t * h..(t + 1) * h].to_vec()
}

#[test]
fn features_group_by_content_not_position() {
    let Ok(weights) = std::env::var("DINOVISION_WEIGHTS") else {
        eprintln!("skipped: set DINOVISION_WEIGHTS to a model.safetensors to run this");
        return;
    };
    let weights = std::path::PathBuf::from(weights);
    if !weights.exists() {
        eprintln!("skipped: {} does not exist", weights.display());
        return;
    }

    let config = Config::vits16().at_resolution(224);
    let gpu = dinovision::init_context(None).expect("GPU context");
    let (mut session, _) = dinovision::bench::build_encoder_session(gpu, &config, None);
    let model = meganeura::data::safetensors::SafeTensorsModel::load(weights).expect("read weights");
    dinovision::weights::load_encoder(&mut session, &model, &config).expect("bind weights");

    let rgb = quadrant_scene(config.image_size);
    let patches = dinovision::preprocess::patches_from_rgb8(&rgb, &config);
    session.set_input("patches", &patches);
    session.step();
    session.wait();
    let features = session.read_output(config.num_tokens() * config.hidden_size);
    assert!(
        features.iter().all(|v| v.is_finite()),
        "features contain non-finite values"
    );

    // Sample two well-separated patches inside each quadrant, staying away
    // from the boundaries where receptive fields mix regions.
    let g = config.grid();
    let q = g / 4;
    let regions = [
        ("red", [(q, q), (q + 1, q + 1)]),
        ("green", [(g - q - 1, q), (g - q - 2, q + 1)]),
        ("blue", [(q, g - q - 1), (q + 1, g - q - 2)]),
        ("checker", [(g - q - 1, g - q - 1), (g - q - 2, g - q - 2)]),
    ];

    // Within a region, two patches of the same material should be close.
    let mut worst_within = 1.0f32;
    for (name, pts) in &regions {
        let a = patch(&features, &config, pts[0].0, pts[0].1);
        let b = patch(&features, &config, pts[1].0, pts[1].1);
        let c = cosine(&a, &b);
        eprintln!("within {name:>8}: {c:.3}");
        worst_within = worst_within.min(c);
    }

    // Across regions, patches of different material should be further apart
    // than any same-material pair.
    let mut best_across = -1.0f32;
    for i in 0..regions.len() {
        for j in (i + 1)..regions.len() {
            let a = patch(&features, &config, regions[i].1[0].0, regions[i].1[0].1);
            let b = patch(&features, &config, regions[j].1[0].0, regions[j].1[0].1);
            let c = cosine(&a, &b);
            eprintln!("across {:>8}/{:<8}: {c:.3}", regions[i].0, regions[j].0);
            best_across = best_across.max(c);
        }
    }

    eprintln!("worst within-region {worst_within:.3}, best across-region {best_across:.3}");
    assert!(
        worst_within > best_across,
        "features do not separate content: the least similar same-region pair \

         ({worst_within:.3}) scored below the most similar different-region pair \

         ({best_across:.3}). A scrambled patch order or transposed projection \

         looks exactly like this."
    );
}