1use std::sync::Arc;
2
3use crate::{
4 Result,
5 algorithms::{AlgorithmIterationMetrics, StepOutput, objective::LossType},
6 backend::Backend,
7 error::Error,
8 measurements::MeasurementRead,
9 reconstruction::{
10 AdaptiveAlternatingProjectionAuxiliaryState, AlgorithmAuxiliaryState, Batch,
11 ReconstructionProblem, ReconstructionState,
12 },
13};
14
15use super::{
16 ReconstructionAlgorithm,
17 common::{ObjectDenominator, UpdateConfiguration, projection_update},
18};
19
20#[derive(Clone, Copy, Debug, Default)]
22pub struct AdaptiveAlternatingProjectionIterationMetrics {
23 object_step: Option<f64>,
24}
25
26impl AdaptiveAlternatingProjectionIterationMetrics {
27 pub const fn object_step(&self) -> Option<f64> {
29 self.object_step
30 }
31}
32
33impl AlgorithmIterationMetrics for AdaptiveAlternatingProjectionIterationMetrics {
34 fn merge(&mut self, other: Self) {
35 if self.object_step.is_none() {
36 self.object_step = other.object_step;
37 }
38 }
39
40 fn append_records(
41 &self,
42 iteration: usize,
43 output: &mut Vec<crate::reconstruction::AlgorithmMetricRecord>,
44 ) {
45 if let Some(value) = self.object_step {
46 output.push(crate::reconstruction::AlgorithmMetricRecord {
47 iteration,
48 namespace: "adaptive_alternating_projection".into(),
49 metric: "object_step".into(),
50 value,
51 });
52 }
53 }
54}
55
56#[derive(Clone, Debug)]
95pub struct AdaptiveAlternatingProjection {
96 pub iterations: usize,
98 pub initial_object_step: f64,
100 pub progress_threshold: f64,
102 pub reduction_factor: f64,
104 pub minimum_object_step: f64,
106 pub batch_size: usize,
108 pub epsilon: f64,
110}
111
112impl Default for AdaptiveAlternatingProjection {
113 fn default() -> Self {
114 Self {
115 iterations: 50,
116 initial_object_step: 1.0,
117 progress_threshold: 0.01,
118 reduction_factor: 0.5,
119 minimum_object_step: 0.001,
120 batch_size: 1,
121 epsilon: 1e-10,
122 }
123 }
124}
125
126impl AdaptiveAlternatingProjection {
127 pub fn iterations(mut self, iterations: usize) -> Self {
129 self.iterations = iterations;
130 self
131 }
132
133 pub fn initial_object_step(mut self, initial_object_step: f64) -> Self {
135 self.initial_object_step = initial_object_step;
136 self
137 }
138
139 pub fn progress_threshold(mut self, progress_threshold: f64) -> Self {
141 self.progress_threshold = progress_threshold;
142 self
143 }
144
145 pub fn reduction_factor(mut self, reduction_factor: f64) -> Self {
147 self.reduction_factor = reduction_factor;
148 self
149 }
150
151 pub fn minimum_object_step(mut self, minimum_object_step: f64) -> Self {
153 self.minimum_object_step = minimum_object_step;
154 self
155 }
156
157 pub fn batch_size(mut self, batch_size: usize) -> Self {
159 self.batch_size = batch_size;
160 self
161 }
162
163 pub fn epsilon(mut self, epsilon: f64) -> Self {
165 self.epsilon = epsilon;
166 self
167 }
168
169 fn new_auxiliary(
170 &self,
171 active_iteration: usize,
172 ) -> AdaptiveAlternatingProjectionAuxiliaryState {
173 AdaptiveAlternatingProjectionAuxiliaryState {
174 active_iteration,
175 current_object_step: self.initial_object_step,
176 previous_objective: None,
177 objective_sum: 0.0,
178 weight_sum: 0.0,
179 frames_accumulated: 0,
180 initial_object_step: self.initial_object_step,
181 progress_threshold: self.progress_threshold,
182 reduction_factor: self.reduction_factor,
183 minimum_object_step: self.minimum_object_step,
184 epsilon: self.epsilon,
185 }
186 }
187
188 fn prepare_auxiliary(
189 &self,
190 state: &mut ReconstructionState,
191 active_iteration: usize,
192 ) -> Result<()> {
193 if state.algorithm_auxiliary.is_none() {
194 state.algorithm_auxiliary =
195 Some(AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(
196 self.new_auxiliary(active_iteration),
197 ));
198 return Ok(());
199 }
200 let auxiliary = match state.algorithm_auxiliary.as_ref() {
201 Some(AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(auxiliary)) => auxiliary,
202 Some(_) => {
203 return Err(Error::InvalidModel(
204 "adaptive alternating projection cannot resume auxiliary state owned by another algorithm"
205 .into(),
206 ));
207 }
208 None => unreachable!("missing state was initialized above"),
209 };
210 for (name, current, stored) in [
211 (
212 "initial_object_step",
213 self.initial_object_step,
214 auxiliary.initial_object_step,
215 ),
216 (
217 "progress_threshold",
218 self.progress_threshold,
219 auxiliary.progress_threshold,
220 ),
221 (
222 "reduction_factor",
223 self.reduction_factor,
224 auxiliary.reduction_factor,
225 ),
226 (
227 "minimum_object_step",
228 self.minimum_object_step,
229 auxiliary.minimum_object_step,
230 ),
231 ("epsilon", self.epsilon, auxiliary.epsilon),
232 ] {
233 if current.to_bits() != stored.to_bits() {
234 return Err(Error::InvalidParameter {
235 name,
236 reason: format!(
237 "value {current} differs from checkpointed adaptive-projection value {stored}"
238 ),
239 });
240 }
241 }
242 Ok(())
243 }
244
245 fn begin_iteration<M: MeasurementRead>(
246 &self,
247 problem: &ReconstructionProblem<M>,
248 state: &mut ReconstructionState,
249 iteration: usize,
250 ) -> Result<f64> {
251 self.prepare_auxiliary(state, iteration)?;
252 let auxiliary = match state.algorithm_auxiliary.as_mut() {
253 Some(AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(auxiliary)) => auxiliary,
254 _ => unreachable!("adaptive auxiliary was prepared above"),
255 };
256 if iteration == auxiliary.active_iteration {
257 return Ok(auxiliary.current_object_step);
258 }
259 if iteration != auxiliary.active_iteration.saturating_add(1) {
260 return Err(Error::InvalidModel(format!(
261 "adaptive alternating projection expected iteration {} or {}, got {iteration}",
262 auxiliary.active_iteration,
263 auxiliary.active_iteration.saturating_add(1)
264 )));
265 }
266 if auxiliary.frames_accumulated != problem.model.frame_count() {
267 return Err(Error::InvalidModel(format!(
268 "adaptive alternating projection accumulated {} of {} frames in iteration {}",
269 auxiliary.frames_accumulated,
270 problem.model.frame_count(),
271 auxiliary.active_iteration
272 )));
273 }
274 if !auxiliary.objective_sum.is_finite()
275 || !auxiliary.weight_sum.is_finite()
276 || auxiliary.weight_sum <= 0.0
277 {
278 return Err(Error::Numerical(
279 "adaptive alternating projection has a non-finite or empty pass objective".into(),
280 ));
281 }
282 let objective = auxiliary.objective_sum / auxiliary.weight_sum;
283 if !objective.is_finite() || objective < 0.0 {
284 return Err(Error::Numerical(
285 "adaptive alternating projection produced an invalid pass objective".into(),
286 ));
287 }
288 if let Some(previous) = auxiliary.previous_objective {
289 let relative_progress = (previous - objective) / previous.max(auxiliary.epsilon);
290 if !relative_progress.is_finite() {
291 return Err(Error::Numerical(
292 "adaptive alternating projection produced non-finite relative progress".into(),
293 ));
294 }
295 if relative_progress <= auxiliary.progress_threshold {
296 auxiliary.current_object_step = (auxiliary.current_object_step
297 * auxiliary.reduction_factor)
298 .max(auxiliary.minimum_object_step);
299 }
300 }
301 auxiliary.previous_objective = Some(objective);
302 auxiliary.objective_sum = 0.0;
303 auxiliary.weight_sum = 0.0;
304 auxiliary.frames_accumulated = 0;
305 auxiliary.active_iteration = iteration;
306 Ok(auxiliary.current_object_step)
307 }
308
309 fn accumulate_batch(
310 &self,
311 state: &mut ReconstructionState,
312 output: &StepOutput<AdaptiveAlternatingProjectionIterationMetrics>,
313 ) -> Result<()> {
314 if !output.summary.objective_sum.is_finite()
315 || !output.summary.weight_sum.is_finite()
316 || output.summary.weight_sum < 0.0
317 {
318 return Err(Error::Numerical(
319 "adaptive alternating projection produced a non-finite batch objective".into(),
320 ));
321 }
322 let auxiliary = match state.algorithm_auxiliary.as_mut() {
323 Some(AlgorithmAuxiliaryState::AdaptiveAlternatingProjection(auxiliary)) => auxiliary,
324 _ => {
325 return Err(Error::InvalidModel(
326 "adaptive alternating projection auxiliary state is missing".into(),
327 ));
328 }
329 };
330 auxiliary.objective_sum += output.summary.objective_sum;
331 auxiliary.weight_sum += output.summary.weight_sum;
332 auxiliary.frames_accumulated = auxiliary
333 .frames_accumulated
334 .checked_add(output.summary.frame_count)
335 .ok_or_else(|| {
336 Error::Numerical("adaptive alternating projection frame counter overflowed".into())
337 })?;
338 if !auxiliary.objective_sum.is_finite() || !auxiliary.weight_sum.is_finite() {
339 return Err(Error::Numerical(
340 "adaptive alternating projection pass objective overflowed".into(),
341 ));
342 }
343 Ok(())
344 }
345}
346
347impl ReconstructionAlgorithm for AdaptiveAlternatingProjection {
348 type IterationMetrics = AdaptiveAlternatingProjectionIterationMetrics;
349
350 fn validate(&self) -> Result<()> {
351 if self.iterations == 0 {
352 return Err(Error::InvalidParameter {
353 name: "iterations",
354 reason: "must be greater than zero".into(),
355 });
356 }
357 if !self.initial_object_step.is_finite() || self.initial_object_step <= 0.0 {
358 return Err(Error::InvalidParameter {
359 name: "initial_object_step",
360 reason: "must be finite and positive".into(),
361 });
362 }
363 if !self.progress_threshold.is_finite() || !(0.0..1.0).contains(&self.progress_threshold) {
364 return Err(Error::InvalidParameter {
365 name: "progress_threshold",
366 reason: "must be finite, at least zero, and less than one".into(),
367 });
368 }
369 if !self.reduction_factor.is_finite()
370 || !(0.0..1.0).contains(&self.reduction_factor)
371 || self.reduction_factor == 0.0
372 {
373 return Err(Error::InvalidParameter {
374 name: "reduction_factor",
375 reason: "must be finite, greater than zero, and less than one".into(),
376 });
377 }
378 if !self.minimum_object_step.is_finite() || self.minimum_object_step <= 0.0 {
379 return Err(Error::InvalidParameter {
380 name: "minimum_object_step",
381 reason: "must be finite and positive".into(),
382 });
383 }
384 if self.minimum_object_step > self.initial_object_step {
385 return Err(Error::InvalidParameter {
386 name: "minimum_object_step",
387 reason: "must not exceed initial_object_step".into(),
388 });
389 }
390 if self.batch_size == 0 {
391 return Err(Error::InvalidParameter {
392 name: "batch_size",
393 reason: "must be greater than zero".into(),
394 });
395 }
396 if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
397 return Err(Error::InvalidParameter {
398 name: "epsilon",
399 reason: "must be finite and positive".into(),
400 });
401 }
402 Ok(())
403 }
404
405 fn initialize<M: MeasurementRead>(
406 &self,
407 problem: &ReconstructionProblem<M>,
408 ) -> Result<ReconstructionState> {
409 let mut state = ReconstructionState::initialize(problem)?;
410 self.prepare_auxiliary(&mut state, 0)?;
411 Ok(state)
412 }
413
414 fn initialize_with_backend<M: MeasurementRead>(
415 &self,
416 problem: &ReconstructionProblem<M>,
417 backend: Arc<dyn Backend>,
418 ) -> Result<ReconstructionState> {
419 let mut state = ReconstructionState::initialize_with_backend(problem, backend)?;
420 self.prepare_auxiliary(&mut state, 0)?;
421 Ok(state)
422 }
423
424 fn supports_joint_reconstruction(&self) -> bool {
425 false
426 }
427
428 fn step<M: MeasurementRead>(
429 &mut self,
430 problem: &ReconstructionProblem<M>,
431 state: &mut ReconstructionState,
432 batch: &Batch,
433 iteration: usize,
434 ) -> Result<StepOutput<Self::IterationMetrics>> {
435 let object_step = self.begin_iteration(problem, state, iteration)?;
436 let summary = projection_update(
437 problem,
438 state,
439 batch,
440 UpdateConfiguration {
441 object_step,
442 pupil_step: None,
443 epsilon: self.epsilon,
444 loss_type: LossType::AmplitudeMse,
445 object_denominator: ObjectDenominator::Local,
446 constrain_pupil: true,
447 gain_update: None,
448 background_update: None,
449 momentum: None,
450 },
451 )?;
452 if state
453 .object_spectrum()
454 .iter()
455 .any(|value| !value.re.is_finite() || !value.im.is_finite())
456 {
457 return Err(Error::Numerical(
458 "adaptive alternating projection produced a non-finite object spectrum".into(),
459 ));
460 }
461 let output = StepOutput {
462 summary,
463 metrics: AdaptiveAlternatingProjectionIterationMetrics {
464 object_step: Some(object_step),
465 },
466 };
467 self.accumulate_batch(state, &output)?;
468 Ok(output)
469 }
470
471 fn iterations(&self) -> usize {
472 self.iterations
473 }
474
475 fn batch_size(&self) -> usize {
476 self.batch_size
477 }
478}