1use ndarray::ArrayView2;
8use num_complex::{Complex, Complex64};
9use num_traits::{Float, ToPrimitive};
10use thiserror::Error;
11
12#[derive(Clone, Copy, Debug, Eq, PartialEq)]
17pub enum ComplexAlignment {
18 None,
20 GlobalPhase,
22 Scale,
24 ComplexGain,
26}
27
28#[derive(Debug, Error, PartialEq)]
30pub enum ComplexMetricError {
31 #[error("reference and estimate shapes differ: {reference:?} and {estimate:?}")]
33 ShapeMismatch {
34 reference: (usize, usize),
36 estimate: (usize, usize),
38 },
39 #[error("mask shape {actual:?} does not match image shape {expected:?}")]
41 MaskShapeMismatch {
42 actual: (usize, usize),
44 expected: (usize, usize),
46 },
47 #[error("at least one selected pixel is required")]
49 EmptyMask,
50 #[error("{input} contains a component that cannot be represented as f64")]
52 UnsupportedScalar {
53 input: &'static str,
55 },
56 #[error("{input} contains a non-finite real or imaginary component")]
58 NonFinite {
59 input: &'static str,
61 },
62 #[error("{metric} is undefined because the reference normalization is zero")]
64 ZeroReferenceNormalization {
65 metric: &'static str,
67 },
68 #[error("{metric} is undefined because the aligned estimate normalization is zero")]
70 ZeroEstimateNormalization {
71 metric: &'static str,
73 },
74 #[error("{alignment:?} alignment is undefined: {reason}")]
76 DegenerateAlignment {
77 alignment: ComplexAlignment,
79 reason: &'static str,
81 },
82}
83
84pub 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
110pub 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
136pub 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
153pub 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
169pub 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
202pub 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
233pub 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
266pub 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}