Skip to main content

fpm_rs/algorithms/
admm.rs

1use num_complex::Complex64;
2
3use crate::{
4    Result,
5    algorithms::{
6        AlgorithmIterationMetrics, StepOutput, StepSummary,
7        objective::{LossType, point_loss},
8    },
9    array_layout::checked_len_2d,
10    backend::FftDirection,
11    error::Error,
12    measurements::MeasurementRead,
13    model::{FourierOffset, ImagePlaneModel, fftshift_copy, ifftshift_copy},
14    reconstruction::{
15        AdmmAuxiliaryState, AlgorithmAuxiliaryState, Batch, ReconstructionProblem,
16        ReconstructionState,
17    },
18};
19
20use super::ReconstructionAlgorithm;
21
22/// ADMM consensus metrics accumulated over every active detector-field mode in
23/// one complete iteration.
24#[derive(Clone, Debug, Default)]
25pub struct AdmmIterationMetrics {
26    primal_residual_sum_squares: f64,
27    dual_residual_sum_squares: f64,
28    residual_count: usize,
29}
30
31impl AdmmIterationMetrics {
32    /// Returns the root-mean-square consensus residual, or `None` before any mode is visited.
33    pub fn primal_residual_rms(&self) -> Option<f64> {
34        (self.residual_count > 0)
35            .then(|| (self.primal_residual_sum_squares / self.residual_count as f64).sqrt())
36    }
37
38    /// Returns the penalty-scaled RMS auxiliary-field change, or `None` when empty.
39    pub fn dual_residual_rms(&self) -> Option<f64> {
40        (self.residual_count > 0)
41            .then(|| (self.dual_residual_sum_squares / self.residual_count as f64).sqrt())
42    }
43
44    fn push_dual_change(&mut self, change: Complex64, penalty: f64) {
45        self.dual_residual_sum_squares += penalty * penalty * change.norm_sqr();
46    }
47
48    fn push_primal_residual(&mut self, residual: Complex64) {
49        self.primal_residual_sum_squares += residual.norm_sqr();
50        self.residual_count += 1;
51    }
52}
53
54impl AlgorithmIterationMetrics for AdmmIterationMetrics {
55    fn merge(&mut self, other: Self) {
56        self.primal_residual_sum_squares += other.primal_residual_sum_squares;
57        self.dual_residual_sum_squares += other.dual_residual_sum_squares;
58        self.residual_count += other.residual_count;
59    }
60
61    fn append_records(
62        &self,
63        iteration: usize,
64        output: &mut Vec<crate::reconstruction::AlgorithmMetricRecord>,
65    ) {
66        if let Some(value) = self.primal_residual_rms() {
67            output.push(crate::reconstruction::AlgorithmMetricRecord {
68                iteration,
69                namespace: "admm".into(),
70                metric: "primal_residual_rms".into(),
71                value,
72            });
73        }
74        if let Some(value) = self.dual_residual_rms() {
75            output.push(crate::reconstruction::AlgorithmMetricRecord {
76                iteration,
77                namespace: "admm".into(),
78                metric: "dual_residual_rms".into(),
79                value,
80            });
81        }
82    }
83}
84
85/// Linearized ADMM reconstruction for Fourier ptychographic microscopy.
86///
87/// # Method
88///
89/// The algorithm splits each predicted detector field from the shared object
90/// by introducing an auxiliary field and a scaled dual variable. Each step
91/// alternates among a measurement-amplitude proximal update of the auxiliary
92/// fields, a pupil-preconditioned linearized update of the common object
93/// spectrum, and a scaled-dual update that drives the auxiliary and predicted
94/// fields toward consensus. This separates the nonlinear measurement
95/// constraint from the overlapping Fourier-patch consistency constraint.
96///
97/// For incoherently multiplexed data, the amplitude proximal is joint across
98/// all source modes in a frame. Auxiliary and scaled-dual fields are stored per
99/// frame-source mode in [`ReconstructionState`] for exact checkpoint resumption.
100/// The fixed-pupil linearization and multiplexed proximal used here are crate
101/// adaptations of the reference ADMM-FPM formulation.
102///
103/// Each iteration reports detector-field RMS residuals. The primal residual is
104/// `r_k = A x_k - z_k`, evaluated after the linearized object update, and the
105/// dual residual is `s_k = rho (z_k - z_{k-1})`. Here `A x` denotes the
106/// concatenated per-mode detector fields, including each mode of a multiplexed
107/// frame. Masked pixels and zero-weight frames are excluded. These residuals
108/// diagnose consensus and auxiliary-field motion; the solver does not use them
109/// as stopping criteria.
110///
111/// # Reference
112///
113/// [A. Wang, Z. Zhang, S. Wang, A. Pan, C. Ma, and B. Yao, “Fourier
114/// Ptychographic Microscopy via Alternating Direction Method of Multipliers”
115/// (2022)](https://doi.org/10.3390/cells11091512), *Cells* **11**(9), 1512.
116#[derive(Clone, Debug)]
117pub struct Admm {
118    /// Number of complete passes through the acquisition schedule.
119    pub iterations: usize,
120    /// Step size of the linearized, pupil-preconditioned object update.
121    pub object_step: f64,
122    /// Positive augmented-Lagrangian penalty tying auxiliary fields to the
123    /// fields predicted by the shared object.
124    pub penalty: f64,
125    /// Scaled-dual update relaxation in the inclusive range `0..=2`.
126    pub dual_relaxation: f64,
127    /// Number of measured frames supplied to each reconstruction step; the
128    /// default processes every frame together.
129    pub batch_size: usize,
130    /// Positive numerical floor used in normalizations and dark-field handling.
131    pub epsilon: f64,
132}
133
134impl Default for Admm {
135    fn default() -> Self {
136        Self {
137            iterations: 100,
138            object_step: 0.8,
139            penalty: 1.0,
140            dual_relaxation: 1.0,
141            batch_size: usize::MAX,
142            epsilon: 1e-10,
143        }
144    }
145}
146
147impl Admm {
148    /// Sets the number of complete acquisition-schedule passes; validation requires non-zero.
149    pub fn iterations(mut self, iterations: usize) -> Self {
150        self.iterations = iterations;
151        self
152    }
153
154    /// Sets the finite positive step size for the linearized object update.
155    pub fn object_step(mut self, step: f64) -> Self {
156        self.object_step = step;
157        self
158    }
159
160    /// Sets the finite positive augmented-Lagrangian consensus penalty.
161    pub fn penalty(mut self, penalty: f64) -> Self {
162        self.penalty = penalty;
163        self
164    }
165
166    /// Sets finite scaled-dual relaxation in the inclusive interval `[0, 2]`.
167    pub fn dual_relaxation(mut self, relaxation: f64) -> Self {
168        self.dual_relaxation = relaxation;
169        self
170    }
171
172    /// Sets the positive acquisition-frame batch size.
173    pub fn batch_size(mut self, batch_size: usize) -> Self {
174        self.batch_size = batch_size;
175        self
176    }
177}
178
179impl ReconstructionAlgorithm for Admm {
180    type IterationMetrics = AdmmIterationMetrics;
181
182    fn validate(&self) -> Result<()> {
183        if !self.object_step.is_finite() || self.object_step <= 0.0 {
184            return Err(Error::InvalidParameter {
185                name: "object_step",
186                reason: "must be finite and positive".into(),
187            });
188        }
189        if !self.penalty.is_finite() || self.penalty <= 0.0 {
190            return Err(Error::InvalidParameter {
191                name: "penalty",
192                reason: "must be finite and positive".into(),
193            });
194        }
195        if !self.dual_relaxation.is_finite() || !(0.0..=2.0).contains(&self.dual_relaxation) {
196            return Err(Error::InvalidParameter {
197                name: "dual_relaxation",
198                reason: "must be finite and between zero and two".into(),
199            });
200        }
201        if self.batch_size == 0 || !self.epsilon.is_finite() || self.epsilon <= 0.0 {
202            return Err(Error::InvalidParameter {
203                name: "batch_size/epsilon",
204                reason: "batch size must be non-zero and epsilon finite and positive".into(),
205            });
206        }
207        Ok(())
208    }
209
210    fn step<M: MeasurementRead>(
211        &mut self,
212        problem: &ReconstructionProblem<M>,
213        state: &mut ReconstructionState,
214        batch: &Batch,
215        _iteration: usize,
216    ) -> Result<StepOutput<Self::IterationMetrics>> {
217        let expected = admm_auxiliary_len(&problem.model)?;
218        let mut auxiliary = match state.algorithm_auxiliary.take() {
219            None => AdmmAuxiliaryState {
220                auxiliary_fields: vec![Complex64::default(); expected],
221                dual_fields: vec![Complex64::default(); expected],
222            },
223            Some(AlgorithmAuxiliaryState::Admm(auxiliary)) => auxiliary,
224            Some(_) => {
225                return Err(Error::InvalidModel(
226                    "ADMM cannot resume auxiliary state owned by another algorithm".into(),
227                ));
228            }
229        };
230        let result = self.step_with_auxiliary(problem, state, batch, &mut auxiliary);
231        state.algorithm_auxiliary = Some(AlgorithmAuxiliaryState::Admm(auxiliary));
232        result
233    }
234
235    fn iterations(&self) -> usize {
236        self.iterations
237    }
238
239    fn batch_size(&self) -> usize {
240        self.batch_size
241    }
242}
243
244impl Admm {
245    fn step_with_auxiliary<M: MeasurementRead>(
246        &self,
247        problem: &ReconstructionProblem<M>,
248        state: &mut ReconstructionState,
249        batch: &Batch,
250        auxiliary: &mut AdmmAuxiliaryState,
251    ) -> Result<StepOutput<AdmmIterationMetrics>> {
252        let model = &problem.model;
253        let shape = model.image_shape;
254        let image_len = checked_len_2d(shape)?;
255        let expected = admm_auxiliary_len(model)?;
256        if auxiliary.auxiliary_fields.len() != expected || auxiliary.dual_fields.len() != expected {
257            return Err(Error::InvalidModel(
258                "ADMM auxiliary state does not match the problem modes".into(),
259            ));
260        }
261        state
262            .scratch
263            .object_gradient
264            .resize(state.object_spectrum.len(), Complex64::default());
265        state.scratch.object_gradient.fill(Complex64::default());
266        let maximum_pupil_power = state
267            .pupil
268            .values
269            .as_slice()
270            .iter()
271            .map(|value| value.norm_sqr())
272            .fold(0.0, f64::max)
273            .max(self.epsilon);
274        let mut summary = StepSummary::default();
275        let mut metrics = AdmmIterationMetrics::default();
276        let mut active_frames = 0;
277
278        for &frame in &batch.indices {
279            let frame_weight = problem.measurements.frame_weight(frame)?;
280            if frame_weight == 0.0 {
281                summary.push_frame(frame, 0.0, 0.0);
282                continue;
283            }
284            active_frames += 1;
285            let single_source = [(frame, 1.0)];
286            let sources = frame_sources(model, frame, &single_source);
287            let source_weight_sum: f64 = sources.iter().map(|&(_, weight)| weight).sum();
288            let mode_start = frame_mode_start(model, frame);
289            let multiplex_len =
290                image_len
291                    .checked_mul(sources.len())
292                    .ok_or_else(|| Error::ShapeOverflow {
293                        shape: vec![sources.len(), shape.0, shape.1],
294                    })?;
295            state
296                .scratch
297                .multiplex_fields
298                .resize(multiplex_len, Complex64::default());
299            state.scratch.multiplex_offsets.clear();
300            state.scratch.projected_field.fill(Complex64::default());
301
302            // The real component accumulates the physical prediction; the
303            // imaginary component accumulates the dual-shifted proximal norm.
304            for (local_mode, &(source, source_weight)) in sources.iter().enumerate() {
305                let offset = state.effective_source_offset(model, source)?;
306                state.scratch.multiplex_offsets.push(offset);
307                compute_source_field(problem, state, source, offset)?;
308                let local_start = local_mode * image_len;
309                state.scratch.multiplex_fields[local_start..local_start + image_len]
310                    .copy_from_slice(&state.scratch.field);
311                let auxiliary_start = (mode_start + local_mode) * image_len;
312                for pixel in 0..image_len {
313                    let field = state.scratch.field[pixel];
314                    let consensus = field + auxiliary.dual_fields[auxiliary_start + pixel];
315                    state.scratch.projected_field[pixel].re += source_weight * field.norm_sqr();
316                    state.scratch.projected_field[pixel].im += source_weight * consensus.norm_sqr();
317                }
318            }
319
320            let measured = problem.measurements.frame(frame)?;
321            let mask = problem.measurements.frame_mask(frame)?;
322            let gain = state
323                .frame_gains
324                .as_ref()
325                .map_or(1.0, |values| values[frame]);
326            if !gain.is_finite() || gain <= 0.0 {
327                return Err(Error::InvalidModel(format!(
328                    "state frame {frame} has invalid gain {gain}"
329                )));
330            }
331            let mut frame_loss = 0.0;
332            let mut valid_pixels = 0;
333            for pixel in 0..image_len {
334                if mask.is_some_and(|values| values[pixel] == 0) {
335                    continue;
336                }
337                valid_pixels += 1;
338                let background = background_value(state, frame, pixel, image_len);
339                let predicted = gain * state.scratch.projected_field[pixel].re + background;
340                frame_loss += point_loss(predicted, measured[pixel], LossType::AmplitudeMse);
341            }
342            if valid_pixels == 0 {
343                return Err(Error::InvalidMeasurements(format!(
344                    "frame {frame} has no unmasked pixels"
345                )));
346            }
347            summary.push_frame(frame, frame_loss / valid_pixels as f64, frame_weight);
348
349            for (local_mode, &(source, source_weight)) in sources.iter().enumerate() {
350                let local_start = local_mode * image_len;
351                state.scratch.field.copy_from_slice(
352                    &state.scratch.multiplex_fields[local_start..local_start + image_len],
353                );
354                let auxiliary_start = (mode_start + local_mode) * image_len;
355                for pixel in 0..image_len {
356                    let index = auxiliary_start + pixel;
357                    if mask.is_some_and(|values| values[pixel] == 0) {
358                        auxiliary.auxiliary_fields[index] = state.scratch.field[pixel];
359                        auxiliary.dual_fields[index] = Complex64::default();
360                        state.scratch.difference[pixel] = Complex64::default();
361                        continue;
362                    }
363                    let background = background_value(state, frame, pixel, image_len);
364                    let target = ((measured[pixel] - background) / gain).max(0.0).sqrt();
365                    let consensus = state.scratch.field[pixel] + auxiliary.dual_fields[index];
366                    let consensus_norm = state.scratch.projected_field[pixel].im;
367                    let projected = if consensus_norm > self.epsilon {
368                        consensus * (target / consensus_norm.sqrt())
369                    } else {
370                        Complex64::new(target / source_weight_sum.sqrt(), 0.0)
371                    };
372                    let auxiliary_value = (self.penalty * consensus + frame_weight * projected)
373                        / (self.penalty + frame_weight);
374                    metrics.push_dual_change(
375                        auxiliary_value - auxiliary.auxiliary_fields[index],
376                        self.penalty,
377                    );
378                    auxiliary.auxiliary_fields[index] = auxiliary_value;
379                    state.scratch.difference[pixel] =
380                        auxiliary_value - auxiliary.dual_fields[index] - state.scratch.field[pixel];
381                }
382                state.backend.fft2(
383                    &mut state.scratch.difference,
384                    shape,
385                    FftDirection::Forward,
386                    &mut state.scratch.column,
387                )?;
388                fftshift_copy(
389                    &state.scratch.difference,
390                    &mut state.scratch.projected_spectrum,
391                    shape,
392                );
393                for pixel in 0..image_len {
394                    state.scratch.difference[pixel] = state.pupil.values.as_slice()[pixel].conj()
395                        * state.scratch.projected_spectrum[pixel]
396                        / (maximum_pupil_power + self.epsilon);
397                }
398                model.insert_patch_adjoint_slice_at_offset(
399                    &mut state.scratch.object_gradient,
400                    source,
401                    &state.scratch.difference,
402                    source_weight / source_weight_sum,
403                    state.scratch.multiplex_offsets[local_mode],
404                )?;
405            }
406        }
407
408        if active_frames > 0 {
409            let step = self.object_step / active_frames as f64;
410            for (object, &update) in state
411                .object_spectrum
412                .as_slice_mut()
413                .iter_mut()
414                .zip(&state.scratch.object_gradient)
415            {
416                *object += step * update;
417            }
418        }
419
420        for &frame in &batch.indices {
421            if problem.measurements.frame_weight(frame)? == 0.0 {
422                continue;
423            }
424            let single_source = [(frame, 1.0)];
425            let sources = frame_sources(model, frame, &single_source);
426            let mode_start = frame_mode_start(model, frame);
427            let mask = problem.measurements.frame_mask(frame)?;
428            for (local_mode, &(source, _)) in sources.iter().enumerate() {
429                let offset = state.effective_source_offset(model, source)?;
430                compute_source_field(problem, state, source, offset)?;
431                let auxiliary_start = (mode_start + local_mode) * image_len;
432                for pixel in 0..image_len {
433                    let index = auxiliary_start + pixel;
434                    if mask.is_some_and(|values| values[pixel] == 0) {
435                        auxiliary.auxiliary_fields[index] = state.scratch.field[pixel];
436                        auxiliary.dual_fields[index] = Complex64::default();
437                    } else {
438                        let primal_residual =
439                            state.scratch.field[pixel] - auxiliary.auxiliary_fields[index];
440                        metrics.push_primal_residual(primal_residual);
441                        auxiliary.dual_fields[index] += self.dual_relaxation * primal_residual;
442                    }
443                }
444            }
445        }
446        Ok(StepOutput { summary, metrics })
447    }
448}
449
450pub(crate) fn admm_auxiliary_len(model: &ImagePlaneModel) -> Result<usize> {
451    let mode_count = model.multiplexing_matrix.as_ref().map_or_else(
452        || Ok(model.frame_count()),
453        |matrix| {
454            matrix.iter().try_fold(0_usize, |count, row| {
455                count
456                    .checked_add(row.len())
457                    .ok_or_else(|| Error::InvalidShape("ADMM source mode count overflows".into()))
458            })
459        },
460    )?;
461    checked_len_2d(model.image_shape)?
462        .checked_mul(mode_count)
463        .ok_or_else(|| Error::InvalidShape("ADMM auxiliary length overflows".into()))
464}
465
466fn frame_mode_start(model: &ImagePlaneModel, frame: usize) -> usize {
467    model
468        .multiplexing_matrix
469        .as_ref()
470        .map_or(frame, |matrix| matrix[..frame].iter().map(Vec::len).sum())
471}
472
473fn frame_sources<'a>(
474    model: &'a ImagePlaneModel,
475    frame: usize,
476    single_source: &'a [(usize, f64); 1],
477) -> &'a [(usize, f64)] {
478    model
479        .multiplexing_matrix
480        .as_ref()
481        .map_or(single_source, |matrix| matrix[frame].as_slice())
482}
483
484fn compute_source_field<M: MeasurementRead>(
485    problem: &ReconstructionProblem<M>,
486    state: &mut ReconstructionState,
487    source: usize,
488    offset: FourierOffset,
489) -> Result<()> {
490    let model = &problem.model;
491    let shape = model.image_shape;
492    model.extract_patch_at_offset(
493        state.object_spectrum.view(),
494        source,
495        offset,
496        &mut state.scratch.patch,
497    )?;
498    for pixel in 0..state.scratch.patch.len() {
499        state.scratch.exit_spectrum[pixel] =
500            state.scratch.patch[pixel] * state.pupil.values.as_slice()[pixel];
501    }
502    ifftshift_copy(
503        &state.scratch.exit_spectrum,
504        &mut state.scratch.field,
505        shape,
506    );
507    state.backend.fft2(
508        &mut state.scratch.field,
509        shape,
510        FftDirection::Inverse,
511        &mut state.scratch.column,
512    )
513}
514
515fn background_value(
516    state: &ReconstructionState,
517    frame: usize,
518    pixel: usize,
519    image_len: usize,
520) -> f64 {
521    state.background.as_ref().map_or(0.0, |values| {
522        values[if values.len() == image_len {
523            pixel
524        } else {
525            frame * image_len + pixel
526        }]
527    })
528}