1use num_complex::Complex64;
2
3use crate::{
4 Result,
5 algorithms::{
6 AlgorithmIterationMetrics, StepOutput, StepSummary,
7 objective::{LossType, point_loss},
8 },
9 array_layout::checked_len_2d,
10 backend::FftDirection,
11 error::Error,
12 measurements::MeasurementRead,
13 model::{FourierOffset, ImagePlaneModel, fftshift_copy, ifftshift_copy},
14 reconstruction::{
15 AdmmAuxiliaryState, AlgorithmAuxiliaryState, Batch, ReconstructionProblem,
16 ReconstructionState,
17 },
18};
19
20use super::ReconstructionAlgorithm;
21
22#[derive(Clone, Debug, Default)]
25pub struct AdmmIterationMetrics {
26 primal_residual_sum_squares: f64,
27 dual_residual_sum_squares: f64,
28 residual_count: usize,
29}
30
31impl AdmmIterationMetrics {
32 pub fn primal_residual_rms(&self) -> Option<f64> {
34 (self.residual_count > 0)
35 .then(|| (self.primal_residual_sum_squares / self.residual_count as f64).sqrt())
36 }
37
38 pub fn dual_residual_rms(&self) -> Option<f64> {
40 (self.residual_count > 0)
41 .then(|| (self.dual_residual_sum_squares / self.residual_count as f64).sqrt())
42 }
43
44 fn push_dual_change(&mut self, change: Complex64, penalty: f64) {
45 self.dual_residual_sum_squares += penalty * penalty * change.norm_sqr();
46 }
47
48 fn push_primal_residual(&mut self, residual: Complex64) {
49 self.primal_residual_sum_squares += residual.norm_sqr();
50 self.residual_count += 1;
51 }
52}
53
54impl AlgorithmIterationMetrics for AdmmIterationMetrics {
55 fn merge(&mut self, other: Self) {
56 self.primal_residual_sum_squares += other.primal_residual_sum_squares;
57 self.dual_residual_sum_squares += other.dual_residual_sum_squares;
58 self.residual_count += other.residual_count;
59 }
60
61 fn append_records(
62 &self,
63 iteration: usize,
64 output: &mut Vec<crate::reconstruction::AlgorithmMetricRecord>,
65 ) {
66 if let Some(value) = self.primal_residual_rms() {
67 output.push(crate::reconstruction::AlgorithmMetricRecord {
68 iteration,
69 namespace: "admm".into(),
70 metric: "primal_residual_rms".into(),
71 value,
72 });
73 }
74 if let Some(value) = self.dual_residual_rms() {
75 output.push(crate::reconstruction::AlgorithmMetricRecord {
76 iteration,
77 namespace: "admm".into(),
78 metric: "dual_residual_rms".into(),
79 value,
80 });
81 }
82 }
83}
84
85#[derive(Clone, Debug)]
117pub struct Admm {
118 pub iterations: usize,
120 pub object_step: f64,
122 pub penalty: f64,
125 pub dual_relaxation: f64,
127 pub batch_size: usize,
130 pub epsilon: f64,
132}
133
134impl Default for Admm {
135 fn default() -> Self {
136 Self {
137 iterations: 100,
138 object_step: 0.8,
139 penalty: 1.0,
140 dual_relaxation: 1.0,
141 batch_size: usize::MAX,
142 epsilon: 1e-10,
143 }
144 }
145}
146
147impl Admm {
148 pub fn iterations(mut self, iterations: usize) -> Self {
150 self.iterations = iterations;
151 self
152 }
153
154 pub fn object_step(mut self, step: f64) -> Self {
156 self.object_step = step;
157 self
158 }
159
160 pub fn penalty(mut self, penalty: f64) -> Self {
162 self.penalty = penalty;
163 self
164 }
165
166 pub fn dual_relaxation(mut self, relaxation: f64) -> Self {
168 self.dual_relaxation = relaxation;
169 self
170 }
171
172 pub fn batch_size(mut self, batch_size: usize) -> Self {
174 self.batch_size = batch_size;
175 self
176 }
177}
178
179impl ReconstructionAlgorithm for Admm {
180 type IterationMetrics = AdmmIterationMetrics;
181
182 fn validate(&self) -> Result<()> {
183 if !self.object_step.is_finite() || self.object_step <= 0.0 {
184 return Err(Error::InvalidParameter {
185 name: "object_step",
186 reason: "must be finite and positive".into(),
187 });
188 }
189 if !self.penalty.is_finite() || self.penalty <= 0.0 {
190 return Err(Error::InvalidParameter {
191 name: "penalty",
192 reason: "must be finite and positive".into(),
193 });
194 }
195 if !self.dual_relaxation.is_finite() || !(0.0..=2.0).contains(&self.dual_relaxation) {
196 return Err(Error::InvalidParameter {
197 name: "dual_relaxation",
198 reason: "must be finite and between zero and two".into(),
199 });
200 }
201 if self.batch_size == 0 || !self.epsilon.is_finite() || self.epsilon <= 0.0 {
202 return Err(Error::InvalidParameter {
203 name: "batch_size/epsilon",
204 reason: "batch size must be non-zero and epsilon finite and positive".into(),
205 });
206 }
207 Ok(())
208 }
209
210 fn step<M: MeasurementRead>(
211 &mut self,
212 problem: &ReconstructionProblem<M>,
213 state: &mut ReconstructionState,
214 batch: &Batch,
215 _iteration: usize,
216 ) -> Result<StepOutput<Self::IterationMetrics>> {
217 let expected = admm_auxiliary_len(&problem.model)?;
218 let mut auxiliary = match state.algorithm_auxiliary.take() {
219 None => AdmmAuxiliaryState {
220 auxiliary_fields: vec![Complex64::default(); expected],
221 dual_fields: vec![Complex64::default(); expected],
222 },
223 Some(AlgorithmAuxiliaryState::Admm(auxiliary)) => auxiliary,
224 Some(_) => {
225 return Err(Error::InvalidModel(
226 "ADMM cannot resume auxiliary state owned by another algorithm".into(),
227 ));
228 }
229 };
230 let result = self.step_with_auxiliary(problem, state, batch, &mut auxiliary);
231 state.algorithm_auxiliary = Some(AlgorithmAuxiliaryState::Admm(auxiliary));
232 result
233 }
234
235 fn iterations(&self) -> usize {
236 self.iterations
237 }
238
239 fn batch_size(&self) -> usize {
240 self.batch_size
241 }
242}
243
244impl Admm {
245 fn step_with_auxiliary<M: MeasurementRead>(
246 &self,
247 problem: &ReconstructionProblem<M>,
248 state: &mut ReconstructionState,
249 batch: &Batch,
250 auxiliary: &mut AdmmAuxiliaryState,
251 ) -> Result<StepOutput<AdmmIterationMetrics>> {
252 let model = &problem.model;
253 let shape = model.image_shape;
254 let image_len = checked_len_2d(shape)?;
255 let expected = admm_auxiliary_len(model)?;
256 if auxiliary.auxiliary_fields.len() != expected || auxiliary.dual_fields.len() != expected {
257 return Err(Error::InvalidModel(
258 "ADMM auxiliary state does not match the problem modes".into(),
259 ));
260 }
261 state
262 .scratch
263 .object_gradient
264 .resize(state.object_spectrum.len(), Complex64::default());
265 state.scratch.object_gradient.fill(Complex64::default());
266 let maximum_pupil_power = state
267 .pupil
268 .values
269 .as_slice()
270 .iter()
271 .map(|value| value.norm_sqr())
272 .fold(0.0, f64::max)
273 .max(self.epsilon);
274 let mut summary = StepSummary::default();
275 let mut metrics = AdmmIterationMetrics::default();
276 let mut active_frames = 0;
277
278 for &frame in &batch.indices {
279 let frame_weight = problem.measurements.frame_weight(frame)?;
280 if frame_weight == 0.0 {
281 summary.push_frame(frame, 0.0, 0.0);
282 continue;
283 }
284 active_frames += 1;
285 let single_source = [(frame, 1.0)];
286 let sources = frame_sources(model, frame, &single_source);
287 let source_weight_sum: f64 = sources.iter().map(|&(_, weight)| weight).sum();
288 let mode_start = frame_mode_start(model, frame);
289 let multiplex_len =
290 image_len
291 .checked_mul(sources.len())
292 .ok_or_else(|| Error::ShapeOverflow {
293 shape: vec![sources.len(), shape.0, shape.1],
294 })?;
295 state
296 .scratch
297 .multiplex_fields
298 .resize(multiplex_len, Complex64::default());
299 state.scratch.multiplex_offsets.clear();
300 state.scratch.projected_field.fill(Complex64::default());
301
302 for (local_mode, &(source, source_weight)) in sources.iter().enumerate() {
305 let offset = state.effective_source_offset(model, source)?;
306 state.scratch.multiplex_offsets.push(offset);
307 compute_source_field(problem, state, source, offset)?;
308 let local_start = local_mode * image_len;
309 state.scratch.multiplex_fields[local_start..local_start + image_len]
310 .copy_from_slice(&state.scratch.field);
311 let auxiliary_start = (mode_start + local_mode) * image_len;
312 for pixel in 0..image_len {
313 let field = state.scratch.field[pixel];
314 let consensus = field + auxiliary.dual_fields[auxiliary_start + pixel];
315 state.scratch.projected_field[pixel].re += source_weight * field.norm_sqr();
316 state.scratch.projected_field[pixel].im += source_weight * consensus.norm_sqr();
317 }
318 }
319
320 let measured = problem.measurements.frame(frame)?;
321 let mask = problem.measurements.frame_mask(frame)?;
322 let gain = state
323 .frame_gains
324 .as_ref()
325 .map_or(1.0, |values| values[frame]);
326 if !gain.is_finite() || gain <= 0.0 {
327 return Err(Error::InvalidModel(format!(
328 "state frame {frame} has invalid gain {gain}"
329 )));
330 }
331 let mut frame_loss = 0.0;
332 let mut valid_pixels = 0;
333 for pixel in 0..image_len {
334 if mask.is_some_and(|values| values[pixel] == 0) {
335 continue;
336 }
337 valid_pixels += 1;
338 let background = background_value(state, frame, pixel, image_len);
339 let predicted = gain * state.scratch.projected_field[pixel].re + background;
340 frame_loss += point_loss(predicted, measured[pixel], LossType::AmplitudeMse);
341 }
342 if valid_pixels == 0 {
343 return Err(Error::InvalidMeasurements(format!(
344 "frame {frame} has no unmasked pixels"
345 )));
346 }
347 summary.push_frame(frame, frame_loss / valid_pixels as f64, frame_weight);
348
349 for (local_mode, &(source, source_weight)) in sources.iter().enumerate() {
350 let local_start = local_mode * image_len;
351 state.scratch.field.copy_from_slice(
352 &state.scratch.multiplex_fields[local_start..local_start + image_len],
353 );
354 let auxiliary_start = (mode_start + local_mode) * image_len;
355 for pixel in 0..image_len {
356 let index = auxiliary_start + pixel;
357 if mask.is_some_and(|values| values[pixel] == 0) {
358 auxiliary.auxiliary_fields[index] = state.scratch.field[pixel];
359 auxiliary.dual_fields[index] = Complex64::default();
360 state.scratch.difference[pixel] = Complex64::default();
361 continue;
362 }
363 let background = background_value(state, frame, pixel, image_len);
364 let target = ((measured[pixel] - background) / gain).max(0.0).sqrt();
365 let consensus = state.scratch.field[pixel] + auxiliary.dual_fields[index];
366 let consensus_norm = state.scratch.projected_field[pixel].im;
367 let projected = if consensus_norm > self.epsilon {
368 consensus * (target / consensus_norm.sqrt())
369 } else {
370 Complex64::new(target / source_weight_sum.sqrt(), 0.0)
371 };
372 let auxiliary_value = (self.penalty * consensus + frame_weight * projected)
373 / (self.penalty + frame_weight);
374 metrics.push_dual_change(
375 auxiliary_value - auxiliary.auxiliary_fields[index],
376 self.penalty,
377 );
378 auxiliary.auxiliary_fields[index] = auxiliary_value;
379 state.scratch.difference[pixel] =
380 auxiliary_value - auxiliary.dual_fields[index] - state.scratch.field[pixel];
381 }
382 state.backend.fft2(
383 &mut state.scratch.difference,
384 shape,
385 FftDirection::Forward,
386 &mut state.scratch.column,
387 )?;
388 fftshift_copy(
389 &state.scratch.difference,
390 &mut state.scratch.projected_spectrum,
391 shape,
392 );
393 for pixel in 0..image_len {
394 state.scratch.difference[pixel] = state.pupil.values.as_slice()[pixel].conj()
395 * state.scratch.projected_spectrum[pixel]
396 / (maximum_pupil_power + self.epsilon);
397 }
398 model.insert_patch_adjoint_slice_at_offset(
399 &mut state.scratch.object_gradient,
400 source,
401 &state.scratch.difference,
402 source_weight / source_weight_sum,
403 state.scratch.multiplex_offsets[local_mode],
404 )?;
405 }
406 }
407
408 if active_frames > 0 {
409 let step = self.object_step / active_frames as f64;
410 for (object, &update) in state
411 .object_spectrum
412 .as_slice_mut()
413 .iter_mut()
414 .zip(&state.scratch.object_gradient)
415 {
416 *object += step * update;
417 }
418 }
419
420 for &frame in &batch.indices {
421 if problem.measurements.frame_weight(frame)? == 0.0 {
422 continue;
423 }
424 let single_source = [(frame, 1.0)];
425 let sources = frame_sources(model, frame, &single_source);
426 let mode_start = frame_mode_start(model, frame);
427 let mask = problem.measurements.frame_mask(frame)?;
428 for (local_mode, &(source, _)) in sources.iter().enumerate() {
429 let offset = state.effective_source_offset(model, source)?;
430 compute_source_field(problem, state, source, offset)?;
431 let auxiliary_start = (mode_start + local_mode) * image_len;
432 for pixel in 0..image_len {
433 let index = auxiliary_start + pixel;
434 if mask.is_some_and(|values| values[pixel] == 0) {
435 auxiliary.auxiliary_fields[index] = state.scratch.field[pixel];
436 auxiliary.dual_fields[index] = Complex64::default();
437 } else {
438 let primal_residual =
439 state.scratch.field[pixel] - auxiliary.auxiliary_fields[index];
440 metrics.push_primal_residual(primal_residual);
441 auxiliary.dual_fields[index] += self.dual_relaxation * primal_residual;
442 }
443 }
444 }
445 }
446 Ok(StepOutput { summary, metrics })
447 }
448}
449
450pub(crate) fn admm_auxiliary_len(model: &ImagePlaneModel) -> Result<usize> {
451 let mode_count = model.multiplexing_matrix.as_ref().map_or_else(
452 || Ok(model.frame_count()),
453 |matrix| {
454 matrix.iter().try_fold(0_usize, |count, row| {
455 count
456 .checked_add(row.len())
457 .ok_or_else(|| Error::InvalidShape("ADMM source mode count overflows".into()))
458 })
459 },
460 )?;
461 checked_len_2d(model.image_shape)?
462 .checked_mul(mode_count)
463 .ok_or_else(|| Error::InvalidShape("ADMM auxiliary length overflows".into()))
464}
465
466fn frame_mode_start(model: &ImagePlaneModel, frame: usize) -> usize {
467 model
468 .multiplexing_matrix
469 .as_ref()
470 .map_or(frame, |matrix| matrix[..frame].iter().map(Vec::len).sum())
471}
472
473fn frame_sources<'a>(
474 model: &'a ImagePlaneModel,
475 frame: usize,
476 single_source: &'a [(usize, f64); 1],
477) -> &'a [(usize, f64)] {
478 model
479 .multiplexing_matrix
480 .as_ref()
481 .map_or(single_source, |matrix| matrix[frame].as_slice())
482}
483
484fn compute_source_field<M: MeasurementRead>(
485 problem: &ReconstructionProblem<M>,
486 state: &mut ReconstructionState,
487 source: usize,
488 offset: FourierOffset,
489) -> Result<()> {
490 let model = &problem.model;
491 let shape = model.image_shape;
492 model.extract_patch_at_offset(
493 state.object_spectrum.view(),
494 source,
495 offset,
496 &mut state.scratch.patch,
497 )?;
498 for pixel in 0..state.scratch.patch.len() {
499 state.scratch.exit_spectrum[pixel] =
500 state.scratch.patch[pixel] * state.pupil.values.as_slice()[pixel];
501 }
502 ifftshift_copy(
503 &state.scratch.exit_spectrum,
504 &mut state.scratch.field,
505 shape,
506 );
507 state.backend.fft2(
508 &mut state.scratch.field,
509 shape,
510 FftDirection::Inverse,
511 &mut state.scratch.column,
512 )
513}
514
515fn background_value(
516 state: &ReconstructionState,
517 frame: usize,
518 pixel: usize,
519 image_len: usize,
520) -> f64 {
521 state.background.as_ref().map_or(0.0, |values| {
522 values[if values.len() == image_len {
523 pixel
524 } else {
525 frame * image_len + pixel
526 }]
527 })
528}