Skip to main content

fpm_rs/reconstruction/
checkpoint.rs

1use std::{
2    fs::File,
3    io::{BufReader, BufWriter},
4    path::Path,
5};
6
7use ndarray::{Array2, ArrayView2};
8use num_complex::Complex64;
9use serde::{Deserialize, Serialize};
10
11use crate::{
12    Result,
13    array_layout::checked_len_2d,
14    array_serde::Array2Data,
15    error::Error,
16    illumination_calibration::IlluminationCalibrationState,
17    measurements::MeasurementRead,
18    model::{ImagePlaneModel, Pupil},
19};
20
21use super::{
22    AlgorithmAuxiliaryState, ReconstructionProblem, ReconstructionState, ReconstructionTrace,
23};
24
25/// Current JSON checkpoint serialization format version.
26pub const CHECKPOINT_FORMAT_VERSION: u32 = 2;
27
28/// Serializable algorithm state used to resume a reconstruction exactly.
29///
30/// Built-in pupil-recovering algorithms capture object and pupil arrays after
31/// their iteration-boundary gauge projection. Older valid checkpoints are
32/// projected immediately after restoration. Problem-aware validation requires
33/// the stored pupil support to equal the compiled model support, not only to
34/// have the same shape. Format version 2 uses `algorithm_auxiliary` as an
35/// extension point for solver state. Readers predating a particular auxiliary
36/// enum variant cannot deserialize checkpoints containing that variant.
37#[derive(Clone, Debug)]
38pub struct ReconstructionCheckpoint {
39    pub(crate) format_version: u32,
40    pub(crate) completed_iterations: usize,
41    pub(crate) object_spectrum: Array2<Complex64>,
42    pub(crate) pupil: Pupil,
43    /// Per-source `(row, column)` corrections in Fourier-grid pixels.
44    pub(crate) illumination_corrections: Option<Vec<(f64, f64)>>,
45    pub(crate) frame_gains: Option<Vec<f64>>,
46    pub(crate) background: Option<Vec<f64>>,
47    pub(crate) physical_illumination_calibration: Option<IlluminationCalibrationState>,
48    pub(crate) calibrated_model: Option<ImagePlaneModel>,
49    pub(crate) algorithm_auxiliary: Option<AlgorithmAuxiliaryState>,
50    pub(crate) trace: ReconstructionTrace,
51}
52
53impl ReconstructionCheckpoint {
54    /// Clones resumable state and trace after `completed_iterations` complete passes.
55    pub fn capture(
56        completed_iterations: usize,
57        state: &ReconstructionState,
58        trace: &ReconstructionTrace,
59    ) -> Self {
60        Self {
61            format_version: CHECKPOINT_FORMAT_VERSION,
62            completed_iterations,
63            object_spectrum: state.object_spectrum.clone().into_inner(),
64            pupil: state.pupil.clone(),
65            illumination_corrections: state.illumination_corrections.clone(),
66            frame_gains: state.frame_gains.clone(),
67            background: state.background.clone(),
68            physical_illumination_calibration: state.physical_illumination_calibration.clone(),
69            calibrated_model: state.calibrated_model.clone(),
70            algorithm_auxiliary: state.algorithm_auxiliary.clone(),
71            trace: trace.clone(),
72        }
73    }
74
75    /// Returns the serialized checkpoint format version.
76    pub const fn format_version(&self) -> u32 {
77        self.format_version
78    }
79
80    /// Returns the number of complete iterations represented by this state.
81    pub const fn completed_iterations(&self) -> usize {
82        self.completed_iterations
83    }
84
85    /// Borrows the centered high-resolution complex object spectrum.
86    pub fn object_spectrum(&self) -> ArrayView2<'_, Complex64> {
87        self.object_spectrum.view()
88    }
89
90    /// Borrows the low-resolution recovered pupil state.
91    pub fn pupil(&self) -> &Pupil {
92        &self.pupil
93    }
94
95    /// Borrows optional per-source `(row, column)` corrections in Fourier-grid pixels.
96    pub fn illumination_corrections(&self) -> Option<&[(f64, f64)]> {
97        self.illumination_corrections.as_deref()
98    }
99
100    /// Borrows optional positive calibration gains in acquisition-frame order.
101    pub fn frame_gains(&self) -> Option<&[f64]> {
102        self.frame_gains.as_deref()
103    }
104
105    /// Borrows optional additive intensity backgrounds in acquisition-frame order.
106    pub fn background(&self) -> Option<&[f64]> {
107        self.background.as_deref()
108    }
109
110    /// Borrows physical planar-array calibration state, when joint calibration is active.
111    pub fn physical_illumination_calibration(&self) -> Option<&IlluminationCalibrationState> {
112        self.physical_illumination_calibration.as_ref()
113    }
114
115    /// Borrows the illumination-refreshed model used at checkpoint capture.
116    pub fn calibrated_model(&self) -> Option<&ImagePlaneModel> {
117        self.calibrated_model.as_ref()
118    }
119
120    /// Borrows optional solver-specific resumable state.
121    pub fn algorithm_auxiliary(&self) -> Option<&AlgorithmAuxiliaryState> {
122        self.algorithm_auxiliary.as_ref()
123    }
124
125    /// Borrows the iteration trace accumulated before capture.
126    pub const fn trace(&self) -> &ReconstructionTrace {
127        &self.trace
128    }
129
130    /// Validates and serializes this checkpoint as JSON.
131    pub fn save(&self, path: impl AsRef<Path>) -> Result<()> {
132        self.validate()?;
133        let writer = BufWriter::new(File::create(path)?);
134        serde_json::to_writer(writer, self)?;
135        Ok(())
136    }
137
138    /// Deserializes and validates a JSON checkpoint independently of a problem.
139    pub fn load(path: impl AsRef<Path>) -> Result<Self> {
140        let reader = BufReader::new(File::open(path)?);
141        let checkpoint: Self = serde_json::from_reader(reader)?;
142        checkpoint.validate()?;
143        Ok(checkpoint)
144    }
145
146    /// Loads a checkpoint and verifies all dimensions, the exact pupil support,
147    /// and calibration counts against the problem that will resume it.
148    pub fn load_for_problem<M: MeasurementRead>(
149        path: impl AsRef<Path>,
150        problem: &ReconstructionProblem<M>,
151    ) -> Result<Self> {
152        let checkpoint = Self::load(path)?;
153        checkpoint.validate_for_problem(problem)?;
154        Ok(checkpoint)
155    }
156
157    /// Validates integrity that does not depend on a reconstruction problem.
158    pub fn validate(&self) -> Result<()> {
159        if self.format_version != CHECKPOINT_FORMAT_VERSION {
160            return Err(Error::InvalidParameter {
161                name: "checkpoint format_version",
162                reason: format!(
163                    "expected {CHECKPOINT_FORMAT_VERSION}, got {}",
164                    self.format_version
165                ),
166            });
167        }
168        if self.pupil.support.len() != self.pupil.values.len() {
169            return Err(Error::InvalidShape(
170                "checkpoint pupil support and values have different lengths".into(),
171            ));
172        }
173        if self
174            .object_spectrum
175            .iter()
176            .chain(self.pupil.values.as_slice())
177            .any(|value| !value.re.is_finite() || !value.im.is_finite())
178        {
179            return Err(Error::InvalidModel(
180                "checkpoint contains non-finite complex values".into(),
181            ));
182        }
183        if self
184            .illumination_corrections
185            .as_ref()
186            .is_some_and(|values| {
187                values
188                    .iter()
189                    .any(|&(row, column)| !row.is_finite() || !column.is_finite())
190            })
191        {
192            return Err(Error::InvalidModel(
193                "checkpoint illumination corrections contain non-finite values".into(),
194            ));
195        }
196        if self.frame_gains.as_ref().is_some_and(|values| {
197            values
198                .iter()
199                .any(|value| !value.is_finite() || *value <= 0.0)
200        }) {
201            return Err(Error::InvalidModel(
202                "checkpoint frame gains must be finite and positive".into(),
203            ));
204        }
205        if self
206            .background
207            .as_ref()
208            .is_some_and(|values| values.iter().any(|value| !value.is_finite()))
209        {
210            return Err(Error::InvalidModel(
211                "checkpoint background contains non-finite values".into(),
212            ));
213        }
214        if self.physical_illumination_calibration.is_some() != self.calibrated_model.is_some() {
215            return Err(Error::InvalidModel(
216                "checkpoint physical calibration and calibrated model must be present together"
217                    .into(),
218            ));
219        }
220        if let Some(model) = &self.calibrated_model {
221            model.validate()?;
222        }
223        if let Some(calibration) = &self.physical_illumination_calibration {
224            calibration.validate()?;
225        }
226        if self
227            .algorithm_auxiliary
228            .as_ref()
229            .is_some_and(|auxiliary| match auxiliary {
230                AlgorithmAuxiliaryState::Admm(admm) => {
231                    admm.auxiliary_fields.len() != admm.dual_fields.len()
232                        || admm
233                            .auxiliary_fields
234                            .iter()
235                            .chain(&admm.dual_fields)
236                            .any(|value| !value.re.is_finite() || !value.im.is_finite())
237                }
238                AlgorithmAuxiliaryState::Mpie(mpie) => {
239                    mpie.velocity.len() != mpie.anchor.len()
240                        || mpie
241                            .velocity
242                            .iter()
243                            .chain(&mpie.anchor)
244                            .any(|value| !value.re.is_finite() || !value.im.is_finite())
245                        || !mpie.object_step.is_finite()
246                        || mpie.object_step <= 0.0
247                        || !mpie.stability.is_finite()
248                        || !(0.0..=1.0).contains(&mpie.stability)
249                        || !mpie.epsilon.is_finite()
250                        || mpie.epsilon <= 0.0
251                        || mpie.momentum_interval == 0
252                        || mpie.effective_frames_since_momentum >= mpie.momentum_interval
253                        || !mpie.momentum_friction.is_finite()
254                        || !(0.0..1.0).contains(&mpie.momentum_friction)
255                        || !mpie.momentum_feedback.is_finite()
256                        || !(0.0..=1.0).contains(&mpie.momentum_feedback)
257                }
258                AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(adaptive) => {
259                    !adaptive.current_object_step.is_finite()
260                        || adaptive.current_object_step <= 0.0
261                        || !adaptive.initial_object_step.is_finite()
262                        || adaptive.initial_object_step <= 0.0
263                        || !adaptive.minimum_object_step.is_finite()
264                        || adaptive.minimum_object_step <= 0.0
265                        || adaptive.minimum_object_step > adaptive.initial_object_step
266                        || adaptive.current_object_step < adaptive.minimum_object_step
267                        || adaptive.current_object_step > adaptive.initial_object_step
268                        || !adaptive.progress_threshold.is_finite()
269                        || !(0.0..1.0).contains(&adaptive.progress_threshold)
270                        || !adaptive.reduction_factor.is_finite()
271                        || adaptive.reduction_factor <= 0.0
272                        || adaptive.reduction_factor >= 1.0
273                        || !adaptive.epsilon.is_finite()
274                        || adaptive.epsilon <= 0.0
275                        || !adaptive.objective_sum.is_finite()
276                        || adaptive.objective_sum < 0.0
277                        || !adaptive.weight_sum.is_finite()
278                        || adaptive.weight_sum < 0.0
279                        || adaptive
280                            .previous_objective
281                            .is_some_and(|value| !value.is_finite() || value < 0.0)
282                        || (adaptive.frames_accumulated == 0
283                            && (adaptive.objective_sum != 0.0 || adaptive.weight_sum != 0.0))
284                        || (adaptive.frames_accumulated > 0 && adaptive.weight_sum <= 0.0)
285                        || (self.completed_iterations == 0
286                            && (adaptive.active_iteration != 0 || adaptive.frames_accumulated != 0))
287                        || (self.completed_iterations > 0
288                            && (adaptive.active_iteration.checked_add(1)
289                                != Some(self.completed_iterations)
290                                || adaptive.frames_accumulated == 0))
291                        || (self.completed_iterations >= 2 && adaptive.previous_objective.is_none())
292                }
293            })
294        {
295            return Err(Error::InvalidModel(
296                "checkpoint algorithm auxiliary state is inconsistent or non-finite".into(),
297            ));
298        }
299        let records = &self.trace.iterations;
300        if records.len() != self.completed_iterations
301            || records.iter().enumerate().any(|(index, record)| {
302                record.iteration != index + 1
303                    || !record.objective.is_finite()
304                    || !record.elapsed_seconds.is_finite()
305                    || record.elapsed_seconds < 0.0
306            })
307            || records
308                .windows(2)
309                .any(|pair| pair[1].elapsed_seconds < pair[0].elapsed_seconds)
310        {
311            return Err(Error::InvalidModel(
312                "checkpoint trace is incomplete, non-finite, or non-monotonic".into(),
313            ));
314        }
315        if self.trace.algorithm_metrics.iter().any(|record| {
316            record.iteration == 0
317                || record.iteration > self.completed_iterations
318                || record.namespace.is_empty()
319                || record.metric.is_empty()
320                || !record.value.is_finite()
321        }) {
322            return Err(Error::InvalidModel(
323                "checkpoint algorithm metrics are invalid".into(),
324            ));
325        }
326        Ok(())
327    }
328
329    /// Validates checkpoint dimensions and calibration variables for `problem`.
330    pub fn validate_for_problem<M: MeasurementRead>(
331        &self,
332        problem: &ReconstructionProblem<M>,
333    ) -> Result<()> {
334        problem.validate()?;
335        self.validate()?;
336        if self.object_spectrum.dim() != problem.model.reconstruction_shape {
337            return Err(Error::InvalidShape(format!(
338                "checkpoint spectrum shape {:?} differs from reconstruction shape {:?}",
339                self.object_spectrum.dim(),
340                problem.model.reconstruction_shape
341            )));
342        }
343        if self.pupil.shape() != problem.model.image_shape {
344            return Err(Error::InvalidShape(
345                "checkpoint pupil does not match the model image shape".into(),
346            ));
347        }
348        if self.pupil.support.as_slice() != problem.model.pupil().support.as_slice() {
349            return Err(Error::InvalidModel(
350                "checkpoint pupil support differs from the reconstruction problem".into(),
351            ));
352        }
353        if self
354            .illumination_corrections
355            .as_ref()
356            .is_some_and(|values| values.len() != problem.model.source_count())
357        {
358            return Err(Error::InvalidModel(
359                "checkpoint illumination correction count does not match the model".into(),
360            ));
361        }
362        if self
363            .frame_gains
364            .as_ref()
365            .is_some_and(|values| values.len() != problem.model.frame_count())
366        {
367            return Err(Error::InvalidModel(
368                "checkpoint frame gain count does not match the model".into(),
369            ));
370        }
371        if let Some(model) = &self.calibrated_model
372            && (model.image_shape() != problem.model.image_shape()
373                || model.reconstruction_shape() != problem.model.reconstruction_shape()
374                || model.source_count() != problem.model.source_count()
375                || model.frame_count() != problem.model.frame_count())
376        {
377            return Err(Error::InvalidModel(
378                "checkpoint calibrated model topology differs from the reconstruction problem"
379                    .into(),
380            ));
381        }
382        let image_len = problem.measurements.frame_len();
383        let stack_len = image_len
384            .checked_mul(problem.model.frame_count())
385            .ok_or_else(|| Error::InvalidShape("checkpoint stack length overflows".into()))?;
386        if self
387            .background
388            .as_ref()
389            .is_some_and(|values| values.len() != image_len && values.len() != stack_len)
390        {
391            return Err(Error::InvalidModel(
392                "checkpoint background dimensions do not match the model".into(),
393            ));
394        }
395        let mode_count = problem.model.multiplexing_matrix.as_ref().map_or_else(
396            || Ok(problem.model.frame_count()),
397            |matrix| {
398                matrix.iter().try_fold(0_usize, |count, row| {
399                    count.checked_add(row.len()).ok_or_else(|| {
400                        Error::InvalidShape("checkpoint source mode count overflows".into())
401                    })
402                })
403            },
404        )?;
405        let auxiliary_len = image_len
406            .checked_mul(mode_count)
407            .ok_or_else(|| Error::InvalidShape("checkpoint auxiliary length overflows".into()))?;
408        let object_len = checked_len_2d(problem.model.reconstruction_shape())?;
409        if self
410            .algorithm_auxiliary
411            .as_ref()
412            .is_some_and(|auxiliary| match auxiliary {
413                AlgorithmAuxiliaryState::Admm(admm) => admm.auxiliary_fields.len() != auxiliary_len,
414                AlgorithmAuxiliaryState::Mpie(mpie) => {
415                    mpie.velocity.len() != object_len || mpie.anchor.len() != object_len
416                }
417                AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(adaptive) => {
418                    self.completed_iterations > 0
419                        && adaptive.frames_accumulated != problem.model.frame_count()
420                }
421            })
422        {
423            return Err(Error::InvalidModel(
424                "checkpoint algorithm auxiliary dimensions do not match the model".into(),
425            ));
426        }
427        Ok(())
428    }
429}
430
431impl Serialize for ReconstructionCheckpoint {
432    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
433    where
434        S: serde::Serializer,
435    {
436        #[derive(Serialize)]
437        struct Representation<'a> {
438            format_version: u32,
439            completed_iterations: usize,
440            object_spectrum: Array2Data<Complex64>,
441            pupil: &'a Pupil,
442            illumination_corrections: &'a Option<Vec<(f64, f64)>>,
443            frame_gains: &'a Option<Vec<f64>>,
444            background: &'a Option<Vec<f64>>,
445            physical_illumination_calibration: &'a Option<IlluminationCalibrationState>,
446            calibrated_model: &'a Option<ImagePlaneModel>,
447            algorithm_auxiliary: &'a Option<AlgorithmAuxiliaryState>,
448            trace: &'a ReconstructionTrace,
449        }
450
451        Representation {
452            format_version: self.format_version,
453            completed_iterations: self.completed_iterations,
454            object_spectrum: Array2Data::from_view(self.object_spectrum.view()),
455            pupil: &self.pupil,
456            illumination_corrections: &self.illumination_corrections,
457            frame_gains: &self.frame_gains,
458            background: &self.background,
459            physical_illumination_calibration: &self.physical_illumination_calibration,
460            calibrated_model: &self.calibrated_model,
461            algorithm_auxiliary: &self.algorithm_auxiliary,
462            trace: &self.trace,
463        }
464        .serialize(serializer)
465    }
466}
467
468impl<'de> Deserialize<'de> for ReconstructionCheckpoint {
469    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
470    where
471        D: serde::Deserializer<'de>,
472    {
473        use serde::de::Error as _;
474
475        #[derive(Deserialize)]
476        #[serde(deny_unknown_fields)]
477        struct Representation {
478            format_version: u32,
479            completed_iterations: usize,
480            object_spectrum: Array2Data<Complex64>,
481            pupil: Pupil,
482            illumination_corrections: Option<Vec<(f64, f64)>>,
483            frame_gains: Option<Vec<f64>>,
484            background: Option<Vec<f64>>,
485            physical_illumination_calibration: Option<IlluminationCalibrationState>,
486            calibrated_model: Option<ImagePlaneModel>,
487            #[serde(default)]
488            algorithm_auxiliary: Option<AlgorithmAuxiliaryState>,
489            trace: ReconstructionTrace,
490        }
491
492        let representation = Representation::deserialize(deserializer)?;
493        let checkpoint = Self {
494            format_version: representation.format_version,
495            completed_iterations: representation.completed_iterations,
496            object_spectrum: representation
497                .object_spectrum
498                .into_array()
499                .map_err(D::Error::custom)?,
500            pupil: representation.pupil,
501            illumination_corrections: representation.illumination_corrections,
502            frame_gains: representation.frame_gains,
503            background: representation.background,
504            physical_illumination_calibration: representation.physical_illumination_calibration,
505            calibrated_model: representation.calibrated_model,
506            algorithm_auxiliary: representation.algorithm_auxiliary,
507            trace: representation.trace,
508        };
509        checkpoint.validate().map_err(D::Error::custom)?;
510        Ok(checkpoint)
511    }
512}