Skip to main content

fpm_rs/algorithms/
joint_reconstruction.rs

1use std::path::Path;
2
3use serde::{Deserialize, Serialize};
4
5use crate::{
6    Result,
7    callbacks::Callback,
8    error::Error,
9    experiment::{Illumination, Optics},
10    illumination_calibration::{
11        CalibrationConditioning, CalibrationConvergenceReason, CalibrationLossHistoryEntry,
12        CalibrationParameterHistoryEntry, IlluminationCalibration, IlluminationCalibrationState,
13        PlanarArrayParameterValues,
14    },
15    measurements::MeasurementRead,
16    model::{ImagePlaneModel, ReconstructionShape},
17    reconstruction::{
18        AlgorithmMetricRecord, Batch, ReconstructionCheckpoint, ReconstructionProblem,
19        ReconstructionResult, ReconstructionState, RunOptions, Runner,
20    },
21};
22
23use super::{AlgorithmIterationMetrics, ReconstructionAlgorithm, StepOutput, StepSummary};
24
25/// Scalar metrics emitted after one alternating object/illumination outer iteration.
26#[derive(Clone, Debug, Default)]
27pub struct JointIterationMetrics {
28    data_loss: Option<f64>,
29    regularization_loss: Option<f64>,
30    total_loss: Option<f64>,
31    geometry_recompilations: usize,
32    multiplicative_updates: usize,
33    rejected_steps: usize,
34    object_metrics: Vec<(String, String, f64)>,
35}
36
37impl AlgorithmIterationMetrics for JointIterationMetrics {
38    fn merge(&mut self, other: Self) {
39        self.data_loss = other.data_loss.or(self.data_loss);
40        self.regularization_loss = other.regularization_loss.or(self.regularization_loss);
41        self.total_loss = other.total_loss.or(self.total_loss);
42        self.geometry_recompilations += other.geometry_recompilations;
43        self.multiplicative_updates += other.multiplicative_updates;
44        self.rejected_steps += other.rejected_steps;
45        self.object_metrics.extend(other.object_metrics);
46    }
47
48    fn append_records(&self, iteration: usize, output: &mut Vec<AlgorithmMetricRecord>) {
49        for (metric, value) in [
50            ("data_loss", self.data_loss),
51            ("regularization_loss", self.regularization_loss),
52            ("total_loss", self.total_loss),
53        ] {
54            if let Some(value) = value {
55                output.push(AlgorithmMetricRecord {
56                    iteration,
57                    namespace: "physical_illumination".into(),
58                    metric: metric.into(),
59                    value,
60                });
61            }
62        }
63        for (metric, value) in [
64            (
65                "geometry_recompilations",
66                self.geometry_recompilations as f64,
67            ),
68            ("multiplicative_updates", self.multiplicative_updates as f64),
69            ("rejected_steps", self.rejected_steps as f64),
70        ] {
71            output.push(AlgorithmMetricRecord {
72                iteration,
73                namespace: "physical_illumination".into(),
74                metric: metric.into(),
75                value,
76            });
77        }
78        output.extend(
79            self.object_metrics
80                .iter()
81                .map(|(namespace, metric, value)| AlgorithmMetricRecord {
82                    iteration,
83                    namespace: format!("object_update.{namespace}"),
84                    metric: metric.clone(),
85                    value: *value,
86                }),
87        );
88    }
89}
90
91/// Alternates an existing object or object/pupil algorithm with physical LED calibration.
92///
93/// One runner iteration is one outer iteration. The complete acquisition order is supplied
94/// to the wrapped algorithm `object_iterations_per_outer` times, then bounded physical
95/// updates refresh only illumination-dependent model state. Generic Fourier-grid correction
96/// remains an independent feature of [`crate::algorithms::GradientDescent`].
97///
98/// The joint-estimation motivation follows [J. Sun, Q. Chen, Y. Zhang, and C. Zuo,
99/// “Efficient positional misalignment correction method for Fourier ptychographic
100/// microscopy,” *Biomedical Optics Express* **7**(4), 1336–1350
101/// (2016)](https://doi.org/10.1364/BOE.7.001336). This implementation uses bounded,
102/// scaled finite differences rather than that work's simulated annealing and nonlinear
103/// regression. It does not implement the brightfield circle detection or spectral
104/// correlation of [R. Eckert, Z. F. Phillips, and L. Waller, “Efficient illumination
105/// angle self-calibration in Fourier ptychography,” *Applied Optics* **57**(19),
106/// 5434–5442 (2018)](https://doi.org/10.1364/AO.57.005434).
107#[derive(Clone, Debug)]
108pub struct JointReconstruction<A> {
109    /// Existing analytic object or object/pupil update algorithm.
110    pub object_algorithm: A,
111    /// Optical configuration used to resolve physical source positions.
112    pub optics: Optics,
113    /// Nominal reusable physical illumination.
114    pub initial_illumination: Illumination,
115    /// Physical parameter selection, priors, loss, and bounded optimizer.
116    pub illumination_calibration: IlluminationCalibration,
117    /// Number of alternating outer iterations.
118    pub outer_iterations: usize,
119    /// Complete object-update passes before each physical phase.
120    pub object_iterations_per_outer: usize,
121    /// Repetitions of the configured bounded physical phase per outer iteration.
122    pub illumination_steps_per_outer: usize,
123}
124
125impl<A> JointReconstruction<A> {
126    /// Creates an alternating reconstruction with one object pass and one calibration phase.
127    pub fn new(
128        object_algorithm: A,
129        optics: Optics,
130        initial_illumination: Illumination,
131        illumination_calibration: IlluminationCalibration,
132        outer_iterations: usize,
133    ) -> Self {
134        Self {
135            object_algorithm,
136            optics,
137            initial_illumination,
138            illumination_calibration,
139            outer_iterations,
140            object_iterations_per_outer: 1,
141            illumination_steps_per_outer: 1,
142        }
143    }
144
145    /// Sets the positive number of complete object passes in each outer iteration.
146    pub fn object_iterations_per_outer(mut self, iterations: usize) -> Self {
147        self.object_iterations_per_outer = iterations;
148        self
149    }
150
151    /// Sets the positive number of bounded illumination phases in each outer iteration.
152    pub fn illumination_steps_per_outer(mut self, steps: usize) -> Self {
153        self.illumination_steps_per_outer = steps;
154        self
155    }
156}
157
158impl<A: ReconstructionAlgorithm> JointReconstruction<A> {
159    /// Runs alternating reconstruction without callbacks and returns structured physical output.
160    pub fn run<M: MeasurementRead>(
161        self,
162        problem: &ReconstructionProblem<M>,
163    ) -> Result<JointReconstructionResult> {
164        let options = RunOptions {
165            max_iterations: self.outer_iterations,
166            batch_size: usize::MAX,
167            ..RunOptions::default()
168        };
169        JointReconstructionResult::from_reconstruction(Runner::new(self, options).run(problem)?)
170    }
171
172    /// Runs alternating reconstruction with ordinary runner callbacks.
173    ///
174    /// Outer-iteration callbacks receive physical phase metrics in the
175    /// `physical_illumination` namespace, and checkpoint callbacks capture the complete
176    /// physical state and refreshed model.
177    pub fn run_with_callbacks<M: MeasurementRead>(
178        self,
179        problem: &ReconstructionProblem<M>,
180        callbacks: Vec<Box<dyn Callback>>,
181    ) -> Result<JointReconstructionResult> {
182        let options = RunOptions {
183            max_iterations: self.outer_iterations,
184            batch_size: usize::MAX,
185            ..RunOptions::default()
186        };
187        JointReconstructionResult::from_reconstruction(
188            Runner::new(self, options)
189                .with_callbacks(callbacks)
190                .run(problem)?,
191        )
192    }
193
194    /// Resumes alternating reconstruction from a checkpoint containing physical state.
195    pub fn run_from_checkpoint<M: MeasurementRead>(
196        self,
197        problem: &ReconstructionProblem<M>,
198        checkpoint: ReconstructionCheckpoint,
199    ) -> Result<JointReconstructionResult> {
200        let options = RunOptions {
201            max_iterations: self.outer_iterations,
202            batch_size: usize::MAX,
203            ..RunOptions::default()
204        };
205        JointReconstructionResult::from_reconstruction(
206            Runner::new(self, options)
207                .resume_from(checkpoint)
208                .run(problem)?,
209        )
210    }
211}
212
213impl<A: ReconstructionAlgorithm> ReconstructionAlgorithm for JointReconstruction<A> {
214    type IterationMetrics = JointIterationMetrics;
215
216    fn validate(&self) -> Result<()> {
217        self.object_algorithm.validate()?;
218        if !self.object_algorithm.supports_joint_reconstruction() {
219            return Err(Error::InvalidParameter {
220                name: "object_algorithm",
221                reason: "the selected algorithm carries state that cannot be transported across physical model recompilation"
222                    .into(),
223            });
224        }
225        self.illumination_calibration
226            .validate_for(&self.initial_illumination)?;
227        self.optics.validate()?;
228        if self.outer_iterations == 0
229            || self.object_iterations_per_outer == 0
230            || self.illumination_steps_per_outer == 0
231        {
232            return Err(Error::InvalidParameter {
233                name: "joint iteration counts",
234                reason: "outer, object, and illumination iteration counts must be positive".into(),
235            });
236        }
237        Ok(())
238    }
239
240    fn validate_problem<M: MeasurementRead>(
241        &self,
242        problem: &ReconstructionProblem<M>,
243    ) -> Result<()> {
244        self.validate()?;
245        self.object_algorithm.validate_problem(problem)?;
246        let expected = ImagePlaneModel::from_experiment(
247            &self.optics,
248            &self.initial_illumination,
249            problem.model.image_shape(),
250            ReconstructionShape::Exact(problem.model.reconstruction_shape()),
251        )?;
252        if expected.source_count() != problem.model.source_count()
253            || expected.frame_count() != problem.model.frame_count()
254            || expected.k_vectors() != problem.model.k_vectors()
255        {
256            return Err(Error::InvalidModel(
257                "joint reconstruction problem model was not compiled from the supplied initial illumination"
258                    .into(),
259            ));
260        }
261        Ok(())
262    }
263
264    fn initialize<M: MeasurementRead>(
265        &self,
266        problem: &ReconstructionProblem<M>,
267    ) -> Result<ReconstructionState> {
268        let mut state = self.object_algorithm.initialize(problem)?;
269        let calibration_state = self
270            .illumination_calibration
271            .initialize(&self.initial_illumination, &self.optics)?;
272        let mut calibrated_model = problem.model.clone();
273        self.illumination_calibration.synchronize_model(
274            &self.optics,
275            &mut calibrated_model,
276            &self.initial_illumination,
277            &calibration_state,
278        )?;
279        state.frame_gains = calibrated_model.frame_gains().map(<[f64]>::to_vec);
280        state.physical_illumination_calibration = Some(calibration_state);
281        state.calibrated_model = Some(calibrated_model);
282        Ok(state)
283    }
284
285    fn initialize_with_backend<M: MeasurementRead>(
286        &self,
287        problem: &ReconstructionProblem<M>,
288        backend: std::sync::Arc<dyn crate::backend::Backend>,
289    ) -> Result<ReconstructionState> {
290        let mut state = self
291            .object_algorithm
292            .initialize_with_backend(problem, backend)?;
293        let calibration_state = self
294            .illumination_calibration
295            .initialize(&self.initial_illumination, &self.optics)?;
296        let mut calibrated_model = problem.model.clone();
297        self.illumination_calibration.synchronize_model(
298            &self.optics,
299            &mut calibrated_model,
300            &self.initial_illumination,
301            &calibration_state,
302        )?;
303        state.frame_gains = calibrated_model.frame_gains().map(<[f64]>::to_vec);
304        state.physical_illumination_calibration = Some(calibration_state);
305        state.calibrated_model = Some(calibrated_model);
306        Ok(state)
307    }
308
309    fn step<M: MeasurementRead>(
310        &mut self,
311        problem: &ReconstructionProblem<M>,
312        state: &mut ReconstructionState,
313        batch: &Batch,
314        iteration: usize,
315    ) -> Result<StepOutput<Self::IterationMetrics>> {
316        let mut effective_model = state
317            .calibrated_model
318            .clone()
319            .ok_or_else(|| Error::InvalidModel("joint calibrated model is missing".into()))?;
320        let local_problem = ReconstructionProblem {
321            measurements: &problem.measurements,
322            model: effective_model.clone(),
323            name: problem.name.clone(),
324        };
325        let mut summary = StepSummary::default();
326        let mut object_metrics = Vec::new();
327        for _ in 0..self.object_iterations_per_outer {
328            let output = self
329                .object_algorithm
330                .step(&local_problem, state, batch, iteration)?;
331            self.object_algorithm
332                .canonicalize_state(&local_problem, state)?;
333            summary.merge(output.summary);
334            let mut records = Vec::new();
335            output.metrics.append_records(iteration + 1, &mut records);
336            object_metrics.extend(
337                records
338                    .into_iter()
339                    .map(|record| (record.namespace, record.metric, record.value)),
340            );
341        }
342
343        let geometry_before = state
344            .physical_illumination_calibration
345            .as_ref()
346            .ok_or_else(|| {
347                Error::InvalidModel("joint physical calibration state is missing".into())
348            })?
349            .geometry_recompilations;
350        let multiplicative_before = state
351            .physical_illumination_calibration
352            .as_ref()
353            .expect("checked above")
354            .multiplicative_updates;
355        let rejected_before = state
356            .physical_illumination_calibration
357            .as_ref()
358            .expect("checked above")
359            .rejected_steps;
360        let mut objective = None;
361        for _ in 0..self.illumination_steps_per_outer {
362            objective = Some(
363                self.illumination_calibration.optimize(
364                    &problem.measurements,
365                    &self.optics,
366                    state.object_spectrum.ndarray_view(),
367                    &state.pupil,
368                    &mut effective_model,
369                    state
370                        .physical_illumination_calibration
371                        .as_mut()
372                        .expect("checked above"),
373                    iteration + 1,
374                )?,
375            );
376        }
377        state.frame_gains = effective_model.frame_gains().map(<[f64]>::to_vec);
378        state.calibrated_model = Some(effective_model);
379        let calibration = state
380            .physical_illumination_calibration
381            .as_ref()
382            .expect("checked above");
383        let objective = objective.expect("positive illumination phase count was validated");
384        Ok(StepOutput {
385            summary,
386            metrics: JointIterationMetrics {
387                data_loss: Some(objective.data_loss),
388                regularization_loss: Some(objective.regularization_loss),
389                total_loss: Some(objective.total_loss),
390                geometry_recompilations: calibration.geometry_recompilations - geometry_before,
391                multiplicative_updates: calibration.multiplicative_updates - multiplicative_before,
392                rejected_steps: calibration.rejected_steps - rejected_before,
393                object_metrics,
394            },
395        })
396    }
397
398    fn canonicalize_state<M: MeasurementRead>(
399        &self,
400        problem: &ReconstructionProblem<M>,
401        state: &mut ReconstructionState,
402    ) -> Result<()> {
403        let model = state
404            .calibrated_model
405            .clone()
406            .unwrap_or_else(|| problem.model.clone());
407        let local_problem = ReconstructionProblem {
408            measurements: &problem.measurements,
409            model,
410            name: problem.name.clone(),
411        };
412        self.object_algorithm
413            .canonicalize_state(&local_problem, state)
414    }
415
416    fn iterations(&self) -> usize {
417        self.outer_iterations
418    }
419
420    fn batch_size(&self) -> usize {
421        usize::MAX
422    }
423}
424
425/// Structured joint result with a normal reconstruction and reusable physical illumination.
426#[derive(Clone, Debug, Serialize, Deserialize)]
427#[serde(deny_unknown_fields)]
428pub struct JointReconstructionResult {
429    /// Ordinary reconstruction fields, trace, diagnostics, and calibration state.
430    pub reconstruction: ReconstructionResult,
431    /// Gauge-normalized initial physical illumination.
432    pub initial_illumination: Illumination,
433    /// Final reusable physical illumination.
434    pub calibrated_illumination: Illumination,
435    /// Final image-plane model refreshed from the calibrated illumination.
436    pub calibrated_model: ImagePlaneModel,
437    /// Immutable initial absolute physical and multiplicative values.
438    pub initial_parameters: PlanarArrayParameterValues,
439    /// Final absolute physical and multiplicative values.
440    pub final_parameters: PlanarArrayParameterValues,
441    /// Accepted and rejected parameter trials.
442    pub parameter_history: Vec<CalibrationParameterHistoryEntry>,
443    /// Data, regularization, and total objective history.
444    pub loss_history: Vec<CalibrationLossHistoryEntry>,
445    /// Stop reason for the last illumination phase.
446    pub convergence_reason: Option<CalibrationConvergenceReason>,
447    /// Practical finite-difference conditioning indicators.
448    pub conditioning: CalibrationConditioning,
449    /// Complete checkpointable physical state and update counters.
450    pub diagnostics: IlluminationCalibrationState,
451}
452
453impl JointReconstructionResult {
454    /// Builds the structured view from a runner result containing physical state.
455    pub fn from_reconstruction(reconstruction: ReconstructionResult) -> Result<Self> {
456        let diagnostics = reconstruction
457            .physical_illumination_calibration
458            .clone()
459            .ok_or_else(|| {
460                Error::InvalidModel("joint result has no physical calibration".into())
461            })?;
462        let calibrated_model = reconstruction
463            .calibrated_model
464            .clone()
465            .ok_or_else(|| Error::InvalidModel("joint result has no calibrated model".into()))?;
466        Ok(Self {
467            initial_illumination: diagnostics.initial_illumination.clone(),
468            calibrated_illumination: diagnostics.current_illumination.clone(),
469            calibrated_model,
470            initial_parameters: diagnostics.initial_parameters.clone(),
471            final_parameters: diagnostics.current_parameters.clone(),
472            parameter_history: diagnostics.parameter_history.clone(),
473            loss_history: diagnostics.loss_history.clone(),
474            convergence_reason: diagnostics.convergence_reason,
475            conditioning: diagnostics.conditioning.clone(),
476            diagnostics,
477            reconstruction,
478        })
479    }
480
481    /// Serializes the complete structured joint result as JSON.
482    pub fn save_json(&self, path: impl AsRef<Path>) -> Result<()> {
483        let writer = std::io::BufWriter::new(std::fs::File::create(path)?);
484        serde_json::to_writer(writer, self)?;
485        Ok(())
486    }
487
488    /// Loads and validates a structured joint JSON result.
489    pub fn load_json(path: impl AsRef<Path>) -> Result<Self> {
490        let reader = std::io::BufReader::new(std::fs::File::open(path)?);
491        let result: Self = serde_json::from_reader(reader)?;
492        result.reconstruction.validate()?;
493        result.calibrated_model.validate()?;
494        Ok(result)
495    }
496
497    /// Writes the ordinary and physical result through the established bundle path.
498    #[cfg(feature = "parquet")]
499    pub fn write_bundle(
500        &self,
501        path: impl AsRef<Path>,
502        options: crate::reconstruction::BundleExportOptions,
503    ) -> Result<crate::reconstruction::ResultBundle> {
504        self.reconstruction.write_bundle(path, options)
505    }
506}