Skip to main content

fpm_rs/reconstruction/
runner.rs

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
24/// Configurable executor for one reconstruction algorithm.
25///
26/// A runner owns the algorithm, run options, callbacks, optional checkpoint, and optional
27/// backend. [`Self::run`] consumes the runner and borrows a validated problem.
28pub 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    /// Creates a runner from an algorithm and algorithm-independent options.
38    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    /// Appends one callback in invocation order.
49    pub fn with_callback(mut self, callback: Box<dyn Callback>) -> Self {
50        self.callbacks.push(callback);
51        self
52    }
53
54    /// Appends callbacks in their vector order.
55    pub fn with_callbacks(mut self, callbacks: Vec<Box<dyn Callback>>) -> Self {
56        self.callbacks.extend(callbacks);
57        self
58    }
59
60    /// Sets a validated checkpoint from which state and elapsed trace are restored.
61    pub fn resume_from(mut self, checkpoint: ReconstructionCheckpoint) -> Self {
62        self.initial_checkpoint = Some(checkpoint);
63        self
64    }
65
66    /// Selects an execution backend used to initialize or restore state.
67    pub fn with_backend(mut self, backend: Arc<dyn Backend>) -> Self {
68        self.backend = Some(backend);
69        self
70    }
71
72    /// Validates inputs, executes scheduled batches and callbacks, and returns owned results.
73    ///
74    /// The run stops at `max_iterations` or when a callback returns a stop action. A supplied
75    /// checkpoint must match the problem and the algorithm's required auxiliary state.
76    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                        // A batch update completes every frame in the batch at
186                        // once. Emit one frame hook per completed frame, with
187                        // the frame's own natural loss and the shared post-batch
188                        // state. Expensive diagnostics are computed only once.
189                        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}