fpm_rs/metrics/complex_field/
comparison.rs1use 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}