1use ndarray::ArrayView2;
9use num_traits::ToPrimitive;
10use thiserror::Error;
11
12const SSIM_WINDOW_SIZE: usize = 11;
13const SSIM_GAUSSIAN_SIGMA: f64 = 1.5;
14const SSIM_K1: f64 = 0.01;
15const SSIM_K2: f64 = 0.03;
16
17#[derive(Debug, Error, PartialEq)]
19pub enum IntensityMetricError {
20 #[error("reference and candidate shapes differ: {reference:?} and {candidate:?}")]
21 ShapeMismatch {
22 reference: (usize, usize),
23 candidate: (usize, usize),
24 },
25 #[error("valid_mask shape {actual:?} does not match image shape {expected:?}")]
26 MaskShapeMismatch {
27 actual: (usize, usize),
28 expected: (usize, usize),
29 },
30 #[error("at least one valid pixel is required")]
31 EmptyValidMask,
32 #[error("{input} contains a value that cannot be represented as f64")]
33 UnsupportedScalar { input: &'static str },
34 #[error("{input} contains a non-finite value")]
35 NonFinite { input: &'static str },
36 #[error("{metric} is undefined because the reference normalization is zero")]
37 ZeroReferenceNormalization { metric: &'static str },
38 #[error("correlation is undefined for a constant input")]
39 ZeroVariance,
40 #[error("{metric} requires non-negative intensities")]
41 NegativeIntensity { metric: &'static str },
42 #[error("data_range must be finite and strictly positive")]
43 InvalidDataRange,
44 #[error("epsilon must be finite and strictly positive")]
45 InvalidEpsilon,
46 #[error(
47 "SSIM requires images at least {SSIM_WINDOW_SIZE} by {SSIM_WINDOW_SIZE}, got {shape:?}"
48 )]
49 SsimImageTooSmall { shape: (usize, usize) },
50 #[error(
51 "SSIM requires at least one fully valid {SSIM_WINDOW_SIZE} by {SSIM_WINDOW_SIZE} window"
52 )]
53 NoValidSsimWindow,
54}
55
56pub fn bias<T>(
58 reference: ArrayView2<'_, T>,
59 candidate: ArrayView2<'_, T>,
60 valid_mask: Option<ArrayView2<'_, bool>>,
61) -> Result<f64, IntensityMetricError>
62where
63 T: ToPrimitive,
64{
65 let mut sum = 0.0;
66 let count = for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
67 sum += candidate - reference;
68 })?;
69 Ok(sum / count as f64)
70}
71
72pub fn mae<T>(
74 reference: ArrayView2<'_, T>,
75 candidate: ArrayView2<'_, T>,
76 valid_mask: Option<ArrayView2<'_, bool>>,
77) -> Result<f64, IntensityMetricError>
78where
79 T: ToPrimitive,
80{
81 let mut sum = 0.0;
82 let count = for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
83 sum += (candidate - reference).abs();
84 })?;
85 Ok(sum / count as f64)
86}
87
88pub fn mse<T>(
90 reference: ArrayView2<'_, T>,
91 candidate: ArrayView2<'_, T>,
92 valid_mask: Option<ArrayView2<'_, bool>>,
93) -> Result<f64, IntensityMetricError>
94where
95 T: ToPrimitive,
96{
97 let (sum, count) = squared_error_sum(reference, candidate, valid_mask)?;
98 Ok(sum / count as f64)
99}
100
101pub fn rmse<T>(
103 reference: ArrayView2<'_, T>,
104 candidate: ArrayView2<'_, T>,
105 valid_mask: Option<ArrayView2<'_, bool>>,
106) -> Result<f64, IntensityMetricError>
107where
108 T: ToPrimitive,
109{
110 Ok(mse(reference, candidate, valid_mask)?.sqrt())
111}
112
113pub fn relative_l1<T>(
115 reference: ArrayView2<'_, T>,
116 candidate: ArrayView2<'_, T>,
117 valid_mask: Option<ArrayView2<'_, bool>>,
118) -> Result<f64, IntensityMetricError>
119where
120 T: ToPrimitive,
121{
122 let mut residual_sum = 0.0;
123 let mut reference_sum = 0.0;
124 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
125 residual_sum += (candidate - reference).abs();
126 reference_sum += reference.abs();
127 })?;
128 if reference_sum == 0.0 {
129 return Err(IntensityMetricError::ZeroReferenceNormalization {
130 metric: "relative_l1",
131 });
132 }
133 Ok(residual_sum / reference_sum)
134}
135
136pub fn nrmse<T>(
138 reference: ArrayView2<'_, T>,
139 candidate: ArrayView2<'_, T>,
140 valid_mask: Option<ArrayView2<'_, bool>>,
141) -> Result<f64, IntensityMetricError>
142where
143 T: ToPrimitive,
144{
145 let mut residual_squared = 0.0;
146 let mut reference_squared = 0.0;
147 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
148 residual_squared += (candidate - reference).powi(2);
149 reference_squared += reference.powi(2);
150 })?;
151 if reference_squared == 0.0 {
152 return Err(IntensityMetricError::ZeroReferenceNormalization { metric: "nrmse" });
153 }
154 Ok(residual_squared.sqrt() / reference_squared.sqrt())
155}
156
157pub fn amplitude_nrmse<T>(
159 reference: ArrayView2<'_, T>,
160 candidate: ArrayView2<'_, T>,
161 valid_mask: Option<ArrayView2<'_, bool>>,
162) -> Result<f64, IntensityMetricError>
163where
164 T: ToPrimitive,
165{
166 validate_non_negative(reference, candidate, valid_mask.clone(), "amplitude_nrmse")?;
167 let mut residual_squared = 0.0;
168 let mut reference_squared = 0.0;
169 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
170 let reference_amplitude = reference.sqrt();
171 let candidate_amplitude = candidate.sqrt();
172 residual_squared += (candidate_amplitude - reference_amplitude).powi(2);
173 reference_squared += reference;
174 })?;
175 if reference_squared == 0.0 {
176 return Err(IntensityMetricError::ZeroReferenceNormalization {
177 metric: "amplitude_nrmse",
178 });
179 }
180 Ok(residual_squared.sqrt() / reference_squared.sqrt())
181}
182
183pub fn correlation<T>(
185 reference: ArrayView2<'_, T>,
186 candidate: ArrayView2<'_, T>,
187 valid_mask: Option<ArrayView2<'_, bool>>,
188) -> Result<f64, IntensityMetricError>
189where
190 T: ToPrimitive,
191{
192 let mut reference_sum = 0.0;
193 let mut candidate_sum = 0.0;
194 let mut reference_squared = 0.0;
195 let mut candidate_squared = 0.0;
196 let mut product_sum = 0.0;
197 let count = for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
198 reference_sum += reference;
199 candidate_sum += candidate;
200 reference_squared += reference * reference;
201 candidate_squared += candidate * candidate;
202 product_sum += reference * candidate;
203 })? as f64;
204 let reference_variance = reference_squared - reference_sum * reference_sum / count;
205 let candidate_variance = candidate_squared - candidate_sum * candidate_sum / count;
206 if reference_variance <= 0.0 || candidate_variance <= 0.0 {
207 return Err(IntensityMetricError::ZeroVariance);
208 }
209 let covariance = product_sum - reference_sum * candidate_sum / count;
210 Ok(covariance / (reference_variance * candidate_variance).sqrt())
211}
212
213pub fn psnr<T>(
215 reference: ArrayView2<'_, T>,
216 candidate: ArrayView2<'_, T>,
217 valid_mask: Option<ArrayView2<'_, bool>>,
218 data_range: f64,
219) -> Result<f64, IntensityMetricError>
220where
221 T: ToPrimitive,
222{
223 validate_data_range(data_range)?;
224 let value = mse(reference, candidate, valid_mask)?;
225 if value == 0.0 {
226 Ok(f64::INFINITY)
227 } else {
228 Ok(10.0 * (data_range * data_range / value).log10())
229 }
230}
231
232pub fn ssim<T>(
234 reference: ArrayView2<'_, T>,
235 candidate: ArrayView2<'_, T>,
236 valid_mask: Option<ArrayView2<'_, bool>>,
237 data_range: f64,
238) -> Result<f64, IntensityMetricError>
239where
240 T: ToPrimitive,
241{
242 validate_data_range(data_range)?;
243 let shape = reference.dim();
244 validate_shapes(
245 shape,
246 candidate.dim(),
247 valid_mask.as_ref().map(|mask| mask.dim()),
248 )?;
249 if shape.0 < SSIM_WINDOW_SIZE || shape.1 < SSIM_WINDOW_SIZE {
250 return Err(IntensityMetricError::SsimImageTooSmall { shape });
251 }
252 for_each_valid_pair(reference, candidate, valid_mask.clone(), |_, _| {})?;
254
255 let weights = gaussian_weights();
256 let c1 = (SSIM_K1 * data_range).powi(2);
257 let c2 = (SSIM_K2 * data_range).powi(2);
258 let mut score_sum = 0.0;
259 let mut window_count = 0usize;
260 for row in 0..=shape.0 - SSIM_WINDOW_SIZE {
261 for column in 0..=shape.1 - SSIM_WINDOW_SIZE {
262 if valid_mask.as_ref().is_some_and(|mask| {
263 (row..row + SSIM_WINDOW_SIZE).any(|window_row| {
264 (column..column + SSIM_WINDOW_SIZE)
265 .any(|window_column| !mask[(window_row, window_column)])
266 })
267 }) {
268 continue;
269 }
270 let (mut reference_mean, mut candidate_mean) = (0.0, 0.0);
271 for window_row in 0..SSIM_WINDOW_SIZE {
272 for window_column in 0..SSIM_WINDOW_SIZE {
273 let weight = weights[window_row * SSIM_WINDOW_SIZE + window_column];
274 reference_mean += weight
275 * value_as_f64(
276 &reference[(row + window_row, column + window_column)],
277 "reference",
278 )?;
279 candidate_mean += weight
280 * value_as_f64(
281 &candidate[(row + window_row, column + window_column)],
282 "candidate",
283 )?;
284 }
285 }
286 let (mut reference_variance, mut candidate_variance, mut covariance) = (0.0, 0.0, 0.0);
287 for window_row in 0..SSIM_WINDOW_SIZE {
288 for window_column in 0..SSIM_WINDOW_SIZE {
289 let weight = weights[window_row * SSIM_WINDOW_SIZE + window_column];
290 let reference_value = value_as_f64(
291 &reference[(row + window_row, column + window_column)],
292 "reference",
293 )? - reference_mean;
294 let candidate_value = value_as_f64(
295 &candidate[(row + window_row, column + window_column)],
296 "candidate",
297 )? - candidate_mean;
298 reference_variance += weight * reference_value * reference_value;
299 candidate_variance += weight * candidate_value * candidate_value;
300 covariance += weight * reference_value * candidate_value;
301 }
302 }
303 score_sum += ((2.0 * reference_mean * candidate_mean + c1) * (2.0 * covariance + c2))
304 / ((reference_mean.powi(2) + candidate_mean.powi(2) + c1)
305 * (reference_variance + candidate_variance + c2));
306 window_count += 1;
307 }
308 }
309 if window_count == 0 {
310 return Err(IntensityMetricError::NoValidSsimWindow);
311 }
312 Ok(score_sum / window_count as f64)
313}
314
315pub fn poisson_deviance<T>(
317 reference: ArrayView2<'_, T>,
318 candidate: ArrayView2<'_, T>,
319 valid_mask: Option<ArrayView2<'_, bool>>,
320 epsilon: f64,
321) -> Result<f64, IntensityMetricError>
322where
323 T: ToPrimitive,
324{
325 validate_epsilon(epsilon)?;
326 validate_non_negative(reference, candidate, valid_mask.clone(), "poisson_deviance")?;
327 let mut sum = 0.0;
328 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
329 if reference == 0.0 && candidate == 0.0 {
330 return;
331 }
332 let candidate = candidate.max(epsilon);
333 let log_term = if reference == 0.0 {
334 0.0
335 } else {
336 reference * (reference / candidate).ln()
337 };
338 sum += 2.0 * (candidate - reference + log_term);
339 })?;
340 Ok(sum)
341}
342
343pub fn mean_poisson_deviance<T>(
345 reference: ArrayView2<'_, T>,
346 candidate: ArrayView2<'_, T>,
347 valid_mask: Option<ArrayView2<'_, bool>>,
348 epsilon: f64,
349) -> Result<f64, IntensityMetricError>
350where
351 T: ToPrimitive,
352{
353 validate_epsilon(epsilon)?;
354 validate_non_negative(
355 reference,
356 candidate,
357 valid_mask.clone(),
358 "mean_poisson_deviance",
359 )?;
360 let mut sum = 0.0;
361 let count = for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
362 if reference == 0.0 && candidate == 0.0 {
363 return;
364 }
365 let candidate = candidate.max(epsilon);
366 let log_term = if reference == 0.0 {
367 0.0
368 } else {
369 reference * (reference / candidate).ln()
370 };
371 sum += 2.0 * (candidate - reference + log_term);
372 })?;
373 Ok(sum / count as f64)
374}
375
376pub fn fitted_gain<T>(
378 reference: ArrayView2<'_, T>,
379 candidate: ArrayView2<'_, T>,
380 valid_mask: Option<ArrayView2<'_, bool>>,
381) -> Result<f64, IntensityMetricError>
382where
383 T: ToPrimitive,
384{
385 let mut numerator = 0.0;
386 let mut denominator = 0.0;
387 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
388 numerator += reference * candidate;
389 denominator += reference * reference;
390 })?;
391 if denominator == 0.0 {
392 return Err(IntensityMetricError::ZeroReferenceNormalization {
393 metric: "fitted_gain",
394 });
395 }
396 Ok(numerator / denominator)
397}
398
399fn squared_error_sum<T>(
400 reference: ArrayView2<'_, T>,
401 candidate: ArrayView2<'_, T>,
402 valid_mask: Option<ArrayView2<'_, bool>>,
403) -> Result<(f64, usize), IntensityMetricError>
404where
405 T: ToPrimitive,
406{
407 let mut sum = 0.0;
408 let count = for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
409 sum += (candidate - reference).powi(2);
410 })?;
411 Ok((sum, count))
412}
413
414fn for_each_valid_pair<T>(
415 reference: ArrayView2<'_, T>,
416 candidate: ArrayView2<'_, T>,
417 valid_mask: Option<ArrayView2<'_, bool>>,
418 mut function: impl FnMut(f64, f64),
419) -> Result<usize, IntensityMetricError>
420where
421 T: ToPrimitive,
422{
423 let shape = reference.dim();
424 validate_shapes(
425 shape,
426 candidate.dim(),
427 valid_mask.as_ref().map(|mask| mask.dim()),
428 )?;
429 let mut count = 0usize;
430 for row in 0..shape.0 {
431 for column in 0..shape.1 {
432 if valid_mask.as_ref().is_some_and(|mask| !mask[(row, column)]) {
433 continue;
434 }
435 let reference = value_as_f64(&reference[(row, column)], "reference")?;
436 let candidate = value_as_f64(&candidate[(row, column)], "candidate")?;
437 function(reference, candidate);
438 count += 1;
439 }
440 }
441 if count == 0 {
442 return Err(IntensityMetricError::EmptyValidMask);
443 }
444 Ok(count)
445}
446
447fn validate_non_negative<T>(
448 reference: ArrayView2<'_, T>,
449 candidate: ArrayView2<'_, T>,
450 valid_mask: Option<ArrayView2<'_, bool>>,
451 metric: &'static str,
452) -> Result<(), IntensityMetricError>
453where
454 T: ToPrimitive,
455{
456 let mut negative = false;
457 for_each_valid_pair(reference, candidate, valid_mask, |reference, candidate| {
458 negative |= reference < 0.0 || candidate < 0.0;
459 })?;
460 if negative {
461 return Err(IntensityMetricError::NegativeIntensity { metric });
462 }
463 Ok(())
464}
465
466fn validate_shapes(
467 reference: (usize, usize),
468 candidate: (usize, usize),
469 valid_mask: Option<(usize, usize)>,
470) -> Result<(), IntensityMetricError> {
471 if reference != candidate {
472 return Err(IntensityMetricError::ShapeMismatch {
473 reference,
474 candidate,
475 });
476 }
477 if let Some(actual) = valid_mask
478 && actual != reference
479 {
480 return Err(IntensityMetricError::MaskShapeMismatch {
481 actual,
482 expected: reference,
483 });
484 }
485 Ok(())
486}
487
488fn value_as_f64<T>(value: &T, input: &'static str) -> Result<f64, IntensityMetricError>
489where
490 T: ToPrimitive,
491{
492 let value = value
493 .to_f64()
494 .ok_or(IntensityMetricError::UnsupportedScalar { input })?;
495 if !value.is_finite() {
496 return Err(IntensityMetricError::NonFinite { input });
497 }
498 Ok(value)
499}
500
501fn validate_data_range(data_range: f64) -> Result<(), IntensityMetricError> {
502 if !data_range.is_finite() || data_range <= 0.0 {
503 return Err(IntensityMetricError::InvalidDataRange);
504 }
505 Ok(())
506}
507
508fn validate_epsilon(epsilon: f64) -> Result<(), IntensityMetricError> {
509 if !epsilon.is_finite() || epsilon <= 0.0 {
510 return Err(IntensityMetricError::InvalidEpsilon);
511 }
512 Ok(())
513}
514
515fn gaussian_weights() -> Vec<f64> {
516 let center = (SSIM_WINDOW_SIZE - 1) as f64 / 2.0;
517 let mut weights = Vec::with_capacity(SSIM_WINDOW_SIZE * SSIM_WINDOW_SIZE);
518 for row in 0..SSIM_WINDOW_SIZE {
519 for column in 0..SSIM_WINDOW_SIZE {
520 let squared_distance = (row as f64 - center).powi(2) + (column as f64 - center).powi(2);
521 weights.push((-squared_distance / (2.0 * SSIM_GAUSSIAN_SIGMA.powi(2))).exp());
522 }
523 }
524 let normalization = weights.iter().sum::<f64>();
525 for weight in &mut weights {
526 *weight /= normalization;
527 }
528 weights
529}
530
531#[cfg(test)]
532mod tests {
533 use approx::assert_relative_eq;
534 use ndarray::array;
535
536 use super::*;
537
538 #[test]
539 fn residual_metrics_use_candidate_minus_reference_and_support_integers() {
540 let reference = array![[1u8, 2u8]];
541 let candidate = array![[2u8, 0u8]];
542 assert_relative_eq!(
543 bias(reference.view(), candidate.view(), None).unwrap(),
544 -0.5
545 );
546 assert_relative_eq!(mae(reference.view(), candidate.view(), None).unwrap(), 1.5);
547 assert_relative_eq!(mse(reference.view(), candidate.view(), None).unwrap(), 2.5);
548 assert_relative_eq!(
549 rmse(reference.view(), candidate.view(), None).unwrap(),
550 2.5f64.sqrt()
551 );
552 assert_relative_eq!(
553 relative_l1(reference.view(), candidate.view(), None).unwrap(),
554 1.0
555 );
556 assert_relative_eq!(
557 nrmse(reference.view(), candidate.view(), None).unwrap(),
558 1.0
559 );
560 assert_relative_eq!(
561 correlation(reference.view(), candidate.view(), None).unwrap(),
562 -1.0
563 );
564 }
565
566 #[test]
567 fn mask_selects_valid_pixels() {
568 let reference = array![[1.0, 20.0]];
569 let candidate = array![[2.0, 0.0]];
570 let mask = array![[true, false]];
571 assert_relative_eq!(
572 bias(reference.view(), candidate.view(), Some(mask.view())).unwrap(),
573 1.0
574 );
575 }
576
577 #[test]
578 fn quality_metrics_have_documented_reference_values() {
579 let reference = array![[0.0, 1.0]];
580 let candidate = array![[0.0, 0.0]];
581 assert_relative_eq!(
582 psnr(reference.view(), candidate.view(), None, 1.0).unwrap(),
583 10.0 * 2.0f64.log10()
584 );
585
586 let image = ndarray::Array2::<f64>::ones((11, 11));
587 assert_relative_eq!(ssim(image.view(), image.view(), None, 1.0).unwrap(), 1.0);
588 }
589
590 #[test]
591 fn poisson_deviance_and_gain_are_well_defined() {
592 let reference = array![[0.0, 2.0]];
593 let candidate = array![[0.0, 2.0]];
594 assert_relative_eq!(
595 poisson_deviance(reference.view(), candidate.view(), None, 1e-12).unwrap(),
596 0.0
597 );
598 assert_relative_eq!(
599 mean_poisson_deviance(reference.view(), candidate.view(), None, 1e-12).unwrap(),
600 0.0
601 );
602
603 let scaled = array![[0.0, 4.0]];
604 assert_relative_eq!(
605 fitted_gain(reference.view(), scaled.view(), None).unwrap(),
606 2.0
607 );
608 }
609
610 #[test]
611 fn validation_covers_masks_ranges_and_intensity_domains() {
612 let reference = array![[1.0, 2.0]];
613 let candidate = array![[2.0, 0.0]];
614 let wrong_mask = array![[true], [false]];
615 assert!(matches!(
616 mae(reference.view(), candidate.view(), Some(wrong_mask.view())),
617 Err(IntensityMetricError::MaskShapeMismatch { .. })
618 ));
619 assert!(matches!(
620 psnr(reference.view(), candidate.view(), None, 0.0),
621 Err(IntensityMetricError::InvalidDataRange)
622 ));
623 assert!(matches!(
624 ssim(reference.view(), candidate.view(), None, 1.0),
625 Err(IntensityMetricError::SsimImageTooSmall { .. })
626 ));
627
628 let negative = array![[-1.0, 2.0]];
629 assert!(matches!(
630 amplitude_nrmse(negative.view(), candidate.view(), None),
631 Err(IntensityMetricError::NegativeIntensity { .. })
632 ));
633 assert!(matches!(
634 poisson_deviance(negative.view(), candidate.view(), None, 1e-12),
635 Err(IntensityMetricError::NegativeIntensity { .. })
636 ));
637 }
638}