Skip to main content

fpm_rs/metrics/
model.rs

1//! Atomic comparisons for calibrated optical-model quantities.
2
3use 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}