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#[derive(Clone, Debug)]
17pub struct DiagnosticRecorderConfig {
18 pub every: usize,
20
21 pub record_iteration_diagnostics: bool,
23 pub record_frame_summaries: bool,
25 pub record_raw_stack_stats: bool,
27 pub record_coverage: bool,
29
30 pub record_object_snapshots: bool,
32 pub record_pupil_snapshots: bool,
34 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#[derive(Clone)]
80pub struct DiagnosticRecorder {
81 config: DiagnosticRecorderConfig,
82 state: Arc<Mutex<DiagnosticRecorderState>>,
83}
84
85type Complex64Proxy = num_complex::Complex64;
86
87impl DiagnosticRecorder {
88 pub fn new(config: DiagnosticRecorderConfig) -> Self {
90 Self {
91 config,
92 state: Arc::new(Mutex::new(DiagnosticRecorderState::default())),
93 }
94 }
95
96 pub fn diagnostics(&self) -> ReconstructionDiagnostics {
98 self.lock_state().diagnostics.clone()
99 }
100
101 pub fn into_diagnostics(self) -> ReconstructionDiagnostics {
103 self.diagnostics()
104 }
105
106 pub fn object_snapshots(&self) -> Vec<(usize, Array2<f64>, Array2<f64>)> {
108 self.lock_state().object_snapshots.clone()
109 }
110
111 pub fn pupil_snapshots(&self) -> Vec<(usize, Array2<f64>, Array2<f64>)> {
113 self.lock_state().pupil_snapshots.clone()
114 }
115
116 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, ¤t) 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}