fpm_rs/metrics/complex_field/
comparison.rs1use 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#[derive(Clone, Debug, Default, Serialize, Deserialize)]
15pub struct ComplexFieldComparisonMetrics {
16 pub amplitude_rmse: f64,
18 pub amplitude_nrmse: f64,
20 pub complex_rmse: f64,
22 pub complex_nrmse: f64,
24 pub phase_rmse: f64,
26 pub phase_mae: f64,
28 pub fourier_nrmse: f64,
30 pub global_phase_offset: f64,
32}
33
34pub 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
42pub 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}