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#[derive(Clone, Debug)]
58pub struct Mpie {
59 pub iterations: usize,
61 pub object_step: f64,
63 pub stability: f64,
66 pub momentum_interval: usize,
68 pub momentum_friction: f64,
70 pub momentum_feedback: f64,
72 pub batch_size: usize,
74 pub epsilon: f64,
76 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 pub fn iterations(mut self, iterations: usize) -> Self {
100 self.iterations = iterations;
101 self
102 }
103
104 pub fn object_step(mut self, object_step: f64) -> Self {
106 self.object_step = object_step;
107 self
108 }
109
110 pub fn stability(mut self, stability: f64) -> Self {
112 self.stability = stability;
113 self
114 }
115
116 pub fn momentum_interval(mut self, momentum_interval: usize) -> Self {
118 self.momentum_interval = momentum_interval;
119 self
120 }
121
122 pub fn momentum_friction(mut self, momentum_friction: f64) -> Self {
124 self.momentum_friction = momentum_friction;
125 self
126 }
127
128 pub fn momentum_feedback(mut self, momentum_feedback: f64) -> Self {
130 self.momentum_feedback = momentum_feedback;
131 self
132 }
133
134 pub fn batch_size(mut self, batch_size: usize) -> Self {
136 self.batch_size = batch_size;
137 self
138 }
139
140 pub fn epsilon(mut self, epsilon: f64) -> Self {
142 self.epsilon = epsilon;
143 self
144 }
145
146 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}