Skip to main content

fpm_rs/metrics/complex_field/
comparison.rs

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