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