Skip to main content

fpm_rs/metrics/intensity/
compare.rs

1//! Domain-agnostic metrics comparing reference and estimate intensity images.
2//!
3//! Signed quantities use the residual `estimate - reference`; `valid_mask`,
4//! when supplied, includes pixels whose value is `true`. These functions are
5//! evaluation metrics, not reconstruction optimization objectives.
6
7use ndarray::ArrayView2;
8use num_traits::ToPrimitive;
9use serde::{Deserialize, Serialize};
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/// Errors returned by intensity comparison metrics.
18#[derive(Debug, Error, PartialEq)]
19pub enum IntensityMetricError {
20    /// Reference and estimate have different `(height, width)` shapes.
21    #[error("reference and estimate shapes differ: {reference:?} and {estimate:?}")]
22    ShapeMismatch {
23        /// Reference image shape.
24        reference: (usize, usize),
25        /// Estimate image shape.
26        estimate: (usize, usize),
27    },
28    /// A validity mask does not match the image shape.
29    #[error("valid_mask shape {actual:?} does not match image shape {expected:?}")]
30    MaskShapeMismatch {
31        /// Supplied mask shape.
32        actual: (usize, usize),
33        /// Required image shape.
34        expected: (usize, usize),
35    },
36    /// No pixel was selected by the optional mask.
37    #[error("at least one valid pixel is required")]
38    EmptyValidMask,
39    /// A generic scalar could not be represented as `f64`.
40    #[error("{input} contains a value that cannot be represented as f64")]
41    UnsupportedScalar {
42        /// Name of the input containing the unsupported scalar.
43        input: &'static str,
44    },
45    /// An input contains NaN or infinity.
46    #[error("{input} contains a non-finite value")]
47    NonFinite {
48        /// Name of the non-finite input.
49        input: &'static str,
50    },
51    /// A relative metric has zero reference norm.
52    #[error("{metric} is undefined because the reference normalization is zero")]
53    ZeroReferenceNormalization {
54        /// Metric whose denominator was zero.
55        metric: &'static str,
56    },
57    /// Pearson correlation is undefined because an input has zero variance.
58    #[error("correlation is undefined for a constant input")]
59    ZeroVariance,
60    /// A square-root or Poisson metric received a negative intensity.
61    #[error("{metric} requires non-negative intensities")]
62    NegativeIntensity {
63        /// Metric requiring non-negative values.
64        metric: &'static str,
65    },
66    /// PSNR or SSIM received a non-positive or non-finite data range.
67    #[error("data_range must be finite and strictly positive")]
68    InvalidDataRange,
69    /// A Poisson metric received a non-positive or non-finite numerical floor.
70    #[error("epsilon must be finite and strictly positive")]
71    InvalidEpsilon,
72    /// An image is smaller than the canonical 11-by-11 SSIM window.
73    #[error(
74        "SSIM requires images at least {SSIM_WINDOW_SIZE} by {SSIM_WINDOW_SIZE}, got {shape:?}"
75    )]
76    SsimImageTooSmall {
77        /// Supplied image shape.
78        shape: (usize, usize),
79    },
80    /// A mask excludes every complete SSIM window.
81    #[error(
82        "SSIM requires at least one fully valid {SSIM_WINDOW_SIZE} by {SSIM_WINDOW_SIZE} window"
83    )]
84    NoValidSsimWindow,
85}
86
87/// Aggregate residual statistics comparing an estimate with a reference image.
88#[derive(Clone, Debug, Serialize, Deserialize)]
89pub struct IntensityComparisonMetrics {
90    /// Sum of selected reference intensities.
91    pub reference_sum: f64,
92    /// Sum of selected estimate intensities.
93    pub estimate_sum: f64,
94    /// Sum of absolute signed residuals.
95    pub residual_l1: f64,
96    /// Euclidean norm of signed residuals.
97    pub residual_l2: f64,
98    /// Mean signed residual `estimate - reference`.
99    pub residual_mean: f64,
100    /// Population standard deviation of signed residuals.
101    pub residual_std: f64,
102    /// Maximum absolute residual.
103    pub residual_max_abs: f64,
104    /// Residual L2 norm divided by reference L2 norm.
105    pub normalized_l2: f64,
106    /// Selected reference pixels at or above the optional saturation threshold.
107    pub saturated_pixels: Option<usize>,
108}
109
110/// Calculate aggregate residual statistics for a reference/estimate pair.
111///
112/// Signed residuals are `estimate - reference`. `saturation_value`, when
113/// present, counts valid reference pixels greater than or equal to that value.
114pub fn compare_intensity<T>(
115    reference: ArrayView2<'_, T>,
116    estimate: ArrayView2<'_, T>,
117    valid_mask: Option<ArrayView2<'_, bool>>,
118    saturation_value: Option<f64>,
119) -> Result<IntensityComparisonMetrics, IntensityMetricError>
120where
121    T: ToPrimitive,
122{
123    let mut accumulator = ComparisonAccumulator::default();
124    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
125        accumulator.push(reference, estimate, saturation_value);
126    })?;
127    Ok(accumulator.finish(count, saturation_value))
128}
129
130/// Adapter for the reconstruction diagnostics' flat slices and byte masks.
131///
132/// This remains crate-private so the public metric API consistently uses
133/// two-dimensional ndarray views and Boolean masks.
134pub(crate) fn compare_intensity_u8_masked(
135    reference: &[f64],
136    estimate: &[f64],
137    valid_mask: Option<&[u8]>,
138    saturation_value: Option<f64>,
139) -> crate::Result<IntensityComparisonMetrics> {
140    if reference.len() != estimate.len() || reference.is_empty() {
141        return Err(crate::Error::InvalidShape(format!(
142            "intensity comparison inputs have lengths {} and {}",
143            reference.len(),
144            estimate.len()
145        )));
146    }
147    if valid_mask.is_some_and(|mask| mask.len() != reference.len()) {
148        return Err(crate::Error::LengthMismatch {
149            actual: valid_mask.map_or(0, <[u8]>::len),
150            expected: reference.len(),
151            shape: (1, reference.len()),
152        });
153    }
154
155    let mut accumulator = ComparisonAccumulator::default();
156    let mut count = 0;
157    for (index, (&reference, &estimate)) in reference.iter().zip(estimate).enumerate() {
158        if valid_mask.is_some_and(|mask| mask[index] == 0) {
159            continue;
160        }
161        if !reference.is_finite() || !estimate.is_finite() {
162            return Err(crate::Error::Numerical(
163                "intensity comparison contains a non-finite value".into(),
164            ));
165        }
166        accumulator.push(reference, estimate, saturation_value);
167        count += 1;
168    }
169    if count == 0 {
170        return Err(crate::Error::InvalidParameter {
171            name: "valid_mask",
172            reason: "must select at least one pixel".into(),
173        });
174    }
175    Ok(accumulator.finish(count, saturation_value))
176}
177
178#[derive(Default)]
179struct ComparisonAccumulator {
180    reference_sum: f64,
181    estimate_sum: f64,
182    residual_l1: f64,
183    residual_squared: f64,
184    residual_sum: f64,
185    residual_max_abs: f64,
186    reference_squared: f64,
187    saturated_pixels: usize,
188}
189
190impl ComparisonAccumulator {
191    fn push(&mut self, reference: f64, estimate: f64, saturation_value: Option<f64>) {
192        let residual = estimate - reference;
193        self.reference_sum += reference;
194        self.estimate_sum += estimate;
195        self.residual_l1 += residual.abs();
196        self.residual_squared += residual * residual;
197        self.residual_sum += residual;
198        self.residual_max_abs = self.residual_max_abs.max(residual.abs());
199        self.reference_squared += reference * reference;
200        self.saturated_pixels +=
201            usize::from(saturation_value.is_some_and(|limit| reference >= limit));
202    }
203
204    fn finish(self, count: usize, saturation_value: Option<f64>) -> IntensityComparisonMetrics {
205        debug_assert!(count > 0);
206        let count = count as f64;
207        let residual_mean = self.residual_sum / count;
208        IntensityComparisonMetrics {
209            reference_sum: self.reference_sum,
210            estimate_sum: self.estimate_sum,
211            residual_l1: self.residual_l1,
212            residual_l2: self.residual_squared.sqrt(),
213            residual_mean,
214            residual_std: (self.residual_squared / count - residual_mean * residual_mean)
215                .max(0.0)
216                .sqrt(),
217            residual_max_abs: self.residual_max_abs,
218            normalized_l2: self.residual_squared.sqrt()
219                / (self.reference_squared.sqrt() + f64::EPSILON),
220            saturated_pixels: saturation_value.map(|_| self.saturated_pixels),
221        }
222    }
223}
224
225/// Return the mean signed residual, `estimate - reference`.
226pub fn bias<T>(
227    reference: ArrayView2<'_, T>,
228    estimate: ArrayView2<'_, T>,
229    valid_mask: Option<ArrayView2<'_, bool>>,
230) -> Result<f64, IntensityMetricError>
231where
232    T: ToPrimitive,
233{
234    let mut sum = 0.0;
235    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
236        sum += estimate - reference;
237    })?;
238    Ok(sum / count as f64)
239}
240
241/// Return the mean absolute error.
242pub fn mae<T>(
243    reference: ArrayView2<'_, T>,
244    estimate: ArrayView2<'_, T>,
245    valid_mask: Option<ArrayView2<'_, bool>>,
246) -> Result<f64, IntensityMetricError>
247where
248    T: ToPrimitive,
249{
250    let mut sum = 0.0;
251    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
252        sum += (estimate - reference).abs();
253    })?;
254    Ok(sum / count as f64)
255}
256
257/// Return the mean squared error.
258pub fn mse<T>(
259    reference: ArrayView2<'_, T>,
260    estimate: ArrayView2<'_, T>,
261    valid_mask: Option<ArrayView2<'_, bool>>,
262) -> Result<f64, IntensityMetricError>
263where
264    T: ToPrimitive,
265{
266    let (sum, count) = squared_error_sum(reference, estimate, valid_mask)?;
267    Ok(sum / count as f64)
268}
269
270/// Return the root mean squared error.
271pub fn rmse<T>(
272    reference: ArrayView2<'_, T>,
273    estimate: ArrayView2<'_, T>,
274    valid_mask: Option<ArrayView2<'_, bool>>,
275) -> Result<f64, IntensityMetricError>
276where
277    T: ToPrimitive,
278{
279    Ok(mse(reference, estimate, valid_mask)?.sqrt())
280}
281
282/// Return `sum(abs(estimate - reference)) / sum(abs(reference))`.
283pub fn relative_l1<T>(
284    reference: ArrayView2<'_, T>,
285    estimate: ArrayView2<'_, T>,
286    valid_mask: Option<ArrayView2<'_, bool>>,
287) -> Result<f64, IntensityMetricError>
288where
289    T: ToPrimitive,
290{
291    let mut residual_sum = 0.0;
292    let mut reference_sum = 0.0;
293    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
294        residual_sum += (estimate - reference).abs();
295        reference_sum += reference.abs();
296    })?;
297    if reference_sum == 0.0 {
298        return Err(IntensityMetricError::ZeroReferenceNormalization {
299            metric: "relative_l1",
300        });
301    }
302    Ok(residual_sum / reference_sum)
303}
304
305/// Return `||estimate - reference||_2 / ||reference||_2`.
306pub fn nrmse<T>(
307    reference: ArrayView2<'_, T>,
308    estimate: ArrayView2<'_, T>,
309    valid_mask: Option<ArrayView2<'_, bool>>,
310) -> Result<f64, IntensityMetricError>
311where
312    T: ToPrimitive,
313{
314    let mut residual_squared = 0.0;
315    let mut reference_squared = 0.0;
316    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
317        residual_squared += (estimate - reference).powi(2);
318        reference_squared += reference.powi(2);
319    })?;
320    if reference_squared == 0.0 {
321        return Err(IntensityMetricError::ZeroReferenceNormalization { metric: "nrmse" });
322    }
323    Ok(residual_squared.sqrt() / reference_squared.sqrt())
324}
325
326/// Compare square-root intensities, normalized by the reference amplitude L2 norm.
327pub fn amplitude_nrmse<T>(
328    reference: ArrayView2<'_, T>,
329    estimate: ArrayView2<'_, T>,
330    valid_mask: Option<ArrayView2<'_, bool>>,
331) -> Result<f64, IntensityMetricError>
332where
333    T: ToPrimitive,
334{
335    validate_non_negative(reference, estimate, valid_mask, "amplitude_nrmse")?;
336    let mut residual_squared = 0.0;
337    let mut reference_squared = 0.0;
338    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
339        let reference_amplitude = reference.sqrt();
340        let estimate_amplitude = estimate.sqrt();
341        residual_squared += (estimate_amplitude - reference_amplitude).powi(2);
342        reference_squared += reference;
343    })?;
344    if reference_squared == 0.0 {
345        return Err(IntensityMetricError::ZeroReferenceNormalization {
346            metric: "amplitude_nrmse",
347        });
348    }
349    Ok(residual_squared.sqrt() / reference_squared.sqrt())
350}
351
352/// Return the Pearson correlation coefficient of valid pixels.
353pub fn correlation<T>(
354    reference: ArrayView2<'_, T>,
355    estimate: ArrayView2<'_, T>,
356    valid_mask: Option<ArrayView2<'_, bool>>,
357) -> Result<f64, IntensityMetricError>
358where
359    T: ToPrimitive,
360{
361    let mut reference_sum = 0.0;
362    let mut estimate_sum = 0.0;
363    let mut reference_squared = 0.0;
364    let mut estimate_squared = 0.0;
365    let mut product_sum = 0.0;
366    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
367        reference_sum += reference;
368        estimate_sum += estimate;
369        reference_squared += reference * reference;
370        estimate_squared += estimate * estimate;
371        product_sum += reference * estimate;
372    })? as f64;
373    let reference_variance = reference_squared - reference_sum * reference_sum / count;
374    let estimate_variance = estimate_squared - estimate_sum * estimate_sum / count;
375    if reference_variance <= 0.0 || estimate_variance <= 0.0 {
376        return Err(IntensityMetricError::ZeroVariance);
377    }
378    let covariance = product_sum - reference_sum * estimate_sum / count;
379    Ok(covariance / (reference_variance * estimate_variance).sqrt())
380}
381
382/// Return peak signal-to-noise ratio in dB for an explicit intensity range.
383pub fn psnr<T>(
384    reference: ArrayView2<'_, T>,
385    estimate: ArrayView2<'_, T>,
386    valid_mask: Option<ArrayView2<'_, bool>>,
387    data_range: f64,
388) -> Result<f64, IntensityMetricError>
389where
390    T: ToPrimitive,
391{
392    validate_data_range(data_range)?;
393    let value = mse(reference, estimate, valid_mask)?;
394    if value == 0.0 {
395        Ok(f64::INFINITY)
396    } else {
397        Ok(10.0 * (data_range * data_range / value).log10())
398    }
399}
400
401/// Return canonical single-scale SSIM using an 11×11 Gaussian window (σ=1.5).
402///
403/// # References
404///
405/// [Z. Wang, A. C. Bovik, H. R. Sheikh, and E. P. Simoncelli, “Image quality
406/// assessment: From error visibility to structural similarity”
407/// (2004)](https://doi.org/10.1109/TIP.2003.819861), *IEEE Transactions on
408/// Image Processing* **13**(4), 600–612.
409pub fn ssim<T>(
410    reference: ArrayView2<'_, T>,
411    estimate: ArrayView2<'_, T>,
412    valid_mask: Option<ArrayView2<'_, bool>>,
413    data_range: f64,
414) -> Result<f64, IntensityMetricError>
415where
416    T: ToPrimitive,
417{
418    validate_data_range(data_range)?;
419    let shape = reference.dim();
420    validate_shapes(
421        shape,
422        estimate.dim(),
423        valid_mask.as_ref().map(|mask| mask.dim()),
424    )?;
425    if shape.0 < SSIM_WINDOW_SIZE || shape.1 < SSIM_WINDOW_SIZE {
426        return Err(IntensityMetricError::SsimImageTooSmall { shape });
427    }
428    // Validate every included pixel even if the mask leaves no complete SSIM window.
429    for_each_valid_pair(reference, estimate, valid_mask, |_, _| {})?;
430
431    let weights = gaussian_weights();
432    let c1 = (SSIM_K1 * data_range).powi(2);
433    let c2 = (SSIM_K2 * data_range).powi(2);
434    let mut score_sum = 0.0;
435    let mut window_count = 0usize;
436    for row in 0..=shape.0 - SSIM_WINDOW_SIZE {
437        for column in 0..=shape.1 - SSIM_WINDOW_SIZE {
438            if valid_mask.as_ref().is_some_and(|mask| {
439                (row..row + SSIM_WINDOW_SIZE).any(|window_row| {
440                    (column..column + SSIM_WINDOW_SIZE)
441                        .any(|window_column| !mask[(window_row, window_column)])
442                })
443            }) {
444                continue;
445            }
446            let (mut reference_mean, mut estimate_mean) = (0.0, 0.0);
447            for window_row in 0..SSIM_WINDOW_SIZE {
448                for window_column in 0..SSIM_WINDOW_SIZE {
449                    let weight = weights[window_row * SSIM_WINDOW_SIZE + window_column];
450                    reference_mean += weight
451                        * value_as_f64(
452                            &reference[(row + window_row, column + window_column)],
453                            "reference",
454                        )?;
455                    estimate_mean += weight
456                        * value_as_f64(
457                            &estimate[(row + window_row, column + window_column)],
458                            "estimate",
459                        )?;
460                }
461            }
462            let (mut reference_variance, mut estimate_variance, mut covariance) = (0.0, 0.0, 0.0);
463            for window_row in 0..SSIM_WINDOW_SIZE {
464                for window_column in 0..SSIM_WINDOW_SIZE {
465                    let weight = weights[window_row * SSIM_WINDOW_SIZE + window_column];
466                    let reference_value = value_as_f64(
467                        &reference[(row + window_row, column + window_column)],
468                        "reference",
469                    )? - reference_mean;
470                    let estimate_value = value_as_f64(
471                        &estimate[(row + window_row, column + window_column)],
472                        "estimate",
473                    )? - estimate_mean;
474                    reference_variance += weight * reference_value * reference_value;
475                    estimate_variance += weight * estimate_value * estimate_value;
476                    covariance += weight * reference_value * estimate_value;
477                }
478            }
479            score_sum += ((2.0 * reference_mean * estimate_mean + c1) * (2.0 * covariance + c2))
480                / ((reference_mean.powi(2) + estimate_mean.powi(2) + c1)
481                    * (reference_variance + estimate_variance + c2));
482            window_count += 1;
483        }
484    }
485    if window_count == 0 {
486        return Err(IntensityMetricError::NoValidSsimWindow);
487    }
488    Ok(score_sum / window_count as f64)
489}
490
491/// Return summed Poisson deviance, flooring estimate intensities at `epsilon`.
492pub fn poisson_deviance<T>(
493    reference: ArrayView2<'_, T>,
494    estimate: ArrayView2<'_, T>,
495    valid_mask: Option<ArrayView2<'_, bool>>,
496    epsilon: f64,
497) -> Result<f64, IntensityMetricError>
498where
499    T: ToPrimitive,
500{
501    validate_epsilon(epsilon)?;
502    validate_non_negative(reference, estimate, valid_mask, "poisson_deviance")?;
503    let mut sum = 0.0;
504    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
505        if reference == 0.0 && estimate == 0.0 {
506            return;
507        }
508        let estimate = estimate.max(epsilon);
509        let log_term = if reference == 0.0 {
510            0.0
511        } else {
512            reference * (reference / estimate).ln()
513        };
514        sum += 2.0 * (estimate - reference + log_term);
515    })?;
516    Ok(sum)
517}
518
519/// Return mean Poisson deviance, flooring estimate intensities at `epsilon`.
520pub fn mean_poisson_deviance<T>(
521    reference: ArrayView2<'_, T>,
522    estimate: ArrayView2<'_, T>,
523    valid_mask: Option<ArrayView2<'_, bool>>,
524    epsilon: f64,
525) -> Result<f64, IntensityMetricError>
526where
527    T: ToPrimitive,
528{
529    validate_epsilon(epsilon)?;
530    validate_non_negative(reference, estimate, valid_mask, "mean_poisson_deviance")?;
531    let mut sum = 0.0;
532    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
533        if reference == 0.0 && estimate == 0.0 {
534            return;
535        }
536        let estimate = estimate.max(epsilon);
537        let log_term = if reference == 0.0 {
538            0.0
539        } else {
540            reference * (reference / estimate).ln()
541        };
542        sum += 2.0 * (estimate - reference + log_term);
543    })?;
544    Ok(sum / count as f64)
545}
546
547/// Fit the least-squares scalar `gain` in `estimate ≈ gain × reference`.
548pub fn fitted_gain<T>(
549    reference: ArrayView2<'_, T>,
550    estimate: ArrayView2<'_, T>,
551    valid_mask: Option<ArrayView2<'_, bool>>,
552) -> Result<f64, IntensityMetricError>
553where
554    T: ToPrimitive,
555{
556    let mut numerator = 0.0;
557    let mut denominator = 0.0;
558    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
559        numerator += reference * estimate;
560        denominator += reference * reference;
561    })?;
562    if denominator == 0.0 {
563        return Err(IntensityMetricError::ZeroReferenceNormalization {
564            metric: "fitted_gain",
565        });
566    }
567    Ok(numerator / denominator)
568}
569
570fn squared_error_sum<T>(
571    reference: ArrayView2<'_, T>,
572    estimate: ArrayView2<'_, T>,
573    valid_mask: Option<ArrayView2<'_, bool>>,
574) -> Result<(f64, usize), IntensityMetricError>
575where
576    T: ToPrimitive,
577{
578    let mut sum = 0.0;
579    let count = for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
580        sum += (estimate - reference).powi(2);
581    })?;
582    Ok((sum, count))
583}
584
585fn for_each_valid_pair<T>(
586    reference: ArrayView2<'_, T>,
587    estimate: ArrayView2<'_, T>,
588    valid_mask: Option<ArrayView2<'_, bool>>,
589    mut function: impl FnMut(f64, f64),
590) -> Result<usize, IntensityMetricError>
591where
592    T: ToPrimitive,
593{
594    let shape = reference.dim();
595    validate_shapes(
596        shape,
597        estimate.dim(),
598        valid_mask.as_ref().map(|mask| mask.dim()),
599    )?;
600    let mut count = 0usize;
601    for row in 0..shape.0 {
602        for column in 0..shape.1 {
603            if valid_mask.as_ref().is_some_and(|mask| !mask[(row, column)]) {
604                continue;
605            }
606            let reference = value_as_f64(&reference[(row, column)], "reference")?;
607            let estimate = value_as_f64(&estimate[(row, column)], "estimate")?;
608            function(reference, estimate);
609            count += 1;
610        }
611    }
612    if count == 0 {
613        return Err(IntensityMetricError::EmptyValidMask);
614    }
615    Ok(count)
616}
617
618fn validate_non_negative<T>(
619    reference: ArrayView2<'_, T>,
620    estimate: ArrayView2<'_, T>,
621    valid_mask: Option<ArrayView2<'_, bool>>,
622    metric: &'static str,
623) -> Result<(), IntensityMetricError>
624where
625    T: ToPrimitive,
626{
627    let mut negative = false;
628    for_each_valid_pair(reference, estimate, valid_mask, |reference, estimate| {
629        negative |= reference < 0.0 || estimate < 0.0;
630    })?;
631    if negative {
632        return Err(IntensityMetricError::NegativeIntensity { metric });
633    }
634    Ok(())
635}
636
637fn validate_shapes(
638    reference: (usize, usize),
639    estimate: (usize, usize),
640    valid_mask: Option<(usize, usize)>,
641) -> Result<(), IntensityMetricError> {
642    if reference != estimate {
643        return Err(IntensityMetricError::ShapeMismatch {
644            reference,
645            estimate,
646        });
647    }
648    if let Some(actual) = valid_mask
649        && actual != reference
650    {
651        return Err(IntensityMetricError::MaskShapeMismatch {
652            actual,
653            expected: reference,
654        });
655    }
656    Ok(())
657}
658
659fn value_as_f64<T>(value: &T, input: &'static str) -> Result<f64, IntensityMetricError>
660where
661    T: ToPrimitive,
662{
663    let value = value
664        .to_f64()
665        .ok_or(IntensityMetricError::UnsupportedScalar { input })?;
666    if !value.is_finite() {
667        return Err(IntensityMetricError::NonFinite { input });
668    }
669    Ok(value)
670}
671
672fn validate_data_range(data_range: f64) -> Result<(), IntensityMetricError> {
673    if !data_range.is_finite() || data_range <= 0.0 {
674        return Err(IntensityMetricError::InvalidDataRange);
675    }
676    Ok(())
677}
678
679fn validate_epsilon(epsilon: f64) -> Result<(), IntensityMetricError> {
680    if !epsilon.is_finite() || epsilon <= 0.0 {
681        return Err(IntensityMetricError::InvalidEpsilon);
682    }
683    Ok(())
684}
685
686fn gaussian_weights() -> Vec<f64> {
687    let center = (SSIM_WINDOW_SIZE - 1) as f64 / 2.0;
688    let mut weights = Vec::with_capacity(SSIM_WINDOW_SIZE * SSIM_WINDOW_SIZE);
689    for row in 0..SSIM_WINDOW_SIZE {
690        for column in 0..SSIM_WINDOW_SIZE {
691            let squared_distance = (row as f64 - center).powi(2) + (column as f64 - center).powi(2);
692            weights.push((-squared_distance / (2.0 * SSIM_GAUSSIAN_SIGMA.powi(2))).exp());
693        }
694    }
695    let normalization = weights.iter().sum::<f64>();
696    for weight in &mut weights {
697        *weight /= normalization;
698    }
699    weights
700}
701
702#[cfg(test)]
703mod tests {
704    use approx::assert_relative_eq;
705    use ndarray::array;
706
707    use super::*;
708
709    #[test]
710    fn residual_metrics_use_estimate_minus_reference_and_support_integers() {
711        let reference = array![[1u8, 2u8]];
712        let estimate = array![[2u8, 0u8]];
713        assert_relative_eq!(bias(reference.view(), estimate.view(), None).unwrap(), -0.5);
714        assert_relative_eq!(mae(reference.view(), estimate.view(), None).unwrap(), 1.5);
715        assert_relative_eq!(mse(reference.view(), estimate.view(), None).unwrap(), 2.5);
716        assert_relative_eq!(
717            rmse(reference.view(), estimate.view(), None).unwrap(),
718            2.5f64.sqrt()
719        );
720        assert_relative_eq!(
721            relative_l1(reference.view(), estimate.view(), None).unwrap(),
722            1.0
723        );
724        assert_relative_eq!(nrmse(reference.view(), estimate.view(), None).unwrap(), 1.0);
725        assert_relative_eq!(
726            correlation(reference.view(), estimate.view(), None).unwrap(),
727            -1.0
728        );
729    }
730
731    #[test]
732    fn mask_selects_valid_pixels() {
733        let reference = array![[1.0, 20.0]];
734        let estimate = array![[2.0, 0.0]];
735        let mask = array![[true, false]];
736        assert_relative_eq!(
737            bias(reference.view(), estimate.view(), Some(mask.view())).unwrap(),
738            1.0
739        );
740    }
741
742    #[test]
743    fn quality_metrics_have_documented_reference_values() {
744        let reference = array![[0.0, 1.0]];
745        let estimate = array![[0.0, 0.0]];
746        assert_relative_eq!(
747            psnr(reference.view(), estimate.view(), None, 1.0).unwrap(),
748            10.0 * 2.0f64.log10()
749        );
750
751        let image = ndarray::Array2::<f64>::ones((11, 11));
752        assert_relative_eq!(ssim(image.view(), image.view(), None, 1.0).unwrap(), 1.0);
753    }
754
755    #[test]
756    fn poisson_deviance_and_gain_are_well_defined() {
757        let reference = array![[0.0, 2.0]];
758        let estimate = array![[0.0, 2.0]];
759        assert_relative_eq!(
760            poisson_deviance(reference.view(), estimate.view(), None, 1e-12).unwrap(),
761            0.0
762        );
763        assert_relative_eq!(
764            mean_poisson_deviance(reference.view(), estimate.view(), None, 1e-12).unwrap(),
765            0.0
766        );
767
768        let scaled = array![[0.0, 4.0]];
769        assert_relative_eq!(
770            fitted_gain(reference.view(), scaled.view(), None).unwrap(),
771            2.0
772        );
773    }
774
775    #[test]
776    fn validation_covers_masks_ranges_and_intensity_domains() {
777        let reference = array![[1.0, 2.0]];
778        let estimate = array![[2.0, 0.0]];
779        let wrong_mask = array![[true], [false]];
780        assert!(matches!(
781            mae(reference.view(), estimate.view(), Some(wrong_mask.view())),
782            Err(IntensityMetricError::MaskShapeMismatch { .. })
783        ));
784        assert!(matches!(
785            psnr(reference.view(), estimate.view(), None, 0.0),
786            Err(IntensityMetricError::InvalidDataRange)
787        ));
788        assert!(matches!(
789            ssim(reference.view(), estimate.view(), None, 1.0),
790            Err(IntensityMetricError::SsimImageTooSmall { .. })
791        ));
792
793        let negative = array![[-1.0, 2.0]];
794        assert!(matches!(
795            amplitude_nrmse(negative.view(), estimate.view(), None),
796            Err(IntensityMetricError::NegativeIntensity { .. })
797        ));
798        assert!(matches!(
799            poisson_deviance(negative.view(), estimate.view(), None, 1e-12),
800            Err(IntensityMetricError::NegativeIntensity { .. })
801        ));
802    }
803}