Skip to main content

fpm_rs/metrics/intensity/
atomic.rs

1//! Atomic, domain-agnostic metrics for comparing scalar intensity images.
2//!
3//! Inputs use `reference` and `candidate` terminology.  Signed quantities use
4//! the residual `candidate - reference`; `valid_mask`, when supplied, includes
5//! pixels whose value is `true`.  These functions are for evaluation, not
6//! optimization objectives.
7
8use 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/// Errors returned by atomic intensity metrics.
18#[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
56/// Return the mean signed residual, `candidate - reference`.
57pub 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
72/// Return the mean absolute error.
73pub 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
88/// Return the mean squared error.
89pub 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
101/// Return the root mean squared error.
102pub 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
113/// Return `sum(abs(candidate - reference)) / sum(abs(reference))`.
114pub 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
136/// Return `||candidate - reference||_2 / ||reference||_2`.
137pub 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
157/// Compare square-root intensities, normalized by the reference amplitude L2 norm.
158pub 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
183/// Return the Pearson correlation coefficient of valid pixels.
184pub 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
213/// Return peak signal-to-noise ratio in dB for an explicit intensity range.
214pub 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
232/// Return canonical single-scale SSIM using an 11×11 Gaussian window (σ=1.5).
233pub 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    // Validate every included pixel even if the mask leaves no complete SSIM window.
253    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
315/// Return summed Poisson deviance, flooring candidate intensities at `epsilon`.
316pub 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
343/// Return mean Poisson deviance, flooring candidate intensities at `epsilon`.
344pub 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
376/// Fit the least-squares scalar `gain` in `candidate ≈ gain × reference`.
377pub 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}