1use ndarray::Array2;
2use std::{collections::BTreeSet, sync::Arc, time::Instant};
3
4use crate::{
5 Result,
6 algorithms::{AlgorithmIterationMetrics, ReconstructionAlgorithm, StepOutput, StepSummary},
7 backend::Backend,
8 callbacks::{Callback, CallbackAction, CallbackHook, StepContext},
9 complex,
10 diagnostics::{
11 DiagnosticRequest, Diagnostics, FrameDiagnosticRecord, RawFrameStatisticsRecord,
12 },
13 error::Error,
14 measurements::MeasurementRead,
15 metrics::intensity::{compare_intensity_u8_masked, stats},
16 model::ForwardModel,
17};
18
19use super::{
20 Batch, ReconstructionCheckpoint, ReconstructionProblem, ReconstructionResult,
21 ReconstructionState, ReconstructionTrace, RunOptions, RuntimeInfo, state_object,
22};
23
24pub struct Runner<A> {
29 algorithm: A,
30 options: RunOptions,
31 callbacks: Vec<Box<dyn Callback>>,
32 initial_checkpoint: Option<ReconstructionCheckpoint>,
33 backend: Option<Arc<dyn Backend>>,
34}
35
36impl<A: ReconstructionAlgorithm> Runner<A> {
37 pub fn new(algorithm: A, options: RunOptions) -> Self {
39 Self {
40 algorithm,
41 options,
42 callbacks: Vec::new(),
43 initial_checkpoint: None,
44 backend: None,
45 }
46 }
47
48 pub fn with_callback(mut self, callback: Box<dyn Callback>) -> Self {
50 self.callbacks.push(callback);
51 self
52 }
53
54 pub fn with_callbacks(mut self, callbacks: Vec<Box<dyn Callback>>) -> Self {
56 self.callbacks.extend(callbacks);
57 self
58 }
59
60 pub fn resume_from(mut self, checkpoint: ReconstructionCheckpoint) -> Self {
62 self.initial_checkpoint = Some(checkpoint);
63 self
64 }
65
66 pub fn with_backend(mut self, backend: Arc<dyn Backend>) -> Self {
68 self.backend = Some(backend);
69 self
70 }
71
72 pub fn run<M: MeasurementRead>(
77 mut self,
78 problem: &ReconstructionProblem<M>,
79 ) -> Result<ReconstructionResult> {
80 self.algorithm.validate()?;
81 problem.validate()?;
82 self.algorithm.validate_problem(problem)?;
83 if self.options.batch_size == 0 {
84 return Err(Error::InvalidParameter {
85 name: "batch_size",
86 reason: "must be greater than zero".into(),
87 });
88 }
89 let started = Instant::now();
90 let algorithm_name = std::any::type_name::<A>()
91 .rsplit("::")
92 .next()
93 .unwrap_or("reconstruction algorithm")
94 .to_owned();
95 let (mut state, mut trace, starting_iteration) =
96 if let Some(checkpoint) = &self.initial_checkpoint {
97 (
98 if let Some(backend) = &self.backend {
99 ReconstructionState::from_checkpoint_with_backend(
100 problem,
101 checkpoint,
102 backend.clone(),
103 )?
104 } else {
105 ReconstructionState::from_checkpoint(problem, checkpoint)?
106 },
107 checkpoint.trace.clone(),
108 checkpoint.completed_iterations,
109 )
110 } else {
111 (
112 if let Some(backend) = &self.backend {
113 self.algorithm
114 .initialize_with_backend(problem, backend.clone())?
115 } else {
116 self.algorithm.initialize(problem)?
117 },
118 ReconstructionTrace::default(),
119 0,
120 )
121 };
122 self.algorithm.canonicalize_state(problem, &mut state)?;
123 state.object_real_space_cache = None;
124 let previous_elapsed = trace
125 .iterations
126 .last()
127 .map_or(0.0, |record| record.elapsed_seconds);
128 let start_requests =
129 callback_requests(&self.callbacks, CallbackHook::Start, starting_iteration);
130 let start_diagnostics = build_diagnostics(
131 problem,
132 &mut state,
133 &start_requests,
134 None,
135 None,
136 starting_iteration,
137 )?;
138 let start_context = StepContext {
139 iteration: starting_iteration,
140 frame_index: None,
141 batch_index: None,
142 state: &state,
143 diagnostics: &start_diagnostics,
144 trace: &trace,
145 current_algorithm_metrics: &[],
146 model: state.calibrated_model.as_ref().unwrap_or(&problem.model),
147 problem_name: problem.name.as_deref(),
148 };
149 let mut stopped_early = false;
150 for callback in &mut self.callbacks {
151 if callback.on_start(&start_context)? == CallbackAction::Stop {
152 stopped_early = true;
153 }
154 }
155 if !stopped_early {
156 for zero_based_iteration in starting_iteration..self.options.max_iterations {
157 let current_iteration = zero_based_iteration + 1;
158 let frame_requests =
159 callback_requests(&self.callbacks, CallbackHook::FrameEnd, current_iteration);
160 let order = self
161 .options
162 .schedule
163 .order_for_problem(problem, zero_based_iteration)?;
164 let mut iteration_step = StepOutput::<A::IterationMetrics>::default();
165 for (batch_index, indices) in order.chunks(self.options.batch_size).enumerate() {
166 let batch = Batch::new(indices.to_vec(), batch_index);
167 let batch_step =
168 self.algorithm
169 .step(problem, &mut state, &batch, zero_based_iteration)?;
170 state.object_real_space_cache = None;
171 if self.options.enable_frame_callbacks {
172 let batch_objective = batch_step.summary.mean_objective();
173 let mut batch_metric_records = Vec::new();
174 batch_step
175 .metrics
176 .append_records(current_iteration, &mut batch_metric_records);
177 let mut diagnostics = build_diagnostics(
178 problem,
179 &mut state,
180 &frame_requests,
181 batch_objective,
182 Some(&batch_step.summary),
183 current_iteration,
184 )?;
185 for &frame in &batch.indices {
190 if frame_requests.contains(&DiagnosticRequest::Objective) {
191 diagnostics.objective =
192 batch_step.summary.per_frame_objective.get(&frame).copied();
193 }
194 let context = StepContext {
195 iteration: current_iteration,
196 frame_index: Some(frame),
197 batch_index: Some(batch_index),
198 state: &state,
199 diagnostics: &diagnostics,
200 trace: &trace,
201 current_algorithm_metrics: &batch_metric_records,
202 model: state.calibrated_model.as_ref().unwrap_or(&problem.model),
203 problem_name: problem.name.as_deref(),
204 };
205 for callback in &mut self.callbacks {
206 if callback.on_frame_end(&context)? == CallbackAction::Stop {
207 stopped_early = true;
208 }
209 }
210 if stopped_early {
211 break;
212 }
213 }
214 }
215 iteration_step.merge(batch_step);
216 if stopped_early {
217 break;
218 }
219 }
220 self.algorithm.canonicalize_state(problem, &mut state)?;
221 state.object_real_space_cache = None;
222 let objective = iteration_step.summary.mean_objective().unwrap_or(f64::NAN);
223 trace.iterations.push(super::IterationRecord {
224 iteration: current_iteration,
225 objective,
226 elapsed_seconds: previous_elapsed + started.elapsed().as_secs_f64(),
227 });
228 let metric_start = trace.algorithm_metrics.len();
229 iteration_step
230 .metrics
231 .append_records(current_iteration, &mut trace.algorithm_metrics);
232 let iteration_requests = callback_requests(
233 &self.callbacks,
234 CallbackHook::IterationEnd,
235 current_iteration,
236 );
237 let diagnostics = build_diagnostics(
238 problem,
239 &mut state,
240 &iteration_requests,
241 Some(objective),
242 Some(&iteration_step.summary),
243 current_iteration,
244 )?;
245 let context = StepContext {
246 iteration: current_iteration,
247 frame_index: None,
248 batch_index: None,
249 state: &state,
250 diagnostics: &diagnostics,
251 trace: &trace,
252 current_algorithm_metrics: &trace.algorithm_metrics[metric_start..],
253 model: state.calibrated_model.as_ref().unwrap_or(&problem.model),
254 problem_name: problem.name.as_deref(),
255 };
256 for callback in &mut self.callbacks {
257 if callback.on_iteration_end(&context)? == CallbackAction::Stop {
258 stopped_early = true;
259 }
260 }
261 if stopped_early {
262 break;
263 }
264 }
265 }
266
267 let runtime = RuntimeInfo {
268 elapsed_seconds: previous_elapsed + started.elapsed().as_secs_f64(),
269 completed_iterations: trace
270 .iterations
271 .last()
272 .map_or(starting_iteration, |record| record.iteration),
273 stopped_early,
274 algorithm: algorithm_name,
275 };
276 let mut result = ReconstructionResult::from_state(&mut state, trace, runtime)?;
277 if let Some(name) = &problem.name {
278 result.metadata.insert("problem_name".into(), name.clone());
279 }
280 for callback in &mut self.callbacks {
281 callback.on_finish(&result)?;
282 }
283 Ok(result)
284 }
285}
286
287fn callback_requests(
288 callbacks: &[Box<dyn Callback>],
289 hook: CallbackHook,
290 iteration: usize,
291) -> BTreeSet<DiagnosticRequest> {
292 callbacks
293 .iter()
294 .flat_map(|callback| callback.requires_for(hook, iteration))
295 .collect()
296}
297
298fn build_diagnostics<M: MeasurementRead>(
299 problem: &ReconstructionProblem<M>,
300 state: &mut ReconstructionState,
301 requests: &BTreeSet<DiagnosticRequest>,
302 natural_objective: Option<f64>,
303 step: Option<&StepSummary>,
304 iteration: usize,
305) -> Result<Diagnostics> {
306 let mut diagnostics = Diagnostics::default();
307 if requests.contains(&DiagnosticRequest::Objective) {
308 diagnostics.objective = natural_objective;
309 }
310 if requests.contains(&DiagnosticRequest::RawFrameStats) {
311 let mut values = Vec::with_capacity(problem.model.frame_count());
312 for frame in 0..problem.model.frame_count() {
313 let measured = problem.measurements.frame(frame)?;
314 values.push(RawFrameStatisticsRecord {
315 frame_index: frame,
316 metrics: stats(&measured, None)?,
317 });
318 }
319 diagnostics.raw_frame_statistics = Some(values);
320 }
321 if requests.contains(&DiagnosticRequest::PerFrameError)
322 && let Some(step) = step
323 {
324 let mut values = vec![f64::NAN; problem.model.frame_count()];
325 for (&frame, &objective) in &step.per_frame_objective {
326 values[frame] = objective;
327 }
328 diagnostics.per_frame_objective = Some(values);
329 }
330 if requests.contains(&DiagnosticRequest::FrameSummaries)
331 || requests.contains(&DiagnosticRequest::ResidualImages)
332 || (requests.contains(&DiagnosticRequest::PerFrameError) && step.is_none())
333 {
334 let diagnostic_model = model_with_state_calibration(problem, state)?;
335 let forward = ForwardModel::with_backend(&diagnostic_model, state.backend.clone())?;
336 let mut workspace = forward.workspace()?;
337 let mut predicted =
338 vec![0.0; crate::array_layout::checked_len_2d(problem.model.image_shape)?];
339 let mut frame_diagnostics = Vec::with_capacity(problem.model.frame_count());
340 let mut residual_images = if requests.contains(&DiagnosticRequest::ResidualImages) {
341 Some(Vec::with_capacity(problem.model.frame_count()))
342 } else {
343 None
344 };
345 let mut per_frame_objective =
346 if requests.contains(&DiagnosticRequest::PerFrameError) && step.is_none() {
347 Some(Vec::with_capacity(problem.model.frame_count()))
348 } else {
349 None
350 };
351 for frame in 0..problem.model.frame_count() {
352 forward.forward_intensity_standard_into(
353 state.object_spectrum_standard_view(),
354 &state.pupil,
355 frame,
356 &mut workspace,
357 &mut predicted,
358 )?;
359 let measured = problem.measurements.frame(frame)?;
360 let mask = problem.measurements.frame_mask(frame)?;
361 let metadata = problem
362 .measurements
363 .frame_metadata()
364 .get(frame)
365 .cloned()
366 .unwrap_or_else(|| crate::measurements::FrameMetadata::new(frame));
367 if let Some(values) = per_frame_objective.as_mut() {
368 let mut loss_sum = 0.0;
369 let mut valid_pixels = 0;
370 for pixel in 0..predicted.len() {
371 if mask.is_none_or(|values| values[pixel] != 0) {
372 let residual =
373 predicted[pixel].max(0.0).sqrt() - measured[pixel].max(0.0).sqrt();
374 loss_sum += residual * residual;
375 valid_pixels += 1;
376 }
377 }
378 values.push(if valid_pixels == 0 {
379 0.0
380 } else {
381 loss_sum / valid_pixels as f64
382 });
383 }
384 if let Some(images) = residual_images.as_mut() {
385 let mut residual: Vec<_> = predicted
386 .iter()
387 .zip(measured.iter())
388 .map(|(&predicted, &measured)| predicted - measured)
389 .collect();
390 if let Some(mask) = mask {
391 for (value, &valid) in residual.iter_mut().zip(mask) {
392 if valid == 0 {
393 *value = 0.0;
394 }
395 }
396 }
397 images.push(Array2::from_shape_vec(problem.model.image_shape, residual)?);
398 }
399 if requests.contains(&DiagnosticRequest::FrameSummaries) {
400 let metrics = compare_intensity_u8_masked(&measured, &predicted, mask, None)?;
401 frame_diagnostics.push(FrameDiagnosticRecord {
402 iteration: Some(iteration),
403 frame_index: frame,
404 illumination_index: metadata.illumination_index.unwrap_or(frame),
405 metrics,
406 });
407 }
408 }
409 if diagnostics.per_frame_objective.is_none() {
410 diagnostics.per_frame_objective = per_frame_objective;
411 }
412 diagnostics.frame_diagnostics = Some(frame_diagnostics);
413 if let Some(images) = residual_images {
414 diagnostics.residual_images = Some(images);
415 }
416 }
417 if requests.contains(&DiagnosticRequest::ObjectAmplitude)
418 || requests.contains(&DiagnosticRequest::ObjectPhase)
419 {
420 let object = state_object(state)?;
421 if requests.contains(&DiagnosticRequest::ObjectAmplitude) {
422 diagnostics.object_amplitude = Some(complex::amplitude(object.view()));
423 }
424 if requests.contains(&DiagnosticRequest::ObjectPhase) {
425 diagnostics.object_phase = Some(complex::phase(object.view()));
426 }
427 }
428 if requests.contains(&DiagnosticRequest::Pupil) {
429 diagnostics.pupil_amplitude = Some(complex::amplitude(state.pupil.values()));
430 diagnostics.pupil_phase = Some(complex::phase(state.pupil.values()));
431 }
432 Ok(diagnostics)
433}
434
435fn model_with_state_calibration<M: MeasurementRead>(
436 problem: &ReconstructionProblem<M>,
437 state: &ReconstructionState,
438) -> Result<crate::model::ImagePlaneModel> {
439 let mut model = state
440 .calibrated_model
441 .clone()
442 .unwrap_or_else(|| problem.model.clone());
443 model.frame_gains = state.frame_gains.clone();
444 model.background = state.background.clone();
445 if state.illumination_corrections.is_some() {
446 model.subpixel_offsets = Some(
447 (0..model.source_count())
448 .map(|source| state.effective_source_offset(&model, source))
449 .collect::<Result<Vec<_>>>()?,
450 );
451 }
452 model.validate()?;
453 Ok(model)
454}