use crate::image::{
Image, ImageView, ImageViewMut, MAX_RADIUS, RasterImage, RasterImageMut, gaussian_kernel_1d,
gaussian_kernel_size,
};
use crate::pixel::{MonoF64, SingleChannel};
use crate::transform::convolve_separable;
use crate::{Error, Sigma, sigma};
use crate::analyze::statistics::StatisticsChannel;
use crate::border::Skip;
use super::PeakValue;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct SsimParams {
peak: PeakValue,
sigma: Sigma,
k1: f64,
k2: f64,
}
impl SsimParams {
pub const REFERENCE_SIGMA: Sigma = sigma!(1.5);
pub const REFERENCE_K1: f64 = 0.01;
pub const REFERENCE_K2: f64 = 0.03;
pub const TRUNCATE: f32 = 3.0;
#[must_use]
pub const fn reference(peak: PeakValue) -> Self {
Self {
peak,
sigma: Self::REFERENCE_SIGMA,
k1: Self::REFERENCE_K1,
k2: Self::REFERENCE_K2,
}
}
pub fn try_new(peak: PeakValue, sigma: Sigma, k1: f64, k2: f64) -> Result<Self, Error> {
for (name, value) in [("k1", k1), ("k2", k2)] {
if !(value.is_finite() && value > 0.0) {
return Err(Error::InvalidParameter(format!(
"ssim {name} must be finite and positive, got {value}"
)));
}
}
let taps = gaussian_kernel_size(sigma, Self::TRUNCATE);
if taps < 3 {
return Err(Error::InvalidParameter(format!(
"ssim sigma {} derives a one-tap window, which has no variance \
and reduces the score to its luminance term",
sigma.get(),
)));
}
if taps > 2 * MAX_RADIUS + 1 {
return Err(Error::InvalidParameter(format!(
"ssim sigma {} needs a {taps}-tap window, above the {}-tap maximum",
sigma.get(),
2 * MAX_RADIUS + 1,
)));
}
Ok(Self {
peak,
sigma,
k1,
k2,
})
}
#[must_use]
pub const fn peak(&self) -> PeakValue {
self.peak
}
#[must_use]
pub const fn sigma(&self) -> Sigma {
self.sigma
}
#[must_use]
pub const fn k1(&self) -> f64 {
self.k1
}
#[must_use]
pub const fn k2(&self) -> f64 {
self.k2
}
#[must_use]
pub fn window_size(&self) -> usize {
gaussian_kernel_size(self.sigma, Self::TRUNCATE)
}
}
pub fn ssim<A, B, P>(a: &A, b: &B, params: SsimParams) -> Result<f64, Error>
where
A: RasterImage<Pixel = P>,
B: RasterImage<Pixel = P>,
P: SingleChannel,
P::Channel: StatisticsChannel,
{
let map = ssim_map(a, b, params)?;
let count = map.width() * map.height();
debug_assert!(count > 0, "ssim_map returned an empty map");
let mut total = 0.0;
for y in 0..map.height() {
for pixel in map.row(y) {
total += pixel.0;
}
}
Ok(total / count as f64)
}
pub fn ssim_map<A, B, P>(a: &A, b: &B, params: SsimParams) -> Result<Image<MonoF64>, Error>
where
A: RasterImage<Pixel = P>,
B: RasterImage<Pixel = P>,
P: SingleChannel,
P::Channel: StatisticsChannel,
{
if a.size() != b.size() {
return Err(Error::SizeMismatch {
expected: a.size(),
actual: b.size(),
});
}
let window = params.window_size();
if a.width() < window || a.height() < window {
return Err(Error::InvalidParameter(format!(
"ssim: a {}x{} window does not fit a {}x{} image; no position has the \
full window inside the frame",
window,
window,
a.width(),
a.height(),
)));
}
let offset_a = finite_mean(a);
let offset_b = finite_mean(b);
let mut plane_a = centered_plane(a, offset_a);
let mut plane_b = centered_plane(b, offset_b);
let mut plane_ab = Image::<MonoF64>::zero(a.width(), a.height());
for y in 0..plane_ab.height() {
for x in 0..plane_ab.width() {
let product = plane_a.pixel_at(x, y).0 * plane_b.pixel_at(x, y).0;
*plane_ab.pixel_at_mut(x, y) = MonoF64(product);
}
}
let kernel = gaussian_kernel_1d(params.sigma, SsimParams::TRUNCATE);
let mean_a: Image<MonoF64> = convolve_separable(&plane_a, &kernel, &Skip);
let mean_b: Image<MonoF64> = convolve_separable(&plane_b, &kernel, &Skip);
let moment_ab: Image<MonoF64> = convolve_separable(&plane_ab, &kernel, &Skip);
drop(plane_ab);
square_in_place(&mut plane_a);
let moment_aa: Image<MonoF64> = convolve_separable(&plane_a, &kernel, &Skip);
drop(plane_a);
square_in_place(&mut plane_b);
let moment_bb: Image<MonoF64> = convolve_separable(&plane_b, &kernel, &Skip);
drop(plane_b);
let peak = params.peak.get();
let c1 = (params.k1 * peak) * (params.k1 * peak);
let c2 = (params.k2 * peak) * (params.k2 * peak);
let mut map = Image::<MonoF64>::zero(mean_a.width(), mean_a.height());
for y in 0..map.height() {
for x in 0..map.width() {
let centered_mean_a = mean_a.pixel_at(x, y).0;
let centered_mean_b = mean_b.pixel_at(x, y).0;
let luminance_a = centered_mean_a + offset_a;
let luminance_b = centered_mean_b + offset_b;
let variance_a =
non_negative(moment_aa.pixel_at(x, y).0 - centered_mean_a * centered_mean_a);
let variance_b =
non_negative(moment_bb.pixel_at(x, y).0 - centered_mean_b * centered_mean_b);
let raw = moment_ab.pixel_at(x, y).0 - centered_mean_a * centered_mean_b;
let bound = (variance_a * variance_b).sqrt();
let covariance = if raw > bound {
bound
} else if raw < -bound {
-bound
} else {
raw
};
let numerator = (2.0 * luminance_a * luminance_b + c1) * (2.0 * covariance + c2);
let denominator = (luminance_a * luminance_a + luminance_b * luminance_b + c1)
* (variance_a + variance_b + c2);
*map.pixel_at_mut(x, y) = MonoF64(numerator / denominator);
}
}
Ok(map)
}
fn finite_mean<I, P>(image: &I) -> f64
where
I: RasterImage<Pixel = P>,
P: SingleChannel,
P::Channel: StatisticsChannel,
{
use crate::analyze::statistics::{ChannelStatistics, image_statistics};
let stats: ChannelStatistics<P::Channel> = image_statistics(image);
match stats.mean() {
Some(mean) if mean.is_finite() => mean,
_ => 0.0,
}
}
fn centered_plane<I, P>(image: &I, offset: f64) -> Image<MonoF64>
where
I: RasterImage<Pixel = P>,
P: SingleChannel,
P::Channel: StatisticsChannel,
{
let mut plane = Image::<MonoF64>::zero(image.width(), image.height());
for y in 0..image.height() {
let source = image.row(y);
let destination = plane.row_mut(y);
for (pixel, cell) in source.iter().zip(destination.iter_mut()) {
*cell = MonoF64(pixel.channel(0).to_f64() - offset);
}
}
plane
}
fn square_in_place(plane: &mut Image<MonoF64>) {
for y in 0..plane.height() {
for cell in plane.row_mut(y) {
cell.0 *= cell.0;
}
}
}
#[inline]
fn non_negative(value: f64) -> f64 {
if value < 0.0 { 0.0 } else { value }
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Size;
use crate::analyze::statistics::{ChannelStatistics, image_statistics};
use crate::image::ContiguousImage;
use crate::peak;
use crate::pixel::{Mono8, Mono16, MonoF32};
fn reference_params() -> SsimParams {
SsimParams::reference(PeakValue::of_pixel::<Mono8>())
}
fn checkerboard(size: usize, block: usize) -> Image<Mono8> {
Image::generate(size, size, |x, y| {
Mono8::new(if (x / block + y / block) % 2 == 0 {
220
} else {
30
})
})
}
#[test]
fn the_reference_parameters_are_the_published_eleven_tap_window() {
let params = reference_params();
assert_eq!(params.sigma(), sigma!(1.5));
assert_eq!(params.k1(), 0.01);
assert_eq!(params.k2(), 0.03);
assert_eq!(params.window_size(), 11);
assert_eq!(gaussian_kernel_size(sigma!(1.5), 4.0), 13);
}
#[test]
fn a_non_positive_stabilizing_constant_is_refused() {
let peak = PeakValue::of_pixel::<Mono8>();
assert!(SsimParams::try_new(peak, sigma!(1.5), 0.0, 0.03).is_err());
assert!(SsimParams::try_new(peak, sigma!(1.5), 0.01, -0.03).is_err());
assert!(SsimParams::try_new(peak, sigma!(1.5), f64::NAN, 0.03).is_err());
assert!(SsimParams::try_new(peak, sigma!(1.5), 0.01, 0.03).is_ok());
}
#[test]
fn a_sigma_that_derives_a_one_tap_window_is_refused() {
let peak = PeakValue::of_pixel::<Mono8>();
assert_eq!(gaussian_kernel_size(sigma!(0.1), SsimParams::TRUNCATE), 1);
let refused = SsimParams::try_new(peak, sigma!(0.1), 0.01, 0.03);
assert!(
matches!(refused, Err(Error::InvalidParameter(_))),
"{refused:?}"
);
let accepted = SsimParams::try_new(peak, sigma!(0.2), 0.01, 0.03).unwrap();
assert_eq!(accepted.window_size(), 3);
}
#[test]
fn a_sigma_past_the_kernel_capacity_is_refused_at_construction() {
let peak = PeakValue::of_pixel::<Mono8>();
assert!(SsimParams::try_new(peak, sigma!(21.0), 0.01, 0.03).is_ok());
let refused = SsimParams::try_new(peak, sigma!(40.0), 0.01, 0.03);
assert!(
matches!(refused, Err(Error::InvalidParameter(_))),
"{refused:?}"
);
}
#[test]
fn an_image_against_itself_is_exactly_one() {
let image = checkerboard(32, 4);
let map = ssim_map(&image, &image, reference_params()).unwrap();
for y in 0..map.height() {
for (x, pixel) in map.row(y).iter().enumerate() {
assert_eq!(pixel.0, 1.0, "at ({x}, {y})");
}
}
assert_eq!(ssim(&image, &image, reference_params()).unwrap(), 1.0);
}
#[test]
fn the_score_is_symmetric_in_its_arguments() {
let a = checkerboard(40, 5);
let b = Image::generate(40, 40, |x, y| Mono8::new(((x * 7 + y * 3) % 256) as u8));
let forward = ssim(&a, &b, reference_params()).unwrap();
let backward = ssim(&b, &a, reference_params()).unwrap();
assert!(
(forward - backward).abs() < 1e-12,
"{forward} vs {backward}"
);
}
#[test]
fn the_score_stays_within_minus_one_and_one() {
let a = checkerboard(48, 6);
let inverted = Image::generate(48, 48, |x, y| Mono8::new(255 - a.pixel_at(x, y).value()));
let noise = Image::generate(48, 48, |x, y| Mono8::new(((x * 37 + y * 101) % 256) as u8));
for other in [&inverted, &noise] {
let map = ssim_map(&a, other, reference_params()).unwrap();
for pixel in map.as_slice() {
assert!(pixel.0 >= -1.0 && pixel.0 <= 1.0, "{}", pixel.0);
}
}
}
#[test]
fn inverting_contrast_drives_the_score_negative() {
let a = checkerboard(48, 6);
let inverted = Image::generate(48, 48, |x, y| Mono8::new(255 - a.pixel_at(x, y).value()));
let score = ssim(&a, &inverted, reference_params()).unwrap();
assert!(score < 0.0, "{score}");
}
#[test]
fn a_flat_pair_of_equal_brightness_is_similar_and_a_different_one_is_not() {
let params = reference_params();
let grey = Image::fill(32, 32, Mono8::new(128));
let same = Image::fill(32, 32, Mono8::new(128));
assert_eq!(ssim(&grey, &same, params).unwrap(), 1.0);
let brighter = Image::fill(32, 32, Mono8::new(200));
let score = ssim(&grey, &brighter, params).unwrap();
assert!(score < 1.0 && score > 0.0, "{score}");
let c1 = (0.01 * 255.0f64).powi(2);
let expected = (2.0 * 128.0 * 200.0 + c1) / (128.0 * 128.0 + 200.0 * 200.0 + c1);
assert!((score - expected).abs() < 1e-9, "{score} vs {expected}");
}
#[test]
fn more_degradation_scores_lower() {
let params = reference_params();
let reference =
Image::generate(64, 64, |x, y| Mono8::new((((x * 5) ^ (y * 3)) % 256) as u8));
let mut previous = 1.000_001;
for amplitude in [0u8, 4, 16, 48] {
let degraded = Image::generate(64, 64, |x, y| {
let base = reference.pixel_at(x, y).value();
let wobble = if (x + y) % 2 == 0 { amplitude } else { 0 };
Mono8::new(base.saturating_add(wobble))
});
let score = ssim(&reference, °raded, params).unwrap();
assert!(
score < previous,
"amplitude {amplitude}: {score} !< {previous}"
);
previous = score;
}
}
#[test]
fn the_map_is_the_valid_window_block_and_is_offset_by_the_radius() {
let params = reference_params();
let radius = params.window_size() / 2;
let reference = Image::generate(48, 40, |x, _| Mono8::new((x * 5) as u8));
let mut damaged = reference.clone();
for y in 18..26 {
for x in 18..26 {
*damaged.pixel_at_mut(x, y) = Mono8::new(0);
}
}
let map = ssim_map(&reference, &damaged, params).unwrap();
assert_eq!(map.size(), Size::new(48 - 2 * radius, 40 - 2 * radius));
assert!(map.pixel_at(22 - radius, 22 - radius).0 < 0.5);
assert!(map.pixel_at(2, 2).0 > 0.99);
}
#[test]
fn an_image_the_window_does_not_fit_is_refused_rather_than_clipped() {
let params = reference_params();
let small: Image<Mono8> = Image::fill(10, 32, Mono8::new(0));
let other: Image<Mono8> = Image::fill(10, 32, Mono8::new(0));
let refused = ssim_map(&small, &other, params);
assert!(
matches!(refused, Err(Error::InvalidParameter(_))),
"{refused:?}"
);
let exact: Image<Mono8> = Image::fill(11, 11, Mono8::new(50));
let map = ssim_map(&exact, &exact, params).unwrap();
assert_eq!(map.size(), Size::new(1, 1));
assert_eq!(ssim(&exact, &exact, params).unwrap(), 1.0);
}
#[test]
fn a_size_mismatch_is_a_tier_two_error() {
let params = reference_params();
let a: Image<Mono8> = Image::fill(32, 32, Mono8::new(0));
let b: Image<Mono8> = Image::fill(32, 33, Mono8::new(0));
assert_eq!(
ssim(&a, &b, params),
Err(Error::SizeMismatch {
expected: Size::new(32, 32),
actual: Size::new(32, 33),
})
);
}
#[test]
fn the_score_is_the_mean_of_the_map() {
let params = reference_params();
let a = checkerboard(40, 5);
let b = Image::generate(40, 40, |x, y| {
Mono8::new(a.pixel_at(x, y).value().wrapping_add(((x * y) % 20) as u8))
});
let map = ssim_map(&a, &b, params).unwrap();
let expected =
map.as_slice().iter().map(|p| p.0).sum::<f64>() / (map.width() * map.height()) as f64;
let score = ssim(&a, &b, params).unwrap();
assert!((score - expected).abs() < 1e-12, "{score} vs {expected}");
}
#[test]
fn a_pedestal_does_not_collapse_the_variance() {
let peak = peak!(1023.0);
let params = SsimParams::reference(peak);
let structure = |pedestal: u16, transpose: bool| {
Image::generate(64, 64, |x, y| {
let (u, v) = if transpose { (y, x) } else { (x, y) };
Mono16::new(pedestal + (((u * 3 + v) % 7) * 2) as u16)
})
};
let on_pedestal =
ssim(&structure(40_000, false), &structure(40_000, true), params).unwrap();
let at_zero = ssim(&structure(0, false), &structure(0, true), params).unwrap();
assert!(
(on_pedestal - at_zero).abs() < 1e-4,
"the pedestal moved the score: {on_pedestal} vs {at_zero}",
);
assert!(on_pedestal > 0.9 && on_pedestal < 1.0, "{on_pedestal}");
}
#[test]
fn a_pedestal_cannot_push_the_score_out_of_range() {
let params = SsimParams::reference(peak!(20.0));
let a = Image::generate(48, 48, |x, y| {
Mono16::new(60_000 + ((x * 3 + y) % 7) as u16)
});
let b = Image::generate(48, 48, |x, y| {
Mono16::new(60_000 + ((y * 3 + x) % 7) as u16)
});
let map = ssim_map(&a, &b, params).unwrap();
for pixel in map.as_slice() {
assert!(pixel.0 >= -1.0 && pixel.0 <= 1.0, "{}", pixel.0);
}
}
#[test]
fn the_variance_never_lands_negative_in_the_map() {
let params = reference_params();
let half_flat = Image::generate(48, 48, |x, _| {
Mono8::new(if x < 24 { 200 } else { (x * 3) as u8 })
});
let map = ssim_map(&half_flat, &half_flat, params).unwrap();
for pixel in map.as_slice() {
assert!(pixel.0 <= 1.0, "{}", pixel.0);
}
}
#[test]
fn nan_propagates_rather_than_being_silently_dropped() {
let params = SsimParams::reference(peak!(1.0));
let mut a = Image::generate(32, 32, |x, _| MonoF32::new(x as f32 / 31.0));
let b = a.clone();
*a.pixel_at_mut(16, 16) = MonoF32::new(f32::NAN);
let score = ssim(&a, &b, params).unwrap();
assert!(score.is_nan(), "{score}");
let stats: ChannelStatistics<_> = image_statistics(&a);
assert_eq!(stats.nan_count, 1);
}
#[test]
fn a_wider_window_smooths_the_map_and_shrinks_it_further() {
let peak = PeakValue::of_pixel::<Mono8>();
let reference =
Image::generate(64, 64, |x, y| Mono8::new((((x * 5) ^ (y * 3)) % 256) as u8));
let degraded = Image::generate(64, 64, |x, y| {
Mono8::new(
reference
.pixel_at(x, y)
.value()
.saturating_add(if (x + y) % 3 == 0 { 20 } else { 0 }),
)
});
let narrow = SsimParams::try_new(peak, sigma!(1.5), 0.01, 0.03).unwrap();
let wide = SsimParams::try_new(peak, sigma!(4.0), 0.01, 0.03).unwrap();
assert_eq!(narrow.window_size(), 11);
assert_eq!(wide.window_size(), 25);
let narrow_map = ssim_map(&reference, °raded, narrow).unwrap();
let wide_map = ssim_map(&reference, °raded, wide).unwrap();
assert_eq!(narrow_map.size(), Size::new(54, 54));
assert_eq!(wide_map.size(), Size::new(40, 40));
let spread = |map: &Image<MonoF64>| {
let stats: ChannelStatistics<_> = image_statistics(map);
stats.std_dev().unwrap()
};
assert!(spread(&wide_map) < spread(&narrow_map));
}
#[test]
fn the_reference_window_matches_a_hand_computed_gaussian_window() {
let params = reference_params();
let a = Image::generate(11, 11, |x, y| Mono8::new(((x * 17 + y * 5) % 200) as u8));
let b = Image::generate(11, 11, |x, y| Mono8::new(((y * 11 + x * 3) % 200) as u8));
let taps: Vec<f64> = {
let sigma = 1.5f64;
let raw: Vec<f64> = (0..11)
.map(|i| {
let d = i as f64 - 5.0;
(-d * d / (2.0 * sigma * sigma)).exp()
})
.collect();
let sum: f64 = raw.iter().sum();
raw.iter().map(|w| w / sum).collect()
};
let mut mean_a = 0.0;
let mut mean_b = 0.0;
let mut moment_aa = 0.0;
let mut moment_bb = 0.0;
let mut moment_ab = 0.0;
for y in 0..11 {
for x in 0..11 {
let weight = taps[x] * taps[y];
let va = f64::from(a.pixel_at(x, y).value());
let vb = f64::from(b.pixel_at(x, y).value());
mean_a += weight * va;
mean_b += weight * vb;
moment_aa += weight * va * va;
moment_bb += weight * vb * vb;
moment_ab += weight * va * vb;
}
}
let variance_a = moment_aa - mean_a * mean_a;
let variance_b = moment_bb - mean_b * mean_b;
let covariance = moment_ab - mean_a * mean_b;
let c1 = (0.01 * 255.0f64).powi(2);
let c2 = (0.03 * 255.0f64).powi(2);
let expected = ((2.0 * mean_a * mean_b + c1) * (2.0 * covariance + c2))
/ ((mean_a * mean_a + mean_b * mean_b + c1) * (variance_a + variance_b + c2));
let map = ssim_map(&a, &b, params).unwrap();
assert_eq!(map.size(), Size::new(1, 1));
let actual = map.pixel_at(0, 0).0;
assert!((actual - expected).abs() < 1e-6, "{actual} vs {expected}");
}
}