1use crate::{
2 Result,
3 algorithms::{NoIterationMetrics, StepOutput, objective::LossType},
4 error::Error,
5 measurements::MeasurementRead,
6 reconstruction::{Batch, ReconstructionProblem, ReconstructionState},
7};
8
9use super::{
10 ReconstructionAlgorithm,
11 common::{ObjectDenominator, UpdateConfiguration, projection_update},
12 gauge::canonicalize_object_pupil,
13};
14
15#[derive(Clone, Debug)]
55pub struct Epry {
56 pub iterations: usize,
58 pub object_step: f64,
60 pub pupil_step: f64,
62 pub batch_size: usize,
64 pub recover_pupil: bool,
66 pub constrain_pupil_support: bool,
68 pub recover_frame_gains: bool,
70 pub gain_step: f64,
72 pub minimum_gain: f64,
74 pub maximum_gain: f64,
76 pub recover_background: bool,
78 pub background_step: f64,
80 pub minimum_background: f64,
82 pub maximum_background: f64,
84 pub epsilon: f64,
86 pub loss_type: LossType,
89}
90
91impl Default for Epry {
92 fn default() -> Self {
93 Self {
94 iterations: 100,
95 object_step: 0.8,
96 pupil_step: 0.1,
97 batch_size: 1,
98 recover_pupil: true,
99 constrain_pupil_support: true,
100 recover_frame_gains: false,
101 gain_step: 0.2,
102 minimum_gain: 1e-6,
103 maximum_gain: 1e6,
104 recover_background: false,
105 background_step: 0.2,
106 minimum_background: 0.0,
107 maximum_background: 1e12,
108 epsilon: 1e-10,
109 loss_type: LossType::AmplitudeMse,
110 }
111 }
112}
113
114impl Epry {
115 pub fn iterations(mut self, iterations: usize) -> Self {
117 self.iterations = iterations;
118 self
119 }
120
121 pub fn object_step(mut self, step: f64) -> Self {
123 self.object_step = step;
124 self
125 }
126
127 pub fn pupil_step(mut self, step: f64) -> Self {
129 self.pupil_step = step;
130 self
131 }
132
133 pub fn recover_pupil(mut self, recover: bool) -> Self {
135 self.recover_pupil = recover;
136 self
137 }
138
139 pub fn constrain_pupil_support(mut self, constrain: bool) -> Self {
141 self.constrain_pupil_support = constrain;
142 self
143 }
144
145 pub fn recover_frame_gains(mut self, recover: bool) -> Self {
147 self.recover_frame_gains = recover;
148 self
149 }
150
151 pub fn gain_step(mut self, step: f64) -> Self {
153 self.gain_step = step;
154 self
155 }
156
157 pub fn gain_bounds(mut self, minimum: f64, maximum: f64) -> Self {
159 self.minimum_gain = minimum;
160 self.maximum_gain = maximum;
161 self
162 }
163
164 pub fn recover_background(mut self, recover: bool) -> Self {
166 self.recover_background = recover;
167 self
168 }
169
170 pub fn background_step(mut self, step: f64) -> Self {
172 self.background_step = step;
173 self
174 }
175
176 pub fn background_bounds(mut self, minimum: f64, maximum: f64) -> Self {
178 self.minimum_background = minimum;
179 self.maximum_background = maximum;
180 self
181 }
182}
183
184impl ReconstructionAlgorithm for Epry {
185 type IterationMetrics = NoIterationMetrics;
186
187 fn validate(&self) -> Result<()> {
188 if !self.object_step.is_finite() || self.object_step <= 0.0 {
189 return Err(Error::InvalidParameter {
190 name: "object_step",
191 reason: "must be finite and positive".into(),
192 });
193 }
194 if !self.pupil_step.is_finite() || self.pupil_step < 0.0 {
195 return Err(Error::InvalidParameter {
196 name: "pupil_step",
197 reason: "must be finite and non-negative".into(),
198 });
199 }
200 if !self.epsilon.is_finite() || self.epsilon <= 0.0 || self.batch_size == 0 {
201 return Err(Error::InvalidParameter {
202 name: "epsilon/batch_size",
203 reason: "epsilon must be positive and batch size non-zero".into(),
204 });
205 }
206 if !self.gain_step.is_finite() || !(0.0..=1.0).contains(&self.gain_step) {
207 return Err(Error::InvalidParameter {
208 name: "gain_step",
209 reason: "must be finite and between zero and one".into(),
210 });
211 }
212 if !self.minimum_gain.is_finite()
213 || !self.maximum_gain.is_finite()
214 || self.minimum_gain <= 0.0
215 || self.maximum_gain <= self.minimum_gain
216 {
217 return Err(Error::InvalidParameter {
218 name: "gain_bounds",
219 reason: "must be finite, positive, and strictly increasing".into(),
220 });
221 }
222 if !self.background_step.is_finite() || !(0.0..=1.0).contains(&self.background_step) {
223 return Err(Error::InvalidParameter {
224 name: "background_step",
225 reason: "must be finite and between zero and one".into(),
226 });
227 }
228 if !self.minimum_background.is_finite()
229 || !self.maximum_background.is_finite()
230 || self.maximum_background <= self.minimum_background
231 {
232 return Err(Error::InvalidParameter {
233 name: "background_bounds",
234 reason: "must be finite and strictly increasing".into(),
235 });
236 }
237 Ok(())
238 }
239
240 fn canonicalize_state<M: MeasurementRead>(
241 &self,
242 problem: &ReconstructionProblem<M>,
243 state: &mut ReconstructionState,
244 ) -> Result<()> {
245 if self.recover_pupil {
246 canonicalize_object_pupil(problem, state)?;
247 }
248 Ok(())
249 }
250
251 fn step<M: MeasurementRead>(
252 &mut self,
253 problem: &ReconstructionProblem<M>,
254 state: &mut ReconstructionState,
255 batch: &Batch,
256 _iteration: usize,
257 ) -> Result<StepOutput<Self::IterationMetrics>> {
258 Ok(projection_update(
259 problem,
260 state,
261 batch,
262 UpdateConfiguration {
263 object_step: self.object_step,
264 pupil_step: self.recover_pupil.then_some(self.pupil_step),
265 epsilon: self.epsilon,
266 loss_type: self.loss_type,
267 object_denominator: ObjectDenominator::Global,
268 constrain_pupil: self.constrain_pupil_support,
269 gain_update: self.recover_frame_gains.then_some(
270 super::common::GainUpdateConfiguration {
271 step: self.gain_step,
272 minimum: self.minimum_gain,
273 maximum: self.maximum_gain,
274 },
275 ),
276 background_update: self.recover_background.then_some(
277 super::common::BackgroundUpdateConfiguration {
278 step: self.background_step,
279 minimum: self.minimum_background,
280 maximum: self.maximum_background,
281 },
282 ),
283 momentum: None,
284 },
285 )?
286 .into())
287 }
288
289 fn iterations(&self) -> usize {
290 self.iterations
291 }
292
293 fn batch_size(&self) -> usize {
294 self.batch_size
295 }
296}