1use ndarray::ArrayView2;
4use num_complex::Complex64;
5use serde::{Deserialize, Serialize};
6
7use crate::{Result, error::Error, model::ImagePlaneModel};
8
9#[derive(Clone, Debug, Serialize, Deserialize)]
11pub struct PupilComparisonMetrics {
12 pub amplitude_rmse: f64,
14 pub phase_rmse: f64,
16}
17
18pub fn compare_pupils(
20 reference: ArrayView2<'_, Complex64>,
21 candidate: ArrayView2<'_, Complex64>,
22 support: ArrayView2<'_, u8>,
23) -> Result<PupilComparisonMetrics> {
24 if reference.dim() != candidate.dim() || support.dim() != reference.dim() {
25 return Err(Error::InvalidShape(
26 "pupil comparison inputs have incompatible shapes".into(),
27 ));
28 }
29 let (numerator, denominator) = candidate
30 .iter()
31 .zip(reference.iter())
32 .zip(support.iter())
33 .filter(|&(_, &inside)| inside != 0)
34 .fold(
35 (Complex64::default(), 0.0),
36 |(numerator, denominator), ((&candidate, &reference), _)| {
37 (
38 numerator + candidate.conj() * reference,
39 denominator + candidate.norm_sqr(),
40 )
41 },
42 );
43 let alignment = if denominator > f64::EPSILON {
44 numerator / denominator
45 } else {
46 Complex64::new(1.0, 0.0)
47 };
48 let mut amplitude = 0.0;
49 let mut phase = 0.0;
50 let mut count = 0usize;
51 for ((&reference, &candidate), &inside) in
52 reference.iter().zip(candidate.iter()).zip(support.iter())
53 {
54 if inside == 0 {
55 continue;
56 }
57 let candidate = candidate * alignment;
58 amplitude += (candidate.norm() - reference.norm()).powi(2);
59 phase += wrap_phase(candidate.arg() - reference.arg()).powi(2);
60 count += 1;
61 }
62 if count == 0 {
63 return Err(Error::InvalidParameter {
64 name: "support",
65 reason: "must select at least one pupil sample".into(),
66 });
67 }
68 Ok(PupilComparisonMetrics {
69 amplitude_rmse: (amplitude / count as f64).sqrt(),
70 phase_rmse: (phase / count as f64).sqrt(),
71 })
72}
73
74#[derive(Clone, Debug, Serialize, Deserialize)]
76pub struct IlluminationPositionMetrics {
77 pub position_rmse: f64,
79}
80
81pub fn compare_illumination_positions(
83 reference: &ImagePlaneModel,
84 candidate: &ImagePlaneModel,
85 candidate_corrections: Option<&[(f64, f64)]>,
86) -> Result<IlluminationPositionMetrics> {
87 if reference.source_count() != candidate.source_count() {
88 return Err(Error::InvalidModel(
89 "reference and candidate models have different source counts".into(),
90 ));
91 }
92 if candidate_corrections.is_some_and(|values| {
93 values.len() != candidate.source_count()
94 || values
95 .iter()
96 .any(|&(row, column)| !row.is_finite() || !column.is_finite())
97 }) {
98 return Err(Error::InvalidModel(
99 "candidate illumination corrections are invalid".into(),
100 ));
101 }
102 let mut squared = 0.0;
103 for source in 0..candidate.source_count() {
104 let candidate_crop = candidate.crop_indices.get(source)?;
105 let reference_crop = reference.crop_indices.get(source)?;
106 let candidate_offset = candidate.source_offset(source)?;
107 let reference_offset = reference.source_offset(source)?;
108 let correction = candidate_corrections.map_or((0.0, 0.0), |values| values[source]);
109 let row = candidate_crop.start_row as f64 + candidate_offset.row + correction.0
110 - reference_crop.start_row as f64
111 - reference_offset.row;
112 let column = candidate_crop.start_col as f64 + candidate_offset.column + correction.1
113 - reference_crop.start_col as f64
114 - reference_offset.column;
115 squared += row * row + column * column;
116 }
117 Ok(IlluminationPositionMetrics {
118 position_rmse: (squared / candidate.source_count() as f64).sqrt(),
119 })
120}
121
122#[derive(Clone, Debug, Serialize, Deserialize)]
124pub struct FrameGainComparisonMetrics {
125 pub relative_error: f64,
127}
128
129pub fn compare_frame_gains(
131 reference: &[f64],
132 candidate: &[f64],
133) -> Result<FrameGainComparisonMetrics> {
134 if reference.len() != candidate.len()
135 || reference.is_empty()
136 || candidate
137 .iter()
138 .any(|value| !value.is_finite() || *value <= 0.0)
139 {
140 return Err(Error::InvalidParameter {
141 name: "candidate_frame_gains",
142 reason: "must be finite, positive, and match reference length".into(),
143 });
144 }
145 let numerator: f64 = candidate
146 .iter()
147 .zip(reference)
148 .map(|(&candidate, &reference)| candidate * reference)
149 .sum();
150 let denominator: f64 = candidate.iter().map(|value| value * value).sum();
151 let alignment = numerator / denominator.max(f64::EPSILON);
152 let error: f64 = candidate
153 .iter()
154 .zip(reference)
155 .map(|(&candidate, &reference)| (alignment * candidate - reference).powi(2))
156 .sum();
157 let reference_power: f64 = reference.iter().map(|value| value * value).sum();
158 Ok(FrameGainComparisonMetrics {
159 relative_error: (error / reference_power.max(f64::EPSILON)).sqrt(),
160 })
161}
162
163fn wrap_phase(value: f64) -> f64 {
164 (value + std::f64::consts::PI).rem_euclid(std::f64::consts::TAU) - std::f64::consts::PI
165}