mad-bot's picture
Publish verified DinoVision case-study artifacts
3ad454b verified
Raw
History Blame Contribute Delete
19.3 kB
//! Evaluate a trained decoder on an immutable held-out dataset manifest.
//!
//! Unlike `train_decoder`, this program never updates parameters and never
//! samples from the training cache. It verifies every input hash and writes
//! per-image metrics so aggregate claims can be regenerated from raw data.
//!
//! ```text
//! cargo run --release --example evaluate_decoder -- \
//! --manifest experiments/dataset.json \
//! --model model.safetensors --decoder decoder.bin \
//! --split test --layers 3 --output artifacts/quality.json
//! ```
use std::path::{Path, PathBuf};
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use dinovision::decoder;
use dinovision::dinov3::Config;
use meganeura::train::{Mode, SessionConfig};
use serde::Serialize;
mod common;
#[derive(Debug)]
struct Args {
manifest: PathBuf,
model: PathBuf,
decoder: PathBuf,
split: String,
layers: usize,
size: usize,
limit: usize,
output: PathBuf,
samples: usize,
}
fn require_value(args: &mut impl Iterator<Item = String>, flag: &str) -> Result<String, String> {
args.next()
.ok_or_else(|| format!("{flag} requires a value"))
}
impl Args {
fn parse() -> Result<Self, String> {
let mut manifest = None;
let mut model = None;
let mut decoder = None;
let mut split = "test".to_string();
let mut layers = 3usize;
let mut size = 224usize;
let mut limit = usize::MAX;
let mut output = PathBuf::from("artifacts/quality.json");
let mut samples = 12usize;
let mut args = std::env::args().skip(1);
while let Some(flag) = args.next() {
match flag.as_str() {
"--manifest" => manifest = Some(PathBuf::from(require_value(&mut args, &flag)?)),
"--model" => model = Some(PathBuf::from(require_value(&mut args, &flag)?)),
"--decoder" => decoder = Some(PathBuf::from(require_value(&mut args, &flag)?)),
"--split" => split = require_value(&mut args, &flag)?,
"--layers" => {
layers = require_value(&mut args, &flag)?
.parse()
.map_err(|_| "invalid --layers")?
}
"--size" => {
size = require_value(&mut args, &flag)?
.parse()
.map_err(|_| "invalid --size")?
}
"--limit" => {
limit = require_value(&mut args, &flag)?
.parse()
.map_err(|_| "invalid --limit")?
}
"--output" => output = PathBuf::from(require_value(&mut args, &flag)?),
"--samples" => {
samples = require_value(&mut args, &flag)?
.parse()
.map_err(|_| "invalid --samples")?
}
"-h" | "--help" => {
return Err(
"usage: evaluate_decoder --manifest FILE --model FILE --decoder FILE \
[--split test] [--layers 3] [--size 224] [--limit N] \
[--output FILE] [--samples N]"
.to_string(),
);
}
_ => return Err(format!("unknown argument {flag:?}")),
}
}
Ok(Self {
manifest: manifest.ok_or("--manifest is required")?,
model: model.ok_or("--model is required")?,
decoder: decoder.ok_or("--decoder is required")?,
split,
layers,
size,
limit,
output,
samples,
})
}
}
#[derive(Serialize)]
struct PerImage {
path: String,
source: String,
group: String,
sha256: String,
pixels: usize,
mse: f64,
mae: f64,
psnr_db: f64,
ssim: f64,
}
#[derive(Serialize)]
struct Distribution {
mean: f64,
median: f64,
p25: f64,
p75: f64,
min: f64,
max: f64,
}
#[derive(Serialize)]
struct Summary {
count: usize,
global_psnr_db: f64,
psnr_db: Distribution,
ssim: Distribution,
mae: Distribution,
}
#[derive(Serialize)]
struct Evaluation {
schema_version: u32,
created_unix_seconds: u64,
dataset_name: String,
manifest: String,
manifest_sha256: String,
split: String,
image_size: usize,
encoder_layers: usize,
model: String,
model_sha256: String,
decoder: String,
decoder_sha256: String,
elapsed_seconds: f64,
summary: Summary,
images: Vec<PerImage>,
}
fn distribution(values: impl IntoIterator<Item = f64>) -> Distribution {
let mut values: Vec<f64> = values.into_iter().collect();
values.sort_by(f64::total_cmp);
let n = values.len();
assert!(n > 0);
let quantile = |q: f64| {
let at = q * (n - 1) as f64;
let lo = at.floor() as usize;
let hi = at.ceil() as usize;
values[lo] + (values[hi] - values[lo]) * (at - lo as f64)
};
Distribution {
mean: values.iter().sum::<f64>() / n as f64,
median: quantile(0.5),
p25: quantile(0.25),
p75: quantile(0.75),
min: values[0],
max: values[n - 1],
}
}
fn target_chw(rgb: &[u8], size: usize) -> Vec<f32> {
let pixels = size * size;
let mut target = vec![0.0f32; 3 * pixels];
for c in 0..3 {
for p in 0..pixels {
target[c * pixels + p] = rgb[p * 3 + c] as f32 / 255.0;
}
}
target
}
fn metrics(recon: &[f32], target: &[f32], size: usize) -> (f64, f64, f64, f64) {
let mut squared = 0.0f64;
let mut absolute = 0.0f64;
for (&a, &b) in recon.iter().zip(target) {
let d = (a - b) as f64;
squared += d * d;
absolute += d.abs();
}
let mse = squared / recon.len() as f64;
let mae = absolute / recon.len() as f64;
let psnr = 10.0 * (1.0 / mse).log10();
(mse, mae, psnr, ssim_rgb(recon, target, size))
}
/// RGB SSIM with the conventional 11x11 Gaussian window (sigma 1.5),
/// computed independently per channel over the valid image region.
fn ssim_rgb(a: &[f32], b: &[f32], size: usize) -> f64 {
const RADIUS: usize = 5;
const SIGMA: f64 = 1.5;
const C1: f64 = 0.01 * 0.01;
const C2: f64 = 0.03 * 0.03;
assert!(size > 2 * RADIUS);
let mut kernel = [0.0f64; 2 * RADIUS + 1];
let mut kernel_sum = 0.0;
for (index, weight) in kernel.iter_mut().enumerate() {
let x = index as isize - RADIUS as isize;
*weight = (-(x * x) as f64 / (2.0 * SIGMA * SIGMA)).exp();
kernel_sum += *weight;
}
for w in &mut kernel {
*w /= kernel_sum;
}
let plane = size * size;
let valid = size - 2 * RADIUS;
let mut total = 0.0;
let mut count = 0usize;
for c in 0..3 {
// Five Gaussian-filtered moments, stored together to reuse both the
// input reads and kernel coefficients. The separable implementation
// is mathematically equivalent to the 11x11 2-D Gaussian window.
let mut horizontal = vec![[0.0f64; 5]; size * valid];
for y in 0..size {
for x in 0..valid {
let mut moments = [0.0f64; 5];
for (kx, &weight) in kernel.iter().enumerate() {
let index = c * plane + y * size + x + kx;
let va = a[index] as f64;
let vb = b[index] as f64;
moments[0] += weight * va;
moments[1] += weight * vb;
moments[2] += weight * va * va;
moments[3] += weight * vb * vb;
moments[4] += weight * va * vb;
}
horizontal[y * valid + x] = moments;
}
}
for y in 0..valid {
for x in 0..valid {
let mut moments = [0.0f64; 5];
for (ky, &weight) in kernel.iter().enumerate() {
for i in 0..5 {
moments[i] += weight * horizontal[(y + ky) * valid + x][i];
}
}
let [mean_a, mean_b, aa, bb, ab] = moments;
let var_a = (aa - mean_a * mean_a).max(0.0);
let var_b = (bb - mean_b * mean_b).max(0.0);
let covariance = ab - mean_a * mean_b;
total += ((2.0 * mean_a * mean_b + C1) * (2.0 * covariance + C2))
/ ((mean_a * mean_a + mean_b * mean_b + C1) * (var_a + var_b + C2));
count += 1;
}
}
}
total / count as f64
}
fn save_sample(
path: &Path,
rgb: &[u8],
patches: &[f32],
encoder_output: &[f32],
features: &[f32],
recon: &[f32],
size: usize,
) {
let mut image = image::RgbImage::new((2 * size) as u32, size as u32);
let pixels = size * size;
for y in 0..size {
for x in 0..size {
let p = y * size + x;
image.put_pixel(
x as u32,
y as u32,
image::Rgb([rgb[p * 3], rgb[p * 3 + 1], rgb[p * 3 + 2]]),
);
let value = |c: usize| (recon[c * pixels + p].clamp(0.0, 1.0) * 255.0) as u8;
image.put_pixel(
(size + x) as u32,
y as u32,
image::Rgb([value(0), value(1), value(2)]),
);
}
}
image.save(path).expect("write reconstruction sample");
let stem = path
.file_stem()
.expect("sample path has a stem")
.to_string_lossy();
let parent = path.parent().unwrap_or_else(|| Path::new("."));
std::fs::write(
parent.join(format!("{stem}-patches.f32")),
bytemuck::cast_slice(patches),
)
.expect("write preprocessed patches");
std::fs::write(
parent.join(format!("{stem}-encoder.f32")),
bytemuck::cast_slice(encoder_output),
)
.expect("write raw encoder output");
std::fs::write(
parent.join(format!("{stem}-features.f32")),
bytemuck::cast_slice(features),
)
.expect("write raw encoder features");
std::fs::write(
parent.join(format!("{stem}-reconstruction.f32")),
bytemuck::cast_slice(recon),
)
.expect("write raw reconstruction sample");
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
env_logger::Builder::from_env(env_logger::Env::default().default_filter_or("info")).init();
let args = Args::parse().map_err(|e| e.to_string())?;
let loaded = common::load_manifest(&args.manifest)?;
let dataset_name = loaded.manifest.name.clone();
let entries = common::images_for_split(&args.manifest, &args.split, args.limit)?;
log::info!(
"evaluating {} images from split {:?}",
entries.len(),
args.split
);
let config = Config::vits16()
.at_resolution(args.size)
.with_layers(args.layers);
let gpu = dinovision::init_context(None).expect("GPU context");
let (mut encoder, _) = dinovision::bench::build_encoder_session(gpu.clone(), &config, None);
let model = meganeura::data::safetensors::SafeTensorsModel::load(args.model.clone())?;
dinovision::weights::load_encoder(&mut encoder, &model, &config)?;
let feat_len = config.hidden_size * config.num_patches();
let img_len = 3 * config.image_size * config.image_size;
let mut graph = meganeura::Graph::new();
let features = graph.input("feat", &[feat_len]);
let reconstruction = decoder::build_decoder(&mut graph, &config, features, 1);
graph.set_outputs(vec![reconstruction]);
let (mut decoder_session, _) = meganeura::train::build(
&graph,
SessionConfig {
mode: Mode::Inference,
gpu: Some(gpu),
..Default::default()
},
);
decoder::load_parameters(&mut decoder_session, &graph, &args.decoder)?;
let sample_dir = args
.output
.parent()
.unwrap_or_else(|| Path::new("."))
.join("quality_samples");
if args.samples > 0 {
std::fs::create_dir_all(&sample_dir)?;
}
if let Some(parent) = args.output.parent()
&& !parent.as_os_str().is_empty()
{
std::fs::create_dir_all(parent)?;
}
let start = Instant::now();
let total_entries = entries.len();
let mut encoder_out = vec![0.0f32; config.num_tokens() * config.hidden_size];
let mut recon = vec![0.0f32; img_len];
let mut results = Vec::new();
let mut total_squared_error = 0.0f64;
let mut total_values = 0usize;
for (index, (path, entry)) in entries.into_iter().enumerate() {
common::verify_image(&path, &entry)?;
let rgb = common::load_frame(&path, config.image_size as u32)
.ok_or_else(|| format!("failed to load {}", path.display()))?;
let target = target_chw(&rgb, config.image_size);
let patches = dinovision::preprocess::patches_from_rgb8(&rgb, &config);
encoder.set_input("patches", &patches);
encoder.step();
encoder.wait();
encoder.read_output_by_index(0, &mut encoder_out);
let features = decoder::patch_features_to_nchw(&encoder_out, &config);
decoder_session.set_input("feat", &features);
decoder_session.step();
decoder_session.wait();
decoder_session.read_output_by_index(0, &mut recon);
let (mse, mae, psnr_db, ssim) = metrics(&recon, &target, config.image_size);
total_squared_error += mse * recon.len() as f64;
total_values += recon.len();
results.push(PerImage {
path: entry.path.to_string_lossy().replace('\\', "/"),
source: entry.source,
group: entry.group,
sha256: entry.sha256,
pixels: config.image_size * config.image_size,
mse,
mae,
psnr_db,
ssim,
});
if index < args.samples {
save_sample(
&sample_dir.join(format!("{index:04}.png")),
&rgb,
&patches,
&encoder_out,
&features,
&recon,
config.image_size,
);
}
log::info!(
"{}/{} PSNR {:.2} dB SSIM {:.4} {}",
index + 1,
total_entries,
psnr_db,
ssim,
path.display()
);
}
let global_mse = total_squared_error / total_values as f64;
let summary = Summary {
count: results.len(),
global_psnr_db: 10.0 * (1.0 / global_mse).log10(),
psnr_db: distribution(results.iter().map(|x| x.psnr_db)),
ssim: distribution(results.iter().map(|x| x.ssim)),
mae: distribution(results.iter().map(|x| x.mae)),
};
let evaluation = Evaluation {
schema_version: 1,
created_unix_seconds: SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(),
dataset_name,
manifest: args.manifest.to_string_lossy().replace('\\', "/"),
manifest_sha256: common::sha256(&args.manifest)?,
split: args.split,
image_size: config.image_size,
encoder_layers: config.num_hidden_layers,
model: args.model.to_string_lossy().replace('\\', "/"),
model_sha256: common::sha256(&args.model)?,
decoder: args.decoder.to_string_lossy().replace('\\', "/"),
decoder_sha256: common::sha256(&args.decoder)?,
elapsed_seconds: start.elapsed().as_secs_f64(),
summary,
images: results,
};
std::fs::write(&args.output, serde_json::to_vec_pretty(&evaluation)?)?;
println!(
"wrote {}: {} images, global PSNR {:.2} dB, median SSIM {:.4}",
args.output.display(),
evaluation.summary.count,
evaluation.summary.global_psnr_db,
evaluation.summary.ssim.median
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn reference_ssim(a: &[f32], b: &[f32], size: usize) -> f64 {
const RADIUS: isize = 5;
const SIGMA: f64 = 1.5;
const C1: f64 = 0.01 * 0.01;
const C2: f64 = 0.03 * 0.03;
let mut kernel = Vec::new();
let mut sum = 0.0;
for y in -RADIUS..=RADIUS {
for x in -RADIUS..=RADIUS {
let weight = (-((x * x + y * y) as f64) / (2.0 * SIGMA * SIGMA)).exp();
kernel.push(weight);
sum += weight;
}
}
kernel.iter_mut().for_each(|weight| *weight /= sum);
let plane = size * size;
let mut total = 0.0;
let mut count = 0;
for c in 0..3 {
for y in RADIUS as usize..size - RADIUS as usize {
for x in RADIUS as usize..size - RADIUS as usize {
let mut moments = [0.0f64; 5];
let mut wi = 0;
for ky in -RADIUS..=RADIUS {
for kx in -RADIUS..=RADIUS {
let index = c * plane
+ (y as isize + ky) as usize * size
+ (x as isize + kx) as usize;
let va = a[index] as f64;
let vb = b[index] as f64;
let weight = kernel[wi];
moments[0] += weight * va;
moments[1] += weight * vb;
moments[2] += weight * va * va;
moments[3] += weight * vb * vb;
moments[4] += weight * va * vb;
wi += 1;
}
}
let [mean_a, mean_b, aa, bb, ab] = moments;
let var_a = (aa - mean_a * mean_a).max(0.0);
let var_b = (bb - mean_b * mean_b).max(0.0);
let covariance = ab - mean_a * mean_b;
total += ((2.0 * mean_a * mean_b + C1) * (2.0 * covariance + C2))
/ ((mean_a * mean_a + mean_b * mean_b + C1) * (var_a + var_b + C2));
count += 1;
}
}
}
total / count as f64
}
#[test]
fn identical_images_have_unit_ssim() {
let size = 16;
let image: Vec<f32> = (0..3 * size * size)
.map(|i| (i % 251) as f32 / 250.0)
.collect();
assert!((ssim_rgb(&image, &image, size) - 1.0).abs() < 1e-10);
}
#[test]
fn separable_ssim_matches_direct_window() {
let size = 16;
let a: Vec<f32> = (0..3 * size * size)
.map(|i| ((i * 37 + 11) % 251) as f32 / 250.0)
.collect();
let b: Vec<f32> = (0..3 * size * size)
.map(|i| ((i * 19 + 7) % 241) as f32 / 240.0)
.collect();
let expected = reference_ssim(&a, &b, size);
let actual = ssim_rgb(&a, &b, size);
assert!((actual - expected).abs() < 1e-12, "{actual} vs {expected}");
}
#[test]
fn distribution_interpolates_quartiles() {
let d = distribution([1.0, 2.0, 3.0, 4.0]);
assert_eq!(d.median, 2.5);
assert_eq!(d.p25, 1.75);
assert_eq!(d.p75, 3.25);
}
}