fpm_rs/diagnostics/
context.rs1use 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 pub loss_sum: f64,
37 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 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 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}