Skip to main content

fpm_rs/reconstruction/
state.rs

1use ndarray::{Array2, ArrayView2, ArrayViewMut2};
2use num_complex::Complex64;
3use serde::{Deserialize, Serialize};
4use std::sync::Arc;
5
6use crate::{
7    Result,
8    array_layout::{StandardArray2, StandardView2, checked_len_2d},
9    backend::{Backend, CpuBackend, FftDirection},
10    error::Error,
11    illumination_calibration::IlluminationCalibrationState,
12    measurements::MeasurementRead,
13    model::{FourierOffset, ImagePlaneModel, Pupil, fftshift_copy},
14};
15
16use super::{ReconstructionCheckpoint, ReconstructionProblem};
17
18/// Resumable per-source auxiliary and scaled-dual fields for ADMM.
19#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
20pub struct AdmmAuxiliaryState {
21    /// Row-major complex auxiliary detector fields for all individual sources.
22    pub auxiliary_fields: Vec<Complex64>,
23    /// Row-major complex scaled-dual fields with the same length and ordering.
24    pub dual_fields: Vec<Complex64>,
25}
26
27/// Resumable object-spectrum momentum state for mPIE.
28///
29/// `velocity` and `anchor` use centered, row-major reconstruction-spectrum
30/// ordering. The stored scalar parameters make the recurrence unambiguous and
31/// allow a resumed [`crate::algorithms::Mpie`] run to reject a configuration
32/// that would reinterpret its pending momentum interval.
33#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
34pub struct MpieAuxiliaryState {
35    /// Centered complex velocity spectrum accumulated at momentum events.
36    pub velocity: Vec<Complex64>,
37    /// Centered object spectrum immediately after the last momentum event.
38    pub anchor: Vec<Complex64>,
39    /// Positive-weight measured frames processed since the last event.
40    pub effective_frames_since_momentum: usize,
41    /// Per-frame rPIE object-correction scale used to create this state.
42    pub object_step: f64,
43    /// Local-to-maximum pupil-power blend used to create this state.
44    pub stability: f64,
45    /// Numerical denominator floor used to create this state.
46    pub epsilon: f64,
47    /// Effective-frame count between momentum events.
48    pub momentum_interval: usize,
49    /// Fraction of the previous velocity retained at each event.
50    pub momentum_friction: f64,
51    /// Fraction of the updated velocity added to the object spectrum.
52    pub momentum_feedback: f64,
53}
54
55/// Resumable cycle-level feedback state for adaptive alternating projection.
56///
57/// Objective sums use the same weighted, mask-aware amplitude-MSE convention
58/// as the reconstruction trace. Controller parameters are stored so resume
59/// cannot silently reinterpret an in-progress feedback cycle.
60#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
61pub struct AdaptiveAlternatingProjectionAuxiliaryState {
62    /// Zero-based iteration whose batches are currently accumulated.
63    pub active_iteration: usize,
64    /// Object relaxation used by the active iteration.
65    pub current_object_step: f64,
66    /// Objective of the pass preceding the active pass, when established.
67    pub previous_objective: Option<f64>,
68    /// Frame-weighted objective sum accumulated in the active pass.
69    pub objective_sum: f64,
70    /// Non-negative frame-weight sum accumulated in the active pass.
71    pub weight_sum: f64,
72    /// Number of scheduled frames accumulated in the active pass.
73    pub frames_accumulated: usize,
74    /// Initial object relaxation that created this controller state.
75    pub initial_object_step: f64,
76    /// Required relative objective decrease that created this state.
77    pub progress_threshold: f64,
78    /// Multiplicative reduction factor that created this state.
79    pub reduction_factor: f64,
80    /// Positive object-step floor that created this state.
81    pub minimum_object_step: f64,
82    /// Numerical floor used for relative-progress division.
83    pub epsilon: f64,
84}
85
86/// Solver-specific state preserved in checkpoints.
87#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
88pub enum AlgorithmAuxiliaryState {
89    /// ADMM auxiliary and scaled-dual fields.
90    Admm(AdmmAuxiliaryState),
91    /// mPIE object-spectrum velocity, anchor, cadence, and parameters.
92    Mpie(MpieAuxiliaryState),
93    /// Adaptive alternating-projection pass objective and controller state.
94    AdaptiveAlternatingProjection(AdaptiveAlternatingProjectionAuxiliaryState),
95}
96
97#[derive(Clone, Debug)]
98pub(crate) struct ReconstructionScratch {
99    pub(crate) patch: Vec<Complex64>,
100    pub(crate) exit_spectrum: Vec<Complex64>,
101    pub(crate) field: Vec<Complex64>,
102    pub(crate) projected_field: Vec<Complex64>,
103    pub(crate) projected_spectrum: Vec<Complex64>,
104    pub(crate) difference: Vec<Complex64>,
105    /// High-resolution accumulator used by mini-batch algorithms.
106    pub(crate) object_gradient: Vec<Complex64>,
107    pub(crate) regularization_field: Vec<Complex64>,
108    pub(crate) pupil_gradient: Vec<Complex64>,
109    pub(crate) calibration_reference: Vec<f64>,
110    /// Per-pixel inclusion flags for masked and robust data-gradient updates.
111    pub(crate) data_gradient_mask: Vec<u8>,
112    /// Shared precomputed truncation scale passed to parallel frame workers.
113    pub(crate) poisson_truncation_scale: Option<f64>,
114    pub(crate) illumination_gradient: Vec<(f64, f64)>,
115    pub(crate) illumination_curvature: Vec<(f64, f64)>,
116    pub(crate) illumination_weight: Vec<f64>,
117    /// Per-source low-resolution fields for an incoherently multiplexed frame.
118    pub(crate) multiplex_fields: Vec<Complex64>,
119    /// Matching pre-update object patches for multiplexed projection updates.
120    pub(crate) multiplex_patches: Vec<Complex64>,
121    pub(crate) multiplex_offsets: Vec<FourierOffset>,
122    pub(crate) column: Vec<Complex64>,
123}
124
125impl ReconstructionScratch {
126    fn new(low_shape: (usize, usize), high_shape: (usize, usize)) -> Result<Self> {
127        let low_len = checked_len_2d(low_shape)?;
128        checked_len_2d(high_shape)?;
129        Ok(Self {
130            patch: vec![Complex64::default(); low_len],
131            exit_spectrum: vec![Complex64::default(); low_len],
132            field: vec![Complex64::default(); low_len],
133            projected_field: vec![Complex64::default(); low_len],
134            projected_spectrum: vec![Complex64::default(); low_len],
135            difference: vec![Complex64::default(); low_len],
136            object_gradient: Vec::new(),
137            regularization_field: Vec::new(),
138            pupil_gradient: vec![Complex64::default(); low_len],
139            calibration_reference: vec![0.0; low_len],
140            data_gradient_mask: vec![1; low_len],
141            poisson_truncation_scale: None,
142            illumination_gradient: Vec::new(),
143            illumination_curvature: Vec::new(),
144            illumination_weight: Vec::new(),
145            multiplex_fields: Vec::new(),
146            multiplex_patches: Vec::new(),
147            multiplex_offsets: Vec::new(),
148            column: vec![Complex64::default(); low_shape.0.max(high_shape.0)],
149        })
150    }
151}
152
153/// Mutable numerical state shared by algorithms during reconstruction.
154///
155/// The object spectrum is centered and shaped on the high-resolution grid. The pupil and
156/// scratch arrays use the low-resolution grid. Accessors borrow storage without copying;
157/// mutating the spectrum invalidates the cached object-domain field.
158#[derive(Clone)]
159pub struct ReconstructionState {
160    pub(crate) object_spectrum: StandardArray2<Complex64>,
161    pub(crate) object_real_space_cache: Option<StandardArray2<Complex64>>,
162    pub(crate) pupil: Pupil,
163    /// Per-source `(row, column)` corrections in Fourier-grid pixels.
164    pub(crate) illumination_corrections: Option<Vec<(f64, f64)>>,
165    pub(crate) frame_gains: Option<Vec<f64>>,
166    pub(crate) background: Option<Vec<f64>>,
167    pub(crate) physical_illumination_calibration: Option<IlluminationCalibrationState>,
168    pub(crate) calibrated_model: Option<ImagePlaneModel>,
169    pub(crate) algorithm_auxiliary: Option<AlgorithmAuxiliaryState>,
170    pub(crate) scratch: ReconstructionScratch,
171    pub(crate) backend: Arc<dyn Backend>,
172}
173
174impl std::fmt::Debug for ReconstructionState {
175    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
176        formatter
177            .debug_struct("ReconstructionState")
178            .field("object_spectrum", &self.object_spectrum)
179            .field("pupil", &self.pupil)
180            .field("illumination_corrections", &self.illumination_corrections)
181            .field("frame_gains", &self.frame_gains)
182            .field("background", &self.background)
183            .field(
184                "physical_illumination_calibration",
185                &self.physical_illumination_calibration,
186            )
187            .field("calibrated_model", &self.calibrated_model)
188            .field("algorithm_auxiliary", &self.algorithm_auxiliary)
189            .finish_non_exhaustive()
190    }
191}
192
193impl ReconstructionState {
194    /// Borrows the centered high-resolution complex object spectrum.
195    pub fn object_spectrum(&self) -> ArrayView2<'_, Complex64> {
196        self.object_spectrum.ndarray_view()
197    }
198
199    /// Mutably borrows the object spectrum and invalidates the object-domain cache.
200    pub fn object_spectrum_mut(&mut self) -> ArrayViewMut2<'_, Complex64> {
201        self.object_real_space_cache = None;
202        self.object_spectrum.ndarray_view_mut()
203    }
204
205    /// Borrows the current low-resolution complex pupil and binary support.
206    pub fn pupil(&self) -> &Pupil {
207        &self.pupil
208    }
209
210    /// Borrows optional per-source `(row, column)` corrections in Fourier-grid pixels.
211    pub fn illumination_corrections(&self) -> Option<&[(f64, f64)]> {
212        self.illumination_corrections.as_deref()
213    }
214
215    /// Borrows optional positive gains in acquisition-frame order.
216    pub fn frame_gains(&self) -> Option<&[f64]> {
217        self.frame_gains.as_deref()
218    }
219
220    /// Borrows optional additive backgrounds in acquisition-frame order.
221    pub fn background(&self) -> Option<&[f64]> {
222        self.background.as_deref()
223    }
224
225    /// Borrows checkpointable physical planar-array calibration state, when active.
226    pub fn physical_illumination_calibration(&self) -> Option<&IlluminationCalibrationState> {
227        self.physical_illumination_calibration.as_ref()
228    }
229
230    /// Borrows the illumination-refreshed model used by joint reconstruction, when active.
231    pub fn calibrated_model(&self) -> Option<&ImagePlaneModel> {
232        self.calibrated_model.as_ref()
233    }
234
235    /// Most recently accumulated per-source illumination gradient.
236    pub fn illumination_gradient(&self) -> &[(f64, f64)] {
237        &self.scratch.illumination_gradient
238    }
239
240    pub(crate) fn object_spectrum_standard_view(&self) -> StandardView2<'_, Complex64> {
241        self.object_spectrum.view()
242    }
243
244    /// Adds the current calibration correction to a model source offset and validates bounds.
245    pub fn effective_source_offset(
246        &self,
247        model: &ImagePlaneModel,
248        source: usize,
249    ) -> Result<FourierOffset> {
250        let base = model.source_offset(source)?;
251        let correction = match &self.illumination_corrections {
252            None => (0.0, 0.0),
253            Some(values) => *values.get(source).ok_or_else(|| {
254                Error::InvalidModel(
255                    "illumination correction count does not match source count".into(),
256                )
257            })?,
258        };
259        if !correction.0.is_finite() || !correction.1.is_finite() {
260            return Err(Error::InvalidModel(
261                "illumination corrections must be finite".into(),
262            ));
263        }
264        let effective = FourierOffset::new(base.row + correction.0, base.column + correction.1);
265        model.validate_source_offset(source, effective)?;
266        Ok(effective)
267    }
268
269    /// Initializes a CPU-backed object estimate from weighted measured amplitudes.
270    pub fn initialize<M: MeasurementRead>(problem: &ReconstructionProblem<M>) -> Result<Self> {
271        let backend: Arc<dyn Backend> = Arc::new(CpuBackend::new(
272            problem.model.image_shape,
273            problem.model.reconstruction_shape,
274        )?);
275        Self::initialize_with_backend(problem, backend)
276    }
277
278    /// Initializes an object estimate, pupil, calibration variables, and scratch storage
279    /// using `backend` after validating `problem`.
280    pub fn initialize_with_backend<M: MeasurementRead>(
281        problem: &ReconstructionProblem<M>,
282        backend: Arc<dyn Backend>,
283    ) -> Result<Self> {
284        problem.validate()?;
285        let low_shape = problem.model.image_shape;
286        let high_shape = problem.model.reconstruction_shape;
287        let low_len = checked_len_2d(low_shape)?;
288        let mut average_amplitude = vec![0.0; low_len];
289        let mut amplitude_weight = vec![0.0; low_len];
290        let mut total_amplitude = 0.0;
291        let mut total_weight = 0.0;
292        for frame_index in 0..problem.measurements.frame_count() {
293            let frame = problem.measurements.frame(frame_index)?;
294            let frame_weight = problem.measurements.frame_weight(frame_index)?;
295            if frame_weight == 0.0 {
296                continue;
297            }
298            let mask = problem.measurements.frame_mask(frame_index)?;
299            let gain = problem.model.frame_gain(frame_index)?;
300            for (pixel, (average, &intensity)) in
301                average_amplitude.iter_mut().zip(frame.iter()).enumerate()
302            {
303                if mask.is_some_and(|values| values[pixel] == 0) {
304                    continue;
305                }
306                let background = problem.model.background_value(frame_index, pixel)?;
307                let amplitude = ((intensity - background) / gain).max(0.0).sqrt();
308                *average += frame_weight * amplitude;
309                amplitude_weight[pixel] += frame_weight;
310                total_amplitude += frame_weight * amplitude;
311                total_weight += frame_weight;
312            }
313        }
314        if total_weight == 0.0 {
315            return Err(Error::InvalidMeasurements(
316                "no positive-weight, unmasked measurements are available".into(),
317            ));
318        }
319        let fallback_amplitude = total_amplitude / total_weight;
320        for (value, &weight) in average_amplitude.iter_mut().zip(&amplitude_weight) {
321            *value = if weight > 0.0 {
322                *value / weight
323            } else {
324                fallback_amplitude
325            };
326        }
327        let high_len = checked_len_2d(high_shape)?;
328        let mut object = vec![Complex64::default(); high_len];
329        for row in 0..high_shape.0 {
330            let low_row = row * low_shape.0 / high_shape.0;
331            for column in 0..high_shape.1 {
332                let low_column = column * low_shape.1 / high_shape.1;
333                object[row * high_shape.1 + column] =
334                    Complex64::new(average_amplitude[low_row * low_shape.1 + low_column], 0.0);
335            }
336        }
337        let mut column = vec![Complex64::default(); low_shape.0.max(high_shape.0)];
338        backend.fft2(&mut object, high_shape, FftDirection::Forward, &mut column)?;
339        let mut centered = vec![Complex64::default(); object.len()];
340        fftshift_copy(&object, &mut centered, high_shape);
341        let object_spectrum = StandardArray2::from_shape_vec(high_shape, centered)?;
342        Ok(Self {
343            object_spectrum,
344            object_real_space_cache: None,
345            pupil: problem.model.pupil.clone(),
346            illumination_corrections: None,
347            frame_gains: problem.model.frame_gains.clone(),
348            background: problem.model.background.clone(),
349            physical_illumination_calibration: None,
350            calibrated_model: None,
351            algorithm_auxiliary: None,
352            scratch: ReconstructionScratch::new(low_shape, high_shape)?,
353            backend,
354        })
355    }
356
357    /// Initializes state from an owned standard-layout high-resolution complex object.
358    pub fn from_object<M: MeasurementRead>(
359        problem: &ReconstructionProblem<M>,
360        object: Array2<Complex64>,
361    ) -> Result<Self> {
362        let mut object = StandardArray2::try_from(object)?;
363        if object.dim() != problem.model.reconstruction_shape {
364            return Err(Error::InvalidShape(format!(
365                "initial object shape {:?} differs from reconstruction shape {:?}",
366                object.dim(),
367                problem.model.reconstruction_shape
368            )));
369        }
370        let mut state = Self::initialize(problem)?;
371        state.backend.fft2(
372            object.as_slice_mut(),
373            problem.model.reconstruction_shape,
374            FftDirection::Forward,
375            &mut state.scratch.column,
376        )?;
377        fftshift_copy(
378            object.as_slice(),
379            state.object_spectrum.as_slice_mut(),
380            problem.model.reconstruction_shape,
381        );
382        Ok(state)
383    }
384
385    /// Restores CPU-backed state from a checkpoint validated against `problem`.
386    pub fn from_checkpoint<M: MeasurementRead>(
387        problem: &ReconstructionProblem<M>,
388        checkpoint: &ReconstructionCheckpoint,
389    ) -> Result<Self> {
390        let backend: Arc<dyn Backend> = Arc::new(CpuBackend::new(
391            problem.model.image_shape,
392            problem.model.reconstruction_shape,
393        )?);
394        Self::from_checkpoint_with_backend(problem, checkpoint, backend)
395    }
396
397    /// Restores state from a compatible checkpoint using `backend`.
398    pub fn from_checkpoint_with_backend<M: MeasurementRead>(
399        problem: &ReconstructionProblem<M>,
400        checkpoint: &ReconstructionCheckpoint,
401        backend: Arc<dyn Backend>,
402    ) -> Result<Self> {
403        checkpoint.validate_for_problem(problem)?;
404        let low_shape = problem.model.image_shape;
405        let high_shape = problem.model.reconstruction_shape;
406        Ok(Self {
407            object_spectrum: StandardArray2::try_from(checkpoint.object_spectrum.clone())?,
408            object_real_space_cache: None,
409            pupil: checkpoint.pupil.clone(),
410            illumination_corrections: checkpoint.illumination_corrections.clone(),
411            frame_gains: checkpoint.frame_gains.clone(),
412            background: checkpoint.background.clone(),
413            physical_illumination_calibration: checkpoint.physical_illumination_calibration.clone(),
414            calibrated_model: checkpoint.calibrated_model.clone(),
415            algorithm_auxiliary: checkpoint.algorithm_auxiliary.clone(),
416            scratch: ReconstructionScratch::new(low_shape, high_shape)?,
417            backend,
418        })
419    }
420}