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