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
25pub const CHECKPOINT_FORMAT_VERSION: u32 = 2;
27
28#[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 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 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 pub const fn format_version(&self) -> u32 {
77 self.format_version
78 }
79
80 pub const fn completed_iterations(&self) -> usize {
82 self.completed_iterations
83 }
84
85 pub fn object_spectrum(&self) -> ArrayView2<'_, Complex64> {
87 self.object_spectrum.view()
88 }
89
90 pub fn pupil(&self) -> &Pupil {
92 &self.pupil
93 }
94
95 pub fn illumination_corrections(&self) -> Option<&[(f64, f64)]> {
97 self.illumination_corrections.as_deref()
98 }
99
100 pub fn frame_gains(&self) -> Option<&[f64]> {
102 self.frame_gains.as_deref()
103 }
104
105 pub fn background(&self) -> Option<&[f64]> {
107 self.background.as_deref()
108 }
109
110 pub fn physical_illumination_calibration(&self) -> Option<&IlluminationCalibrationState> {
112 self.physical_illumination_calibration.as_ref()
113 }
114
115 pub fn calibrated_model(&self) -> Option<&ImagePlaneModel> {
117 self.calibrated_model.as_ref()
118 }
119
120 pub fn algorithm_auxiliary(&self) -> Option<&AlgorithmAuxiliaryState> {
122 self.algorithm_auxiliary.as_ref()
123 }
124
125 pub const fn trace(&self) -> &ReconstructionTrace {
127 &self.trace
128 }
129
130 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 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 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 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 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}