fpm_rs/algorithms/
alternating_projection.rs1use 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};
13
14#[derive(Clone, Debug)]
37pub struct AlternatingProjection {
38 pub iterations: usize,
40 pub object_step: f64,
42 pub batch_size: usize,
44 pub epsilon: f64,
46 pub loss_type: LossType,
49}
50
51impl Default for AlternatingProjection {
52 fn default() -> Self {
53 Self {
54 iterations: 50,
55 object_step: 1.0,
56 batch_size: 1,
57 epsilon: 1e-10,
58 loss_type: LossType::AmplitudeMse,
59 }
60 }
61}
62
63impl AlternatingProjection {
64 pub fn iterations(mut self, iterations: usize) -> Self {
66 self.iterations = iterations;
67 self
68 }
69
70 pub fn object_step(mut self, object_step: f64) -> Self {
72 self.object_step = object_step;
73 self
74 }
75
76 pub fn batch_size(mut self, batch_size: usize) -> Self {
78 self.batch_size = batch_size;
79 self
80 }
81}
82
83impl ReconstructionAlgorithm for AlternatingProjection {
84 type IterationMetrics = NoIterationMetrics;
85
86 fn validate(&self) -> Result<()> {
87 if !self.object_step.is_finite() || self.object_step <= 0.0 {
88 return Err(Error::InvalidParameter {
89 name: "object_step",
90 reason: "must be finite and positive".into(),
91 });
92 }
93 if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
94 return Err(Error::InvalidParameter {
95 name: "epsilon",
96 reason: "must be finite and positive".into(),
97 });
98 }
99 if self.batch_size == 0 {
100 return Err(Error::InvalidParameter {
101 name: "batch_size",
102 reason: "must be greater than zero".into(),
103 });
104 }
105 Ok(())
106 }
107
108 fn step<M: MeasurementRead>(
109 &mut self,
110 problem: &ReconstructionProblem<M>,
111 state: &mut ReconstructionState,
112 batch: &Batch,
113 _iteration: usize,
114 ) -> Result<StepOutput<Self::IterationMetrics>> {
115 Ok(projection_update(
116 problem,
117 state,
118 batch,
119 UpdateConfiguration {
120 object_step: self.object_step,
121 pupil_step: None,
122 epsilon: self.epsilon,
123 loss_type: self.loss_type,
124 object_denominator: ObjectDenominator::Local,
125 constrain_pupil: true,
126 gain_update: None,
127 background_update: None,
128 momentum: None,
129 },
130 )?
131 .into())
132 }
133
134 fn iterations(&self) -> usize {
135 self.iterations
136 }
137
138 fn batch_size(&self) -> usize {
139 self.batch_size
140 }
141}