Skip to main content

fpm_rs/diagnostics/
context.rs

1use num_complex::Complex64;
2use serde::{Deserialize, Serialize};
3
4use crate::Array2;
5
6use super::{FrameDiagnosticRecord, RawFrameStatisticsRecord};
7
8#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
9pub enum DiagnosticRequest {
10    Loss,
11    PerFrameError,
12    RawFrameStats,
13    FrameSummaries,
14    ObjectAmplitude,
15    ObjectPhase,
16    Pupil,
17    ResidualImages,
18}
19
20#[derive(Clone, Debug, Default)]
21pub struct Diagnostics {
22    pub loss: Option<f64>,
23    pub per_frame_error: Option<Vec<f64>>,
24    pub raw_frame_stats: Option<Vec<RawFrameStatisticsRecord>>,
25    pub frame_diagnostics: Option<Vec<FrameDiagnosticRecord>>,
26    pub object_amplitude: Option<Array2<f64>>,
27    pub object_phase: Option<Array2<f64>>,
28    pub pupil_amplitude: Option<Array2<f64>>,
29    pub pupil_phase: Option<Array2<f64>>,
30    pub residual_images: Option<Vec<Array2<f64>>>,
31}
32
33#[derive(Clone, Debug, Default, Serialize, Deserialize)]
34pub struct StepDiagnostics {
35    /// Sum of per-frame losses multiplied by frame weights.
36    pub loss_sum: f64,
37    /// Number of frames visited, including zero-weight frames.
38    pub frame_count: usize,
39    pub weight_sum: f64,
40    pub per_frame_loss: std::collections::BTreeMap<usize, f64>,
41    #[serde(default)]
42    admm_primal_residual_sum_squares: f64,
43    #[serde(default)]
44    admm_dual_residual_sum_squares: f64,
45    #[serde(default)]
46    admm_residual_count: usize,
47}
48
49impl StepDiagnostics {
50    pub fn push_frame(&mut self, frame: usize, loss: f64, weight: f64) {
51        self.loss_sum += weight * loss;
52        self.frame_count += 1;
53        self.weight_sum += weight;
54        self.per_frame_loss.insert(frame, loss);
55    }
56
57    pub fn mean_loss(&self) -> Option<f64> {
58        (self.weight_sum > 0.0).then(|| self.loss_sum / self.weight_sum)
59    }
60
61    /// Root-mean-square ADMM consensus residual, `A x - z`, over active modes.
62    pub fn admm_primal_residual_rms(&self) -> Option<f64> {
63        (self.admm_residual_count > 0).then(|| {
64            (self.admm_primal_residual_sum_squares / self.admm_residual_count as f64).sqrt()
65        })
66    }
67
68    /// Root-mean-square ADMM dual residual, `rho (z_k - z_{k-1})`, over active modes.
69    pub fn admm_dual_residual_rms(&self) -> Option<f64> {
70        (self.admm_residual_count > 0)
71            .then(|| (self.admm_dual_residual_sum_squares / self.admm_residual_count as f64).sqrt())
72    }
73
74    pub(crate) fn push_admm_dual_change(&mut self, change: Complex64, penalty: f64) {
75        self.admm_dual_residual_sum_squares += penalty * penalty * change.norm_sqr();
76    }
77
78    pub(crate) fn push_admm_primal_residual(&mut self, residual: Complex64) {
79        self.admm_primal_residual_sum_squares += residual.norm_sqr();
80        self.admm_residual_count += 1;
81    }
82
83    pub fn merge(&mut self, other: Self) {
84        self.loss_sum += other.loss_sum;
85        self.frame_count += other.frame_count;
86        self.weight_sum += other.weight_sum;
87        self.per_frame_loss.extend(other.per_frame_loss);
88        self.admm_primal_residual_sum_squares += other.admm_primal_residual_sum_squares;
89        self.admm_dual_residual_sum_squares += other.admm_dual_residual_sum_squares;
90        self.admm_residual_count += other.admm_residual_count;
91    }
92}