Skip to main content

fpm_rs/diagnostics/
recorder.rs

1use std::sync::{Arc, Mutex, MutexGuard};
2
3use ndarray::{Array2, ArrayView2};
4
5use crate::{
6    Result,
7    callbacks::{Callback, CallbackAction, StepContext},
8    reconstruction::ReconstructionResult,
9};
10
11use super::{
12    DiagnosticRequest, IterationDiagnostics, ReconstructionDiagnostics, compute_fourier_coverage,
13};
14
15/// Controls which diagnostics a [`DiagnosticRecorder`] captures and how often.
16#[derive(Clone, Debug)]
17pub struct DiagnosticRecorderConfig {
18    /// Iteration-recording frequency; zero is normalized to one by the recorder.
19    pub every: usize,
20
21    /// Record convergence scalars at the selected frequency.
22    pub record_iteration_diagnostics: bool,
23    /// Record predicted-versus-measured metrics for every acquisition frame.
24    pub record_frame_summaries: bool,
25    /// Record raw measured-intensity statistics once at run start.
26    pub record_raw_stack_stats: bool,
27    /// Record model Fourier coverage once at run start.
28    pub record_coverage: bool,
29
30    /// Retain owned object amplitude and phase snapshots.
31    pub record_object_snapshots: bool,
32    /// Retain owned pupil amplitude and phase snapshots.
33    pub record_pupil_snapshots: bool,
34    /// Snapshot frequency; zero is normalized to one.
35    pub snapshot_every: usize,
36}
37
38impl Default for DiagnosticRecorderConfig {
39    fn default() -> Self {
40        Self {
41            every: 1,
42            record_iteration_diagnostics: true,
43            record_frame_summaries: false,
44            record_raw_stack_stats: false,
45            record_coverage: false,
46            record_object_snapshots: false,
47            record_pupil_snapshots: false,
48            snapshot_every: 10,
49        }
50    }
51}
52
53#[derive(Default)]
54struct DiagnosticRecorderState {
55    diagnostics: ReconstructionDiagnostics,
56    previous_object: Option<Array2<Complex64Proxy>>,
57    previous_pupil: Option<Array2<Complex64Proxy>>,
58    object_snapshots: Vec<(usize, Array2<f64>, Array2<f64>)>,
59    pupil_snapshots: Vec<(usize, Array2<f64>, Array2<f64>)>,
60}
61
62/// A cloneable callback whose recorded output remains accessible after a run.
63///
64/// Pass a clone to a runner and retain the original handle:
65///
66/// ```no_run
67/// # use fpm_rs::diagnostics::{DiagnosticRecorder, DiagnosticRecorderConfig};
68/// let recorder = DiagnosticRecorder::new(DiagnosticRecorderConfig::default());
69/// let callback = Box::new(recorder.clone());
70/// # let _ = callback;
71/// // Run with `callback`, then read `recorder.diagnostics()`.
72/// ```
73///
74/// Clones share one recorder state and may be read from another thread while a
75/// run is active, but they must not be installed in multiple concurrent runs:
76/// each run resets the shared state at its start. Recorder state is
77/// best-effort diagnostics, so a poisoned mutex is recovered and reset on the
78/// next run rather than preventing later diagnostic reads or reuse.
79#[derive(Clone)]
80pub struct DiagnosticRecorder {
81    config: DiagnosticRecorderConfig,
82    state: Arc<Mutex<DiagnosticRecorderState>>,
83}
84
85type Complex64Proxy = num_complex::Complex64;
86
87impl DiagnosticRecorder {
88    /// Creates a recorder with shared state so clones observe the same captured data.
89    pub fn new(config: DiagnosticRecorderConfig) -> Self {
90        Self {
91            config,
92            state: Arc::new(Mutex::new(DiagnosticRecorderState::default())),
93        }
94    }
95
96    /// Returns a cloned, serializable snapshot of all collected diagnostic records.
97    pub fn diagnostics(&self) -> ReconstructionDiagnostics {
98        self.lock_state().diagnostics.clone()
99    }
100
101    /// Consumes this handle and returns a cloned diagnostics snapshot shared with any clones.
102    pub fn into_diagnostics(self) -> ReconstructionDiagnostics {
103        self.diagnostics()
104    }
105
106    /// Returns cloned `(iteration, amplitude, wrapped_phase)` object snapshots.
107    pub fn object_snapshots(&self) -> Vec<(usize, Array2<f64>, Array2<f64>)> {
108        self.lock_state().object_snapshots.clone()
109    }
110
111    /// Returns cloned `(iteration, amplitude, wrapped_phase)` pupil snapshots.
112    pub fn pupil_snapshots(&self) -> Vec<(usize, Array2<f64>, Array2<f64>)> {
113        self.lock_state().pupil_snapshots.clone()
114    }
115
116    /// Clears all values collected by this recorder and its clones.
117    pub fn reset(&self) {
118        *self.lock_state() = DiagnosticRecorderState::default();
119    }
120
121    fn lock_state(&self) -> MutexGuard<'_, DiagnosticRecorderState> {
122        self.state
123            .lock()
124            .unwrap_or_else(|poisoned| poisoned.into_inner())
125    }
126}
127
128impl Callback for DiagnosticRecorder {
129    fn requires_for(
130        &self,
131        hook: crate::callbacks::CallbackHook,
132        iteration: usize,
133    ) -> Vec<DiagnosticRequest> {
134        let should_record = cadence_matches(iteration, self.config.every);
135        let should_snapshot = cadence_matches(iteration, self.config.snapshot_every);
136        match hook {
137            crate::callbacks::CallbackHook::Start => {
138                let mut requests = Vec::new();
139                if self.config.record_raw_stack_stats {
140                    requests.push(DiagnosticRequest::RawFrameStats);
141                }
142                requests
143            }
144            crate::callbacks::CallbackHook::IterationEnd if should_record || should_snapshot => {
145                let mut requests = Vec::new();
146                if should_record && self.config.record_iteration_diagnostics {
147                    requests.push(DiagnosticRequest::Objective);
148                    requests.push(DiagnosticRequest::PerFrameError);
149                }
150                if should_record && self.config.record_frame_summaries {
151                    requests.push(DiagnosticRequest::FrameSummaries);
152                }
153                if should_snapshot && self.config.record_object_snapshots {
154                    requests.push(DiagnosticRequest::ObjectAmplitude);
155                    requests.push(DiagnosticRequest::ObjectPhase);
156                }
157                if should_snapshot && self.config.record_pupil_snapshots {
158                    requests.push(DiagnosticRequest::Pupil);
159                }
160                requests
161            }
162            _ => Vec::new(),
163        }
164    }
165
166    fn on_start(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
167        let mut state = self.lock_state();
168        *state = DiagnosticRecorderState::default();
169        if self.config.record_coverage {
170            state.diagnostics.coverage = Some(compute_fourier_coverage(context.model)?);
171        }
172        if let Some(raw_frame_statistics) = &context.diagnostics.raw_frame_statistics {
173            state
174                .diagnostics
175                .raw_frame_statistics
176                .extend(raw_frame_statistics.iter().cloned());
177        }
178        Ok(CallbackAction::Continue)
179    }
180
181    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
182        let should_record = cadence_matches(context.iteration, self.config.every);
183        let should_snapshot = cadence_matches(context.iteration, self.config.snapshot_every);
184        let mut state = self.lock_state();
185
186        if should_record && self.config.record_iteration_diagnostics {
187            let elapsed_seconds = context
188                .trace
189                .iterations
190                .last()
191                .map(|record| record.elapsed_seconds);
192            let object_relative_change = relative_change(
193                state.previous_object.as_ref(),
194                context.state.object_spectrum.ndarray_view(),
195            );
196            let pupil_relative_change =
197                relative_change(state.previous_pupil.as_ref(), context.state.pupil.values());
198            state
199                .diagnostics
200                .iteration_diagnostics
201                .push(IterationDiagnostics {
202                    iteration: context.iteration,
203                    total_objective: context.diagnostics.objective,
204                    data_objective: context.diagnostics.objective,
205                    regularization_objective: None,
206                    object_relative_change,
207                    pupil_relative_change,
208                    median_frame_objective: context
209                        .diagnostics
210                        .per_frame_objective
211                        .as_ref()
212                        .and_then(|values| median(values)),
213                    worst_frame_objective: context
214                        .diagnostics
215                        .per_frame_objective
216                        .as_ref()
217                        .and_then(|values| values.iter().copied().reduce(f64::max)),
218                    elapsed_seconds,
219                });
220        }
221        if should_record {
222            state.previous_object = Some(context.state.object_spectrum.clone().into_inner());
223            state.previous_pupil = Some(context.state.pupil.values.clone().into_inner());
224        }
225
226        if should_record && let Some(frame_diagnostics) = &context.diagnostics.frame_diagnostics {
227            state
228                .diagnostics
229                .frame_diagnostics
230                .extend(frame_diagnostics.iter().cloned());
231        }
232        if self.config.record_object_snapshots
233            && should_snapshot
234            && let (Some(amplitude), Some(phase)) = (
235                context.diagnostics.object_amplitude.clone(),
236                context.diagnostics.object_phase.clone(),
237            )
238        {
239            state
240                .object_snapshots
241                .push((context.iteration, amplitude, phase));
242        }
243        if self.config.record_pupil_snapshots
244            && should_snapshot
245            && let (Some(amplitude), Some(phase)) = (
246                context.diagnostics.pupil_amplitude.clone(),
247                context.diagnostics.pupil_phase.clone(),
248            )
249        {
250            state
251                .pupil_snapshots
252                .push((context.iteration, amplitude, phase));
253        }
254        Ok(CallbackAction::Continue)
255    }
256
257    fn on_finish(&mut self, _result: &ReconstructionResult) -> Result<()> {
258        Ok(())
259    }
260}
261
262fn cadence_matches(iteration: usize, every: usize) -> bool {
263    every > 0 && iteration.is_multiple_of(every)
264}
265
266fn relative_change(
267    previous: Option<&Array2<Complex64Proxy>>,
268    current: ArrayView2<'_, Complex64Proxy>,
269) -> Option<f64> {
270    let previous = previous?;
271    if previous.dim() != current.dim() {
272        return None;
273    }
274    let mut difference: f64 = 0.0;
275    let mut reference: f64 = 0.0;
276    for (&previous, &current) in previous.iter().zip(current.iter()) {
277        difference += (current - previous).norm_sqr();
278        reference += previous.norm_sqr();
279    }
280    Some(difference.sqrt() / reference.sqrt().max(f64::EPSILON))
281}
282
283fn median(values: &[f64]) -> Option<f64> {
284    if values.is_empty() {
285        return None;
286    }
287    let mut sorted = values.to_vec();
288    sorted.sort_by(|a, b| a.total_cmp(b));
289    let mid = sorted.len() / 2;
290    Some(if sorted.len().is_multiple_of(2) {
291        0.5 * (sorted[mid - 1] + sorted[mid])
292    } else {
293        sorted[mid]
294    })
295}
296
297#[cfg(test)]
298mod tests {
299    use std::panic::{AssertUnwindSafe, catch_unwind};
300
301    use super::*;
302
303    #[test]
304    fn poisoned_recorder_state_remains_resettable_and_readable() {
305        let recorder = DiagnosticRecorder::new(DiagnosticRecorderConfig::default());
306        let state = recorder.state.clone();
307        let result = catch_unwind(AssertUnwindSafe(|| {
308            let _guard = state.lock().unwrap();
309            panic!("intentional recorder-lock poison for test");
310        }));
311        assert!(result.is_err());
312
313        recorder.reset();
314        assert!(recorder.diagnostics().iteration_diagnostics.is_empty());
315        assert!(recorder.object_snapshots().is_empty());
316        assert!(recorder.pupil_snapshots().is_empty());
317    }
318}