Skip to main content

fpm_rs/metrics/complex_field/
compare.rs

1//! Metrics comparing reference and estimate complex-valued images.
2//!
3//! Every metric applies the requested alignment to the estimate without
4//! allocating an aligned image. A mask value of `true` includes that pixel in
5//! both the alignment fit and metric accumulation.
6
7use ndarray::ArrayView2;
8use num_complex::{Complex, Complex64};
9use num_traits::{Float, ToPrimitive};
10use thiserror::Error;
11
12/// Alignment applied to an estimate before evaluating a complex-image metric.
13///
14/// Every fitted mode minimizes the complex least-squares residual over the
15/// selected pixels. Alignment is always explicit at the call site.
16#[derive(Clone, Copy, Debug, Eq, PartialEq)]
17pub enum ComplexAlignment {
18    /// Compare the estimate directly with the reference.
19    None,
20    /// Fit a unit-magnitude complex scalar, removing one global phase offset.
21    GlobalPhase,
22    /// Fit a non-negative real scalar, preserving phase.
23    Scale,
24    /// Fit an unconstrained complex scalar, removing scale and global phase.
25    ComplexGain,
26}
27
28/// Errors returned by explicit-alignment complex-image metrics.
29#[derive(Debug, Error, PartialEq)]
30pub enum ComplexMetricError {
31    /// Reference and estimate have different `(height, width)` shapes.
32    #[error("reference and estimate shapes differ: {reference:?} and {estimate:?}")]
33    ShapeMismatch {
34        /// Reference field shape.
35        reference: (usize, usize),
36        /// Estimate field shape.
37        estimate: (usize, usize),
38    },
39    /// A selection mask does not match the field shape.
40    #[error("mask shape {actual:?} does not match image shape {expected:?}")]
41    MaskShapeMismatch {
42        /// Supplied mask shape.
43        actual: (usize, usize),
44        /// Required field shape.
45        expected: (usize, usize),
46    },
47    /// No pixel was selected by the optional mask.
48    #[error("at least one selected pixel is required")]
49    EmptyMask,
50    /// A generic complex component could not be represented as `f64`.
51    #[error("{input} contains a component that cannot be represented as f64")]
52    UnsupportedScalar {
53        /// Name of the input containing the unsupported component.
54        input: &'static str,
55    },
56    /// An input contains a non-finite real or imaginary component.
57    #[error("{input} contains a non-finite real or imaginary component")]
58    NonFinite {
59        /// Name of the non-finite input.
60        input: &'static str,
61    },
62    /// A relative metric has zero reference norm.
63    #[error("{metric} is undefined because the reference normalization is zero")]
64    ZeroReferenceNormalization {
65        /// Metric whose reference denominator was zero.
66        metric: &'static str,
67    },
68    /// A normalized correlation has zero aligned-estimate norm.
69    #[error("{metric} is undefined because the aligned estimate normalization is zero")]
70    ZeroEstimateNormalization {
71        /// Metric whose estimate denominator was zero.
72        metric: &'static str,
73    },
74    /// The selected alignment cannot be fitted to the selected pixels.
75    #[error("{alignment:?} alignment is undefined: {reason}")]
76    DegenerateAlignment {
77        /// Requested alignment mode.
78        alignment: ComplexAlignment,
79        /// Explanation of the degenerate fit.
80        reason: &'static str,
81    },
82}
83
84/// Return the arithmetic mean of `aligned_estimate - reference`.
85///
86/// The selected [`ComplexAlignment`] is fitted over exactly the pixels
87/// included by `mask`; `true` includes a pixel.
88pub fn bias<T>(
89    reference: ArrayView2<'_, Complex<T>>,
90    estimate: ArrayView2<'_, Complex<T>>,
91    mask: Option<ArrayView2<'_, bool>>,
92    alignment: ComplexAlignment,
93) -> Result<Complex64, ComplexMetricError>
94where
95    T: Float + ToPrimitive,
96{
97    let mut residual_sum = Complex64::default();
98    let count = for_each_aligned_pair(
99        reference,
100        estimate,
101        mask,
102        alignment,
103        |reference, estimate| {
104            residual_sum += estimate - reference;
105        },
106    )?;
107    Ok(residual_sum / count as f64)
108}
109
110/// Return the mean residual magnitude, `mean(|aligned_estimate - reference|)`.
111///
112/// `aligned_estimate` uses the requested [`ComplexAlignment`] fitted over the
113/// selected pixels.
114pub fn mae<T>(
115    reference: ArrayView2<'_, Complex<T>>,
116    estimate: ArrayView2<'_, Complex<T>>,
117    mask: Option<ArrayView2<'_, bool>>,
118    alignment: ComplexAlignment,
119) -> Result<f64, ComplexMetricError>
120where
121    T: Float + ToPrimitive,
122{
123    let mut residual_sum = 0.0;
124    let count = for_each_aligned_pair(
125        reference,
126        estimate,
127        mask,
128        alignment,
129        |reference, estimate| {
130            residual_sum += (estimate - reference).norm();
131        },
132    )?;
133    Ok(residual_sum / count as f64)
134}
135
136/// Return `mean(|aligned_estimate - reference|²)`.
137///
138/// `aligned_estimate` uses the requested [`ComplexAlignment`] fitted over the
139/// selected pixels.
140pub fn mse<T>(
141    reference: ArrayView2<'_, Complex<T>>,
142    estimate: ArrayView2<'_, Complex<T>>,
143    mask: Option<ArrayView2<'_, bool>>,
144    alignment: ComplexAlignment,
145) -> Result<f64, ComplexMetricError>
146where
147    T: Float + ToPrimitive,
148{
149    let (sum, count) = squared_error_sum(reference, estimate, mask, alignment)?;
150    Ok(sum / count as f64)
151}
152
153/// Return `sqrt(mean(|aligned_estimate - reference|²))`.
154///
155/// `aligned_estimate` uses the requested [`ComplexAlignment`] fitted over the
156/// selected pixels.
157pub fn rmse<T>(
158    reference: ArrayView2<'_, Complex<T>>,
159    estimate: ArrayView2<'_, Complex<T>>,
160    mask: Option<ArrayView2<'_, bool>>,
161    alignment: ComplexAlignment,
162) -> Result<f64, ComplexMetricError>
163where
164    T: Float + ToPrimitive,
165{
166    Ok(mse(reference, estimate, mask, alignment)?.sqrt())
167}
168
169/// Return `sum(|aligned_estimate - reference|) / sum(|reference|)`.
170///
171/// `aligned_estimate` uses the requested [`ComplexAlignment`] fitted over the
172/// selected pixels.
173pub fn relative_l1<T>(
174    reference: ArrayView2<'_, Complex<T>>,
175    estimate: ArrayView2<'_, Complex<T>>,
176    mask: Option<ArrayView2<'_, bool>>,
177    alignment: ComplexAlignment,
178) -> Result<f64, ComplexMetricError>
179where
180    T: Float + ToPrimitive,
181{
182    let mut residual_sum = 0.0;
183    let mut reference_sum = 0.0;
184    for_each_aligned_pair(
185        reference,
186        estimate,
187        mask,
188        alignment,
189        |reference, estimate| {
190            residual_sum += (estimate - reference).norm();
191            reference_sum += reference.norm();
192        },
193    )?;
194    if reference_sum == 0.0 {
195        return Err(ComplexMetricError::ZeroReferenceNormalization {
196            metric: "relative_l1",
197        });
198    }
199    Ok(residual_sum / reference_sum)
200}
201
202/// Return `||aligned_estimate - reference||₂ / ||reference||₂`.
203///
204/// `aligned_estimate` uses the requested [`ComplexAlignment`] fitted over the
205/// selected pixels.
206pub fn nrmse<T>(
207    reference: ArrayView2<'_, Complex<T>>,
208    estimate: ArrayView2<'_, Complex<T>>,
209    mask: Option<ArrayView2<'_, bool>>,
210    alignment: ComplexAlignment,
211) -> Result<f64, ComplexMetricError>
212where
213    T: Float + ToPrimitive,
214{
215    let mut residual_squared = 0.0;
216    let mut reference_squared = 0.0;
217    for_each_aligned_pair(
218        reference,
219        estimate,
220        mask,
221        alignment,
222        |reference, estimate| {
223            residual_squared += (estimate - reference).norm_sqr();
224            reference_squared += reference.norm_sqr();
225        },
226    )?;
227    if reference_squared == 0.0 {
228        return Err(ComplexMetricError::ZeroReferenceNormalization { metric: "nrmse" });
229    }
230    Ok((residual_squared / reference_squared).sqrt())
231}
232
233/// Compare magnitudes after alignment, normalized by reference magnitude energy.
234///
235/// This returns `|| |aligned_estimate| - |reference| ||₂ / ||reference||₂`.
236/// The requested [`ComplexAlignment`] is fitted over the selected pixels.
237pub fn amplitude_nrmse<T>(
238    reference: ArrayView2<'_, Complex<T>>,
239    estimate: ArrayView2<'_, Complex<T>>,
240    mask: Option<ArrayView2<'_, bool>>,
241    alignment: ComplexAlignment,
242) -> Result<f64, ComplexMetricError>
243where
244    T: Float + ToPrimitive,
245{
246    let mut residual_squared = 0.0;
247    let mut reference_squared = 0.0;
248    for_each_aligned_pair(
249        reference,
250        estimate,
251        mask,
252        alignment,
253        |reference, estimate| {
254            residual_squared += (estimate.norm() - reference.norm()).powi(2);
255            reference_squared += reference.norm_sqr();
256        },
257    )?;
258    if reference_squared == 0.0 {
259        return Err(ComplexMetricError::ZeroReferenceNormalization {
260            metric: "amplitude_nrmse",
261        });
262    }
263    Ok((residual_squared / reference_squared).sqrt())
264}
265
266/// Return normalized complex correlation after alignment.
267///
268/// The conjugation convention is
269/// `sum(conj(reference) * aligned_estimate) /
270/// sqrt(sum(|reference|²) * sum(|aligned_estimate|²))`. Consequently, without
271/// alignment, an estimate equal to `reference * exp(iθ)` has correlation
272/// `exp(iθ)`.
273pub fn correlation<T>(
274    reference: ArrayView2<'_, Complex<T>>,
275    estimate: ArrayView2<'_, Complex<T>>,
276    mask: Option<ArrayView2<'_, bool>>,
277    alignment: ComplexAlignment,
278) -> Result<Complex64, ComplexMetricError>
279where
280    T: Float + ToPrimitive,
281{
282    let mut product_sum = Complex64::default();
283    let mut reference_squared = 0.0;
284    let mut estimate_squared = 0.0;
285    for_each_aligned_pair(
286        reference,
287        estimate,
288        mask,
289        alignment,
290        |reference, estimate| {
291            product_sum += reference.conj() * estimate;
292            reference_squared += reference.norm_sqr();
293            estimate_squared += estimate.norm_sqr();
294        },
295    )?;
296    if reference_squared == 0.0 {
297        return Err(ComplexMetricError::ZeroReferenceNormalization {
298            metric: "correlation",
299        });
300    }
301    if estimate_squared == 0.0 {
302        return Err(ComplexMetricError::ZeroEstimateNormalization {
303            metric: "correlation",
304        });
305    }
306    Ok(product_sum / (reference_squared.sqrt() * estimate_squared.sqrt()))
307}
308
309fn squared_error_sum<T>(
310    reference: ArrayView2<'_, Complex<T>>,
311    estimate: ArrayView2<'_, Complex<T>>,
312    mask: Option<ArrayView2<'_, bool>>,
313    alignment: ComplexAlignment,
314) -> Result<(f64, usize), ComplexMetricError>
315where
316    T: Float + ToPrimitive,
317{
318    let mut sum = 0.0;
319    let count = for_each_aligned_pair(
320        reference,
321        estimate,
322        mask,
323        alignment,
324        |reference, estimate| {
325            sum += (estimate - reference).norm_sqr();
326        },
327    )?;
328    Ok((sum, count))
329}
330
331fn for_each_aligned_pair<T>(
332    reference: ArrayView2<'_, Complex<T>>,
333    estimate: ArrayView2<'_, Complex<T>>,
334    mask: Option<ArrayView2<'_, bool>>,
335    alignment: ComplexAlignment,
336    mut function: impl FnMut(Complex64, Complex64),
337) -> Result<usize, ComplexMetricError>
338where
339    T: Float + ToPrimitive,
340{
341    let (factor, count) = alignment_factor(reference, estimate, mask, alignment)?;
342    let shape = reference.dim();
343    for row in 0..shape.0 {
344        for column in 0..shape.1 {
345            if mask.as_ref().is_some_and(|mask| !mask[(row, column)]) {
346                continue;
347            }
348            let reference = value_as_complex64(reference[(row, column)], "reference")?;
349            let estimate = value_as_complex64(estimate[(row, column)], "estimate")? * factor;
350            function(reference, estimate);
351        }
352    }
353    Ok(count)
354}
355
356fn alignment_factor<T>(
357    reference: ArrayView2<'_, Complex<T>>,
358    estimate: ArrayView2<'_, Complex<T>>,
359    mask: Option<ArrayView2<'_, bool>>,
360    alignment: ComplexAlignment,
361) -> Result<(Complex64, usize), ComplexMetricError>
362where
363    T: Float + ToPrimitive,
364{
365    let shape = reference.dim();
366    validate_shapes(shape, estimate.dim(), mask.as_ref().map(|mask| mask.dim()))?;
367
368    let mut cross = Complex64::default();
369    let mut estimate_energy = 0.0;
370    let mut count = 0;
371    for row in 0..shape.0 {
372        for column in 0..shape.1 {
373            if mask.as_ref().is_some_and(|mask| !mask[(row, column)]) {
374                continue;
375            }
376            let reference = value_as_complex64(reference[(row, column)], "reference")?;
377            let estimate = value_as_complex64(estimate[(row, column)], "estimate")?;
378            if alignment != ComplexAlignment::None {
379                cross += estimate.conj() * reference;
380                estimate_energy += estimate.norm_sqr();
381            }
382            count += 1;
383        }
384    }
385    if count == 0 {
386        return Err(ComplexMetricError::EmptyMask);
387    }
388
389    let factor = match alignment {
390        ComplexAlignment::None => Complex64::new(1.0, 0.0),
391        ComplexAlignment::GlobalPhase => {
392            validate_estimate_energy(estimate_energy, alignment)?;
393            let magnitude = cross.norm();
394            if magnitude == 0.0 {
395                return Err(ComplexMetricError::DegenerateAlignment {
396                    alignment,
397                    reason: "reference/estimate cross-correlation is zero",
398                });
399            }
400            cross / magnitude
401        }
402        ComplexAlignment::Scale => {
403            validate_estimate_energy(estimate_energy, alignment)?;
404            Complex64::new((cross.re / estimate_energy).max(0.0), 0.0)
405        }
406        ComplexAlignment::ComplexGain => {
407            validate_estimate_energy(estimate_energy, alignment)?;
408            cross / estimate_energy
409        }
410    };
411    Ok((factor, count))
412}
413
414fn validate_estimate_energy(
415    estimate_energy: f64,
416    alignment: ComplexAlignment,
417) -> Result<(), ComplexMetricError> {
418    if estimate_energy == 0.0 {
419        return Err(ComplexMetricError::DegenerateAlignment {
420            alignment,
421            reason: "estimate energy is zero",
422        });
423    }
424    Ok(())
425}
426
427fn validate_shapes(
428    reference: (usize, usize),
429    estimate: (usize, usize),
430    mask: Option<(usize, usize)>,
431) -> Result<(), ComplexMetricError> {
432    if reference != estimate {
433        return Err(ComplexMetricError::ShapeMismatch {
434            reference,
435            estimate,
436        });
437    }
438    if let Some(actual) = mask
439        && actual != reference
440    {
441        return Err(ComplexMetricError::MaskShapeMismatch {
442            actual,
443            expected: reference,
444        });
445    }
446    Ok(())
447}
448
449fn value_as_complex64<T>(
450    value: Complex<T>,
451    input: &'static str,
452) -> Result<Complex64, ComplexMetricError>
453where
454    T: Float + ToPrimitive,
455{
456    let real = value
457        .re
458        .to_f64()
459        .ok_or(ComplexMetricError::UnsupportedScalar { input })?;
460    let imaginary = value
461        .im
462        .to_f64()
463        .ok_or(ComplexMetricError::UnsupportedScalar { input })?;
464    if !real.is_finite() || !imaginary.is_finite() {
465        return Err(ComplexMetricError::NonFinite { input });
466    }
467    Ok(Complex64::new(real, imaginary))
468}