//! 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, flag: &str) -> Result { args.next() .ok_or_else(|| format!("{flag} requires a value")) } impl Args { fn parse() -> Result { 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, } fn distribution(values: impl IntoIterator) -> Distribution { let mut values: Vec = 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::() / 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 { 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> { 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 = (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 = (0..3 * size * size) .map(|i| ((i * 37 + 11) % 251) as f32 / 250.0) .collect(); let b: Vec = (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); } }