use crate::Error;
use crate::analyze::statistics::StatisticsChannel;
use crate::image::RasterImage;
use crate::pixel::HomogeneousPixel;
use super::PeakValue;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ChannelSquaredError {
pub count: u64,
pub nan_count: u64,
sum_squared_error: f64,
max_absolute_error: f64,
}
impl ChannelSquaredError {
#[inline]
pub(crate) fn empty() -> Self {
Self {
count: 0,
nan_count: 0,
sum_squared_error: 0.0,
max_absolute_error: 0.0,
}
}
#[inline]
pub(crate) fn push<C: StatisticsChannel>(&mut self, a: C, b: C) {
let difference = a.to_f64() - b.to_f64();
if difference.is_nan() {
self.nan_count += 1;
return;
}
self.count += 1;
self.sum_squared_error += difference * difference;
let magnitude = difference.abs();
if magnitude > self.max_absolute_error {
self.max_absolute_error = magnitude;
}
}
#[inline]
fn merge(&mut self, other: &Self) {
self.count += other.count;
self.nan_count += other.nan_count;
self.sum_squared_error += other.sum_squared_error;
if other.max_absolute_error > self.max_absolute_error {
self.max_absolute_error = other.max_absolute_error;
}
}
#[doc(alias = "SSE")]
#[must_use]
pub fn sum_squared_error(&self) -> f64 {
self.sum_squared_error
}
#[doc(alias = "mse")]
#[doc(alias = "MSE")]
#[must_use]
pub fn mean_squared_error(&self) -> Option<f64> {
(self.count > 0).then(|| self.sum_squared_error / self.count as f64)
}
#[doc(alias = "rmse")]
#[doc(alias = "RMSE")]
#[must_use]
pub fn root_mean_squared_error(&self) -> Option<f64> {
self.mean_squared_error().map(f64::sqrt)
}
#[must_use]
pub fn max_absolute_error(&self) -> Option<f64> {
(self.count > 0).then_some(self.max_absolute_error)
}
#[doc(alias = "psnr")]
#[doc(alias = "PSNR")]
#[must_use]
pub fn peak_signal_to_noise_ratio(&self, peak: PeakValue) -> Option<f64> {
let mse = self.mean_squared_error()?;
if mse == 0.0 {
return Some(f64::INFINITY);
}
let peak = peak.get();
Some(10.0 * (peak * peak / mse).log10())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct SquaredError {
channels: Vec<ChannelSquaredError>,
}
impl SquaredError {
#[must_use]
pub fn channel_count(&self) -> usize {
self.channels.len()
}
#[must_use]
pub fn channel(&self, index: usize) -> Option<ChannelSquaredError> {
self.channels.get(index).copied()
}
#[must_use]
pub fn channels(&self) -> &[ChannelSquaredError] {
&self.channels
}
#[must_use]
pub fn pooled(&self) -> ChannelSquaredError {
let mut total = ChannelSquaredError::empty();
for channel in &self.channels {
total.merge(channel);
}
total
}
}
pub fn squared_error<A, B, P>(a: &A, b: &B) -> Result<SquaredError, Error>
where
A: RasterImage<Pixel = P>,
B: RasterImage<Pixel = P>,
P: HomogeneousPixel,
P::Channel: StatisticsChannel,
{
if a.size() != b.size() {
return Err(Error::SizeMismatch {
expected: a.size(),
actual: b.size(),
});
}
let mut channels = vec![ChannelSquaredError::empty(); P::CHANNEL_COUNT];
for y in 0..a.height() {
for (pa, pb) in a.row(y).iter().zip(b.row(y).iter()) {
for (index, accumulator) in channels.iter_mut().enumerate() {
accumulator.push(pa.channel(index), pb.channel(index));
}
}
}
Ok(SquaredError { channels })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::image::{Image, SubView};
use crate::peak;
use crate::pixel::{Mono8, Mono16, MonoF32, MonoF64, Rgb8};
use crate::{Coordinate, Rectangle, Size};
#[test]
fn an_identical_pair_has_no_error_and_infinite_psnr() {
let image = Image::generate(16, 16, |x, y| Mono8::new((x * y) as u8));
let error = squared_error(&image, &image).unwrap().pooled();
assert_eq!(error.count, 256);
assert_eq!(error.nan_count, 0);
assert_eq!(error.sum_squared_error(), 0.0);
assert_eq!(error.mean_squared_error(), Some(0.0));
assert_eq!(error.root_mean_squared_error(), Some(0.0));
assert_eq!(error.max_absolute_error(), Some(0.0));
let psnr = error
.peak_signal_to_noise_ratio(PeakValue::of_pixel::<Mono8>())
.unwrap();
assert!(psnr.is_infinite() && psnr > 0.0, "{psnr}");
assert!(!psnr.is_finite());
assert!(psnr > 40.0);
}
#[test]
fn a_known_offset_gives_the_textbook_mse_and_psnr() {
let a = Image::fill(10, 10, Mono8::new(100));
let b = Image::fill(10, 10, Mono8::new(110));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.mean_squared_error(), Some(100.0));
assert_eq!(error.root_mean_squared_error(), Some(10.0));
assert_eq!(error.max_absolute_error(), Some(10.0));
let expected = 10.0 * (255.0f64 * 255.0 / 100.0).log10();
let psnr = error
.peak_signal_to_noise_ratio(PeakValue::of_pixel::<Mono8>())
.unwrap();
assert!((psnr - expected).abs() < 1e-12, "{psnr} vs {expected}");
assert!((psnr - 28.13).abs() < 0.01, "{psnr}");
}
#[test]
fn the_error_is_symmetric_in_its_arguments() {
let a = Image::generate(9, 7, |x, _| Mono8::new((x * 20) as u8));
let b = Image::generate(9, 7, |_, y| Mono8::new((y * 30) as u8));
let forward = squared_error(&a, &b).unwrap().pooled();
let backward = squared_error(&b, &a).unwrap().pooled();
assert_eq!(forward, backward);
}
#[test]
fn each_channel_is_measured_independently() {
let a = Image::fill(4, 4, Rgb8::new(50, 50, 50));
let b = Image::fill(4, 4, Rgb8::new(51, 53, 57));
let error = squared_error(&a, &b).unwrap();
assert_eq!(error.channel_count(), 3);
assert_eq!(error.channel(0).unwrap().mean_squared_error(), Some(1.0));
assert_eq!(error.channel(1).unwrap().mean_squared_error(), Some(9.0));
assert_eq!(error.channel(2).unwrap().mean_squared_error(), Some(49.0));
assert_eq!(error.channels().len(), 3);
assert_eq!(error.channel(3), None);
}
#[test]
fn pooling_sums_accumulators_rather_than_averaging_results() {
let a = Image::fill(4, 4, Rgb8::new(50, 50, 50));
let b = Image::fill(4, 4, Rgb8::new(51, 53, 57));
let error = squared_error(&a, &b).unwrap();
let pooled = error.pooled();
assert_eq!(pooled.count, 48);
assert_eq!(pooled.sum_squared_error(), 16.0 * (1.0 + 9.0 + 49.0));
assert_eq!(pooled.mean_squared_error(), Some((1.0 + 9.0 + 49.0) / 3.0));
assert_eq!(pooled.max_absolute_error(), Some(7.0));
let peak = PeakValue::of_pixel::<Rgb8>();
let pooled_psnr = pooled.peak_signal_to_noise_ratio(peak).unwrap();
let mean_of_psnrs = error
.channels()
.iter()
.map(|c| c.peak_signal_to_noise_ratio(peak).unwrap())
.sum::<f64>()
/ 3.0;
assert!(
(pooled_psnr - mean_of_psnrs).abs() > 1.0,
"{pooled_psnr} vs {mean_of_psnrs}"
);
}
#[test]
fn pooling_a_single_channel_image_is_that_channel() {
let a = Image::fill(3, 3, Mono8::new(7));
let b = Image::fill(3, 3, Mono8::new(9));
let error = squared_error(&a, &b).unwrap();
assert_eq!(error.pooled(), error.channel(0).unwrap());
}
#[test]
fn a_size_mismatch_is_a_tier_two_error() {
let a: Image<Mono8> = Image::fill(4, 4, Mono8::new(0));
let b: Image<Mono8> = Image::fill(4, 5, Mono8::new(0));
assert_eq!(
squared_error(&a, &b),
Err(Error::SizeMismatch {
expected: Size::new(4, 4),
actual: Size::new(4, 5),
})
);
}
#[test]
fn an_empty_pair_reports_absence_not_zero() {
let a: Image<Mono8> = Image::generate(0, 0, |_, _| Mono8::new(0));
let b: Image<Mono8> = Image::generate(0, 0, |_, _| Mono8::new(0));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.count, 0);
assert_eq!(error.sum_squared_error(), 0.0);
assert_eq!(error.mean_squared_error(), None);
assert_eq!(error.root_mean_squared_error(), None);
assert_eq!(error.max_absolute_error(), None);
assert_eq!(
error.peak_signal_to_noise_ratio(PeakValue::of_pixel::<Mono8>()),
None
);
}
#[test]
fn a_nan_difference_is_counted_and_excluded() {
let a = Image::generate(4, 1, |x, _| {
MonoF32::new(if x == 2 { f32::NAN } else { 1.0 })
});
let b = Image::fill(4, 1, MonoF32::new(3.0));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.count, 3);
assert_eq!(error.nan_count, 1);
assert_eq!(error.mean_squared_error(), Some(4.0));
assert_eq!(error.max_absolute_error(), Some(2.0));
}
#[test]
fn an_all_nan_pair_reports_absence() {
let a = Image::fill(2, 2, MonoF64::new(f64::NAN));
let b = Image::fill(2, 2, MonoF64::new(0.0));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.count, 0);
assert_eq!(error.nan_count, 4);
assert_eq!(error.mean_squared_error(), None);
}
#[test]
fn two_equal_infinities_have_no_difference_and_are_excluded() {
let a = Image::fill(2, 1, MonoF64::new(f64::INFINITY));
let b = Image::fill(2, 1, MonoF64::new(f64::INFINITY));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.count, 0);
assert_eq!(error.nan_count, 2);
assert_eq!(error.mean_squared_error(), None);
}
#[test]
fn a_lone_infinite_difference_is_kept_and_reported_as_infinite() {
let a = Image::generate(2, 1, |x, _| {
MonoF64::new(if x == 0 { f64::INFINITY } else { 1.0 })
});
let b = Image::fill(2, 1, MonoF64::new(0.0));
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.count, 2);
assert_eq!(error.nan_count, 0);
assert_eq!(error.mean_squared_error(), Some(f64::INFINITY));
assert_eq!(error.max_absolute_error(), Some(f64::INFINITY));
let psnr = error.peak_signal_to_noise_ratio(peak!(1.0)).unwrap();
assert_eq!(psnr, f64::NEG_INFINITY);
}
#[test]
fn sixteen_bit_differences_keep_their_precision_in_the_sum() {
let a = Image::fill(1024, 1024, Mono16::new(65_535));
let b = Image::fill(1024, 1024, Mono16::new(0));
let error = squared_error(&a, &b).unwrap().pooled();
let expected_sum = 65_535.0f64 * 65_535.0 * 1024.0 * 1024.0;
assert_eq!(error.sum_squared_error(), expected_sum);
assert_eq!(error.mean_squared_error(), Some(65_535.0 * 65_535.0));
let psnr = error
.peak_signal_to_noise_ratio(PeakValue::of_pixel::<Mono16>())
.unwrap();
assert!(psnr.abs() < 1e-12, "{psnr}");
}
#[test]
fn the_max_absolute_error_survives_a_small_mean() {
let a: Image<Mono8> = Image::fill(64, 64, Mono8::new(0));
let mut b = a.clone();
{
use crate::image::ImageViewMut;
*b.pixel_at_mut(31, 31) = Mono8::new(255);
}
let error = squared_error(&a, &b).unwrap().pooled();
assert!(error.mean_squared_error().unwrap() < 16.0);
assert_eq!(error.max_absolute_error(), Some(255.0));
}
#[test]
fn a_region_of_view_compares_against_an_owned_image() {
let big = Image::generate(16, 16, |x, y| Mono8::new((x + y) as u8));
let tile = Image::generate(4, 4, |x, y| Mono8::new((x + 2 + y + 3) as u8));
let view = big
.roi(Rectangle::new(Coordinate::new(2, 3), Size::new(4, 4)))
.expect("in bounds");
let error = squared_error(&view, &tile).unwrap().pooled();
assert_eq!(error.count, 16);
assert_eq!(error.mean_squared_error(), Some(0.0));
}
#[test]
fn signed_channels_are_compared_too() {
let a: Image<i16> = Image::fill(2, 2, -100);
let b: Image<i16> = Image::fill(2, 2, 100);
let error = squared_error(&a, &b).unwrap().pooled();
assert_eq!(error.mean_squared_error(), Some(40_000.0));
assert_eq!(error.max_absolute_error(), Some(200.0));
}
}