Skip to main content

fpm_rs/algorithms/
mpie.rs

1use std::sync::Arc;
2
3use num_complex::Complex64;
4
5use crate::{
6    Result,
7    algorithms::{NoIterationMetrics, StepOutput, objective::LossType},
8    backend::Backend,
9    error::Error,
10    measurements::MeasurementRead,
11    reconstruction::{
12        AlgorithmAuxiliaryState, Batch, MpieAuxiliaryState, ReconstructionProblem,
13        ReconstructionState,
14    },
15};
16
17use super::{
18    ReconstructionAlgorithm,
19    common::{MomentumConfiguration, ObjectDenominator, UpdateConfiguration, projection_update},
20};
21
22/// Momentum-accelerated regularized PIE adapted to image-plane FPM.
23///
24/// # Method
25///
26/// `Mpie` applies the same sequential, object-only rPIE projection as
27/// [`super::Fpie`], then periodically accelerates the centered high-resolution
28/// object spectrum. After `momentum_interval` positive-weight measured-frame
29/// updates, it updates the complex velocity and object as
30///
31/// `V <- momentum_friction * V + (O_rpie - O_anchor)`
32///
33/// `O <- O_rpie + momentum_feedback * V`.
34///
35/// A multiplexed measurement counts once after all of its source modes are
36/// inserted. Zero-weight frames do not advance the interval. The counter,
37/// velocity, anchor, and defining parameters are checkpointed, so changing
38/// `batch_size` does not change the numerical path and a matching checkpoint
39/// resumes exactly.
40///
41/// # Adaptation and assumptions
42///
43/// The cited method was formulated and tested for scanned ptychography and
44/// applies momentum to both object and probe. This implementation accelerates
45/// only the fixed-pupil Fourier-ptychographic object spectrum. It separates the
46/// paper's single `eta_obj` into friction and feedback controls; setting them
47/// equal reproduces its Eqs. (19) and (21). `object_step` corresponds to the
48/// paper's `gamma_obj` in Eq. (22). Physical joint calibration is unsupported
49/// because recompiling the model would require an explicit rule for resetting
50/// or transporting momentum.
51///
52/// # Reference
53///
54/// [A. Maiden, D. Johnson, and P. Li, “Further improvements to the
55/// ptychographical iterative engine” (2017)](https://doi.org/10.1364/OPTICA.4.000736),
56/// *Optica* **4**(7), 736–745.
57#[derive(Clone, Debug)]
58pub struct Mpie {
59    /// Number of complete passes through the acquisition schedule.
60    pub iterations: usize,
61    /// Relaxation factor applied to each rPIE object-spectrum correction.
62    pub object_step: f64,
63    /// Blend between local pupil power (`0`) and maximum pupil power (`1`) in
64    /// the rPIE denominator.
65    pub stability: f64,
66    /// Positive-weight measured-frame updates between momentum events.
67    pub momentum_interval: usize,
68    /// Fraction of the previous velocity retained at each momentum event.
69    pub momentum_friction: f64,
70    /// Fraction of the updated velocity added to the object spectrum.
71    pub momentum_feedback: f64,
72    /// Number of measured frames supplied to each reconstruction step.
73    pub batch_size: usize,
74    /// Positive numerical floor added to the rPIE denominator.
75    pub epsilon: f64,
76    /// Loss used for diagnostics; the projection itself always enforces the
77    /// measured amplitude.
78    pub loss_type: LossType,
79}
80
81impl Default for Mpie {
82    fn default() -> Self {
83        Self {
84            iterations: 50,
85            object_step: 0.2,
86            stability: 0.05,
87            momentum_interval: 30,
88            momentum_friction: 0.9,
89            momentum_feedback: 0.9,
90            batch_size: 1,
91            epsilon: 1e-10,
92            loss_type: LossType::AmplitudeMse,
93        }
94    }
95}
96
97impl Mpie {
98    /// Sets the number of complete acquisition-schedule passes.
99    pub fn iterations(mut self, iterations: usize) -> Self {
100        self.iterations = iterations;
101        self
102    }
103
104    /// Sets the finite positive relaxation applied to rPIE corrections.
105    pub fn object_step(mut self, object_step: f64) -> Self {
106        self.object_step = object_step;
107        self
108    }
109
110    /// Sets the rPIE pupil-power blend; validation requires `[0, 1]`.
111    pub fn stability(mut self, stability: f64) -> Self {
112        self.stability = stability;
113        self
114    }
115
116    /// Sets the positive number of effective frame updates per momentum event.
117    pub fn momentum_interval(mut self, momentum_interval: usize) -> Self {
118        self.momentum_interval = momentum_interval;
119        self
120    }
121
122    /// Sets the retained-velocity fraction; validation requires `[0, 1)`.
123    pub fn momentum_friction(mut self, momentum_friction: f64) -> Self {
124        self.momentum_friction = momentum_friction;
125        self
126    }
127
128    /// Sets the velocity-feedback fraction; validation requires `[0, 1]`.
129    pub fn momentum_feedback(mut self, momentum_feedback: f64) -> Self {
130        self.momentum_feedback = momentum_feedback;
131        self
132    }
133
134    /// Sets the positive number of acquisition frames supplied per step.
135    pub fn batch_size(mut self, batch_size: usize) -> Self {
136        self.batch_size = batch_size;
137        self
138    }
139
140    /// Sets the finite positive numerical floor used by the rPIE denominator.
141    pub fn epsilon(mut self, epsilon: f64) -> Self {
142        self.epsilon = epsilon;
143        self
144    }
145
146    /// Sets the loss reported by diagnostics.
147    pub fn loss_type(mut self, loss_type: LossType) -> Self {
148        self.loss_type = loss_type;
149        self
150    }
151
152    fn auxiliary_from_state(&self, state: &ReconstructionState) -> MpieAuxiliaryState {
153        MpieAuxiliaryState {
154            velocity: vec![Complex64::default(); state.object_spectrum.len()],
155            anchor: state.object_spectrum.as_slice().to_vec(),
156            effective_frames_since_momentum: 0,
157            object_step: self.object_step,
158            stability: self.stability,
159            epsilon: self.epsilon,
160            momentum_interval: self.momentum_interval,
161            momentum_friction: self.momentum_friction,
162            momentum_feedback: self.momentum_feedback,
163        }
164    }
165
166    fn prepare_auxiliary(&self, state: &mut ReconstructionState) -> Result<()> {
167        if state.algorithm_auxiliary.is_none() {
168            state.algorithm_auxiliary = Some(AlgorithmAuxiliaryState::Mpie(
169                self.auxiliary_from_state(state),
170            ));
171            return Ok(());
172        }
173        let auxiliary = match state.algorithm_auxiliary.as_ref() {
174            Some(AlgorithmAuxiliaryState::Mpie(auxiliary)) => auxiliary,
175            Some(_) => {
176                return Err(Error::InvalidModel(
177                    "mPIE cannot resume auxiliary state owned by another algorithm".into(),
178                ));
179            }
180            None => unreachable!("missing state was initialized above"),
181        };
182        if auxiliary.velocity.len() != state.object_spectrum.len()
183            || auxiliary.anchor.len() != state.object_spectrum.len()
184        {
185            return Err(Error::InvalidModel(
186                "mPIE auxiliary dimensions do not match the object spectrum".into(),
187            ));
188        }
189        for (name, current, stored) in [
190            ("object_step", self.object_step, auxiliary.object_step),
191            ("stability", self.stability, auxiliary.stability),
192            ("epsilon", self.epsilon, auxiliary.epsilon),
193            (
194                "momentum_friction",
195                self.momentum_friction,
196                auxiliary.momentum_friction,
197            ),
198            (
199                "momentum_feedback",
200                self.momentum_feedback,
201                auxiliary.momentum_feedback,
202            ),
203        ] {
204            if current.to_bits() != stored.to_bits() {
205                return Err(Error::InvalidParameter {
206                    name,
207                    reason: format!(
208                        "value {current} differs from checkpointed mPIE value {stored}"
209                    ),
210                });
211            }
212        }
213        if self.momentum_interval != auxiliary.momentum_interval {
214            return Err(Error::InvalidParameter {
215                name: "momentum_interval",
216                reason: format!(
217                    "value {} differs from checkpointed mPIE value {}",
218                    self.momentum_interval, auxiliary.momentum_interval
219                ),
220            });
221        }
222        Ok(())
223    }
224}
225
226impl ReconstructionAlgorithm for Mpie {
227    type IterationMetrics = NoIterationMetrics;
228
229    fn validate(&self) -> Result<()> {
230        if !self.object_step.is_finite() || self.object_step <= 0.0 {
231            return Err(Error::InvalidParameter {
232                name: "object_step",
233                reason: "must be finite and positive".into(),
234            });
235        }
236        if !self.stability.is_finite() || !(0.0..=1.0).contains(&self.stability) {
237            return Err(Error::InvalidParameter {
238                name: "stability",
239                reason: "must be finite and between zero and one".into(),
240            });
241        }
242        if self.momentum_interval == 0 {
243            return Err(Error::InvalidParameter {
244                name: "momentum_interval",
245                reason: "must be greater than zero".into(),
246            });
247        }
248        if !self.momentum_friction.is_finite() || !(0.0..1.0).contains(&self.momentum_friction) {
249            return Err(Error::InvalidParameter {
250                name: "momentum_friction",
251                reason: "must be finite, at least zero, and less than one".into(),
252            });
253        }
254        if !self.momentum_feedback.is_finite() || !(0.0..=1.0).contains(&self.momentum_feedback) {
255            return Err(Error::InvalidParameter {
256                name: "momentum_feedback",
257                reason: "must be finite and between zero and one".into(),
258            });
259        }
260        if self.batch_size == 0 {
261            return Err(Error::InvalidParameter {
262                name: "batch_size",
263                reason: "must be greater than zero".into(),
264            });
265        }
266        if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
267            return Err(Error::InvalidParameter {
268                name: "epsilon",
269                reason: "must be finite and positive".into(),
270            });
271        }
272        Ok(())
273    }
274
275    fn initialize<M: MeasurementRead>(
276        &self,
277        problem: &ReconstructionProblem<M>,
278    ) -> Result<ReconstructionState> {
279        let mut state = ReconstructionState::initialize(problem)?;
280        self.prepare_auxiliary(&mut state)?;
281        Ok(state)
282    }
283
284    fn initialize_with_backend<M: MeasurementRead>(
285        &self,
286        problem: &ReconstructionProblem<M>,
287        backend: Arc<dyn Backend>,
288    ) -> Result<ReconstructionState> {
289        let mut state = ReconstructionState::initialize_with_backend(problem, backend)?;
290        self.prepare_auxiliary(&mut state)?;
291        Ok(state)
292    }
293
294    fn supports_joint_reconstruction(&self) -> bool {
295        false
296    }
297
298    fn step<M: MeasurementRead>(
299        &mut self,
300        problem: &ReconstructionProblem<M>,
301        state: &mut ReconstructionState,
302        batch: &Batch,
303        _iteration: usize,
304    ) -> Result<StepOutput<Self::IterationMetrics>> {
305        self.prepare_auxiliary(state)?;
306        Ok(projection_update(
307            problem,
308            state,
309            batch,
310            UpdateConfiguration {
311                object_step: self.object_step,
312                pupil_step: None,
313                epsilon: self.epsilon,
314                loss_type: self.loss_type,
315                object_denominator: ObjectDenominator::Rpie(self.stability),
316                constrain_pupil: true,
317                gain_update: None,
318                background_update: None,
319                momentum: Some(MomentumConfiguration {
320                    interval: self.momentum_interval,
321                    friction: self.momentum_friction,
322                    feedback: self.momentum_feedback,
323                }),
324            },
325        )?
326        .into())
327    }
328
329    fn iterations(&self) -> usize {
330        self.iterations
331    }
332
333    fn batch_size(&self) -> usize {
334        self.batch_size
335    }
336}