Skip to main content

fpm_rs/metrics/complex_field/
comparison.rs

1//! Metrics comparing a reference and candidate complex field.
2
3use ndarray::{Array2, ArrayView2};
4use num_complex::Complex64;
5use serde::{Deserialize, Serialize};
6
7use crate::{
8    Result,
9    backend::{Backend, CpuBackend, FftDirection},
10    error::Error,
11};
12
13/// Aggregate errors between reference and globally phase-aligned complex fields.
14#[derive(Clone, Debug, Default, Serialize, Deserialize)]
15pub struct ComplexFieldComparisonMetrics {
16    /// Root-mean-square amplitude error.
17    pub amplitude_rmse: f64,
18    /// Amplitude L2 error normalized by reference-field L2 norm.
19    pub amplitude_nrmse: f64,
20    /// Root-mean-square complex residual after global-phase alignment.
21    pub complex_rmse: f64,
22    /// Complex residual L2 norm divided by reference-field L2 norm.
23    pub complex_nrmse: f64,
24    /// Root-mean-square wrapped phase error in radians over non-dark reference pixels.
25    pub phase_rmse: f64,
26    /// Mean absolute wrapped phase error in radians over non-dark reference pixels.
27    pub phase_mae: f64,
28    /// Centered Fourier-spectrum residual L2 norm divided by reference-spectrum norm.
29    pub fourier_nrmse: f64,
30    /// Fitted global candidate-to-reference phase offset in radians.
31    pub global_phase_offset: f64,
32}
33
34/// Compares equal-shaped fields after fitting one global phase offset.
35pub fn compare_complex_fields(
36    reference: ArrayView2<'_, Complex64>,
37    candidate: ArrayView2<'_, Complex64>,
38) -> Result<ComplexFieldComparisonMetrics> {
39    compare_complex_fields_masked(reference, candidate, None)
40}
41
42/// Compares equal-shaped fields over non-zero entries of an optional byte mask.
43pub fn compare_complex_fields_masked(
44    reference: ArrayView2<'_, Complex64>,
45    candidate: ArrayView2<'_, Complex64>,
46    valid_mask: Option<ArrayView2<'_, u8>>,
47) -> Result<ComplexFieldComparisonMetrics> {
48    if reference.dim() != candidate.dim() || reference.is_empty() {
49        return Err(Error::InvalidShape(format!(
50            "reference shape {:?} differs from candidate {:?}",
51            reference.dim(),
52            candidate.dim()
53        )));
54    }
55    if valid_mask.is_some_and(|mask| mask.dim() != reference.dim()) {
56        return Err(Error::InvalidShape(
57            "complex-field mask shape differs from inputs".into(),
58        ));
59    }
60    let mask_values: Vec<u8> = valid_mask
61        .map(|mask| mask.iter().copied().collect())
62        .unwrap_or_else(|| vec![1; reference.len()]);
63    let count = mask_values.iter().filter(|&&value| value != 0).count();
64    if count == 0 {
65        return Err(Error::InvalidParameter {
66            name: "valid_mask",
67            reason: "must select at least one field sample".into(),
68        });
69    }
70    let cross: Complex64 = reference
71        .iter()
72        .zip(candidate.iter())
73        .zip(&mask_values)
74        .filter(|&(_, &valid)| valid != 0)
75        .map(|((&reference, &candidate), _)| candidate * reference.conj())
76        .sum();
77    let global_phase_offset = cross.arg();
78    let correction = Complex64::from_polar(1.0, -global_phase_offset);
79    let mut amplitude_squared = 0.0;
80    let mut complex_squared = 0.0;
81    let mut phase_squared = 0.0;
82    let mut phase_absolute = 0.0;
83    let mut reference_amplitude_squared = 0.0;
84    let mut reference_complex_squared = 0.0;
85    for ((&reference, &candidate), &valid) in
86        reference.iter().zip(candidate.iter()).zip(&mask_values)
87    {
88        if valid == 0 {
89            continue;
90        }
91        let aligned = candidate * correction;
92        amplitude_squared += (aligned.norm() - reference.norm()).powi(2);
93        complex_squared += (aligned - reference).norm_sqr();
94        reference_amplitude_squared += reference.norm_sqr();
95        reference_complex_squared += reference.norm_sqr();
96        let phase_error = wrap_phase(aligned.arg() - reference.arg());
97        phase_squared += phase_error * phase_error;
98        phase_absolute += phase_error.abs();
99    }
100    let reference_spectrum = spectrum(masked(reference, valid_mask)?.view())?;
101    let candidate_spectrum = spectrum(masked(candidate, valid_mask)?.view())?;
102    let fourier_squared: f64 = reference_spectrum
103        .iter()
104        .zip(candidate_spectrum.iter())
105        .map(|(&reference, &candidate)| (candidate * correction - reference).norm_sqr())
106        .sum();
107    let fourier_reference_squared: f64 = reference_spectrum
108        .iter()
109        .map(|value| value.norm_sqr())
110        .sum();
111    let count = count as f64;
112    let amplitude_rmse = (amplitude_squared / count).sqrt();
113    let complex_rmse = (complex_squared / count).sqrt();
114    Ok(ComplexFieldComparisonMetrics {
115        amplitude_rmse,
116        amplitude_nrmse: amplitude_rmse
117            / (reference_amplitude_squared / count)
118                .sqrt()
119                .max(f64::EPSILON),
120        complex_rmse,
121        complex_nrmse: complex_rmse / (reference_complex_squared / count).sqrt().max(f64::EPSILON),
122        phase_rmse: (phase_squared / count).sqrt(),
123        phase_mae: phase_absolute / count,
124        fourier_nrmse: (fourier_squared / fourier_reference_squared.max(f64::EPSILON)).sqrt(),
125        global_phase_offset,
126    })
127}
128
129fn masked(
130    values: ArrayView2<'_, Complex64>,
131    mask: Option<ArrayView2<'_, u8>>,
132) -> Result<Array2<Complex64>> {
133    let data = match mask {
134        Some(mask) => values
135            .iter()
136            .zip(mask.iter())
137            .map(|(&value, &valid)| {
138                if valid != 0 {
139                    value
140                } else {
141                    Complex64::default()
142                }
143            })
144            .collect(),
145        None => values.iter().copied().collect(),
146    };
147    Ok(Array2::from_shape_vec(values.dim(), data)?)
148}
149
150fn spectrum(values: ArrayView2<'_, Complex64>) -> Result<Array2<Complex64>> {
151    let shape = values.dim();
152    let backend = CpuBackend::new(shape, shape)?;
153    let mut spectrum: Vec<_> = values.iter().copied().collect();
154    let mut column = vec![Complex64::default(); shape.0];
155    backend.fft2(&mut spectrum, shape, FftDirection::Forward, &mut column)?;
156    Ok(Array2::from_shape_vec(shape, spectrum)?)
157}
158
159fn wrap_phase(value: f64) -> f64 {
160    (value + std::f64::consts::PI).rem_euclid(std::f64::consts::TAU) - std::f64::consts::PI
161}