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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
20pub struct AdmmAuxiliaryState {
21 pub auxiliary_fields: Vec<Complex64>,
23 pub dual_fields: Vec<Complex64>,
25}
26
27#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
34pub struct MpieAuxiliaryState {
35 pub velocity: Vec<Complex64>,
37 pub anchor: Vec<Complex64>,
39 pub effective_frames_since_momentum: usize,
41 pub object_step: f64,
43 pub stability: f64,
45 pub epsilon: f64,
47 pub momentum_interval: usize,
49 pub momentum_friction: f64,
51 pub momentum_feedback: f64,
53}
54
55#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
61pub struct AdaptiveAlternatingProjectionAuxiliaryState {
62 pub active_iteration: usize,
64 pub current_object_step: f64,
66 pub previous_objective: Option<f64>,
68 pub objective_sum: f64,
70 pub weight_sum: f64,
72 pub frames_accumulated: usize,
74 pub initial_object_step: f64,
76 pub progress_threshold: f64,
78 pub reduction_factor: f64,
80 pub minimum_object_step: f64,
82 pub epsilon: f64,
84}
85
86#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
88pub enum AlgorithmAuxiliaryState {
89 Admm(AdmmAuxiliaryState),
91 Mpie(MpieAuxiliaryState),
93 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 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 pub(crate) data_gradient_mask: Vec<u8>,
112 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 pub(crate) multiplex_fields: Vec<Complex64>,
119 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#[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 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 pub fn object_spectrum(&self) -> ArrayView2<'_, Complex64> {
196 self.object_spectrum.ndarray_view()
197 }
198
199 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 pub fn pupil(&self) -> &Pupil {
207 &self.pupil
208 }
209
210 pub fn illumination_corrections(&self) -> Option<&[(f64, f64)]> {
212 self.illumination_corrections.as_deref()
213 }
214
215 pub fn frame_gains(&self) -> Option<&[f64]> {
217 self.frame_gains.as_deref()
218 }
219
220 pub fn background(&self) -> Option<&[f64]> {
222 self.background.as_deref()
223 }
224
225 pub fn physical_illumination_calibration(&self) -> Option<&IlluminationCalibrationState> {
227 self.physical_illumination_calibration.as_ref()
228 }
229
230 pub fn calibrated_model(&self) -> Option<&ImagePlaneModel> {
232 self.calibrated_model.as_ref()
233 }
234
235 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 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 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 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(&litude_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 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 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 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}