1use num_complex::Complex64;
2use std::thread;
3
4use crate::{
5 Result,
6 algorithms::{
7 AlgorithmIterationMetrics, StepOutput, StepSummary,
8 objective::{LossType, point_loss},
9 },
10 array_layout::checked_len_2d,
11 backend::FftDirection,
12 error::Error,
13 measurements::MeasurementRead,
14 model::{FourierOffset, fftshift_copy, ifftshift_copy},
15 reconstruction::{Batch, ReconstructionProblem, ReconstructionState},
16};
17
18use super::{
19 ReconstructionAlgorithm,
20 gauge::canonicalize_object_pupil,
21 regularization::{apply_complex_tv_step, apply_quadratic_smoothing_step},
22};
23
24#[derive(Clone, Copy, Debug, Default)]
30pub struct GradientDescentIterationMetrics {
31 retained_weight: f64,
32 eligible_weight: f64,
33 truncation_enabled: bool,
34}
35
36impl GradientDescentIterationMetrics {
37 pub fn retained_pixel_fraction(&self) -> Option<f64> {
39 (self.truncation_enabled && self.eligible_weight > 0.0)
40 .then(|| self.retained_weight / self.eligible_weight)
41 }
42}
43
44impl AlgorithmIterationMetrics for GradientDescentIterationMetrics {
45 fn merge(&mut self, other: Self) {
46 self.retained_weight += other.retained_weight;
47 self.eligible_weight += other.eligible_weight;
48 self.truncation_enabled |= other.truncation_enabled;
49 }
50
51 fn append_records(
52 &self,
53 iteration: usize,
54 output: &mut Vec<crate::reconstruction::AlgorithmMetricRecord>,
55 ) {
56 if let Some(value) = self.retained_pixel_fraction() {
57 output.push(crate::reconstruction::AlgorithmMetricRecord {
58 iteration,
59 namespace: "gradient_descent".into(),
60 metric: "retained_pixel_fraction".into(),
61 value,
62 });
63 }
64 }
65}
66
67#[derive(Clone, Debug)]
129pub struct GradientDescent {
130 pub iterations: usize,
132 pub object_step: f64,
134 pub batch_size: usize,
136 pub epsilon: f64,
138 pub loss_type: LossType,
140 pub poisson_truncation_threshold: Option<f64>,
145 pub recover_illumination: bool,
147 pub illumination_step: f64,
149 pub illumination_finite_difference: f64,
151 pub maximum_illumination_correction: f64,
153 pub recover_pupil: bool,
155 pub pupil_step: f64,
157 pub constrain_pupil_support: bool,
159 pub object_tv_weight: f64,
162 pub object_tv_epsilon: f64,
164 pub pupil_smoothing_weight: f64,
167 pub parallel_workers: usize,
169}
170
171impl Default for GradientDescent {
172 fn default() -> Self {
173 Self {
174 iterations: 100,
175 object_step: 0.5,
176 batch_size: 1,
177 epsilon: 1e-10,
178 loss_type: LossType::AmplitudeMse,
179 poisson_truncation_threshold: None,
180 recover_illumination: false,
181 illumination_step: 0.1,
182 illumination_finite_difference: 0.05,
183 maximum_illumination_correction: 1.0,
184 recover_pupil: false,
185 pupil_step: 0.05,
186 constrain_pupil_support: true,
187 object_tv_weight: 0.0,
188 object_tv_epsilon: 1e-6,
189 pupil_smoothing_weight: 0.0,
190 parallel_workers: std::thread::available_parallelism().map_or(1, |count| count.get()),
191 }
192 }
193}
194
195impl GradientDescent {
196 pub fn iterations(mut self, iterations: usize) -> Self {
198 self.iterations = iterations;
199 self
200 }
201
202 pub fn object_step(mut self, step: f64) -> Self {
204 self.object_step = step;
205 self
206 }
207
208 pub fn batch_size(mut self, batch_size: usize) -> Self {
210 self.batch_size = batch_size;
211 self
212 }
213
214 pub fn loss_type(mut self, loss_type: LossType) -> Self {
216 self.loss_type = loss_type;
217 self
218 }
219
220 pub fn poisson_truncation_threshold(mut self, threshold: f64) -> Self {
225 self.poisson_truncation_threshold = Some(threshold);
226 self
227 }
228
229 pub fn recover_illumination(mut self, recover: bool) -> Self {
231 self.recover_illumination = recover;
232 self
233 }
234
235 pub fn illumination_step(mut self, step: f64) -> Self {
237 self.illumination_step = step;
238 self
239 }
240
241 pub fn illumination_finite_difference(mut self, distance: f64) -> Self {
243 self.illumination_finite_difference = distance;
244 self
245 }
246
247 pub fn illumination_bounds(mut self, maximum_absolute_correction: f64) -> Self {
249 self.maximum_illumination_correction = maximum_absolute_correction;
250 self
251 }
252
253 pub fn recover_pupil(mut self, recover: bool) -> Self {
255 self.recover_pupil = recover;
256 self
257 }
258
259 pub fn pupil_step(mut self, step: f64) -> Self {
261 self.pupil_step = step;
262 self
263 }
264
265 pub fn constrain_pupil_support(mut self, constrain: bool) -> Self {
267 self.constrain_pupil_support = constrain;
268 self
269 }
270
271 pub fn object_tv(mut self, weight: f64) -> Self {
273 self.object_tv_weight = weight;
274 self
275 }
276
277 pub fn object_tv_epsilon(mut self, epsilon: f64) -> Self {
279 self.object_tv_epsilon = epsilon;
280 self
281 }
282
283 pub fn pupil_smoothing(mut self, weight: f64) -> Self {
285 self.pupil_smoothing_weight = weight;
286 self
287 }
288
289 pub fn parallel_workers(mut self, workers: usize) -> Self {
292 self.parallel_workers = workers;
293 self
294 }
295}
296
297impl ReconstructionAlgorithm for GradientDescent {
298 type IterationMetrics = GradientDescentIterationMetrics;
299
300 fn validate(&self) -> Result<()> {
301 if !self.object_step.is_finite() || self.object_step <= 0.0 {
302 return Err(Error::InvalidParameter {
303 name: "object_step",
304 reason: "must be finite and positive".into(),
305 });
306 }
307 if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
308 return Err(Error::InvalidParameter {
309 name: "epsilon",
310 reason: "must be finite and positive".into(),
311 });
312 }
313 if self.batch_size == 0 {
314 return Err(Error::InvalidParameter {
315 name: "batch_size",
316 reason: "must be greater than zero".into(),
317 });
318 }
319 if let Some(threshold) = self.poisson_truncation_threshold {
320 if !threshold.is_finite() || threshold <= 0.0 {
321 return Err(Error::InvalidParameter {
322 name: "poisson_truncation_threshold",
323 reason: "must be finite and positive when provided".into(),
324 });
325 }
326 if self.loss_type != LossType::PoissonNegativeLogLikelihood {
327 return Err(Error::InvalidParameter {
328 name: "poisson_truncation_threshold",
329 reason: "requires Poisson negative log likelihood".into(),
330 });
331 }
332 }
333 if !self.illumination_step.is_finite() || self.illumination_step <= 0.0 {
334 return Err(Error::InvalidParameter {
335 name: "illumination_step",
336 reason: "must be finite and positive".into(),
337 });
338 }
339 if !self.illumination_finite_difference.is_finite()
340 || self.illumination_finite_difference <= 0.0
341 {
342 return Err(Error::InvalidParameter {
343 name: "illumination_finite_difference",
344 reason: "must be finite and positive".into(),
345 });
346 }
347 if !self.maximum_illumination_correction.is_finite()
348 || self.maximum_illumination_correction <= 0.0
349 {
350 return Err(Error::InvalidParameter {
351 name: "maximum_illumination_correction",
352 reason: "must be finite and positive".into(),
353 });
354 }
355 if !self.pupil_step.is_finite() || self.pupil_step < 0.0 {
356 return Err(Error::InvalidParameter {
357 name: "pupil_step",
358 reason: "must be finite and non-negative".into(),
359 });
360 }
361 if !self.object_tv_weight.is_finite() || self.object_tv_weight < 0.0 {
362 return Err(Error::InvalidParameter {
363 name: "object_tv_weight",
364 reason: "must be finite and non-negative".into(),
365 });
366 }
367 if !self.object_tv_epsilon.is_finite() || self.object_tv_epsilon <= 0.0 {
368 return Err(Error::InvalidParameter {
369 name: "object_tv_epsilon",
370 reason: "must be finite and positive".into(),
371 });
372 }
373 if !self.pupil_smoothing_weight.is_finite() || self.pupil_smoothing_weight < 0.0 {
374 return Err(Error::InvalidParameter {
375 name: "pupil_smoothing_weight",
376 reason: "must be finite and non-negative".into(),
377 });
378 }
379 if self.pupil_smoothing_weight > 0.0 && !self.recover_pupil {
380 return Err(Error::InvalidParameter {
381 name: "pupil_smoothing_weight",
382 reason: "requires pupil recovery to be enabled".into(),
383 });
384 }
385 if self.parallel_workers == 0 {
386 return Err(Error::InvalidParameter {
387 name: "parallel_workers",
388 reason: "must be greater than zero".into(),
389 });
390 }
391 Ok(())
392 }
393
394 fn canonicalize_state<M: MeasurementRead>(
395 &self,
396 problem: &ReconstructionProblem<M>,
397 state: &mut ReconstructionState,
398 ) -> Result<()> {
399 if self.recover_pupil {
400 canonicalize_object_pupil(problem, state)?;
401 }
402 Ok(())
403 }
404
405 fn step<M: MeasurementRead>(
406 &mut self,
407 problem: &ReconstructionProblem<M>,
408 state: &mut ReconstructionState,
409 batch: &Batch,
410 iteration: usize,
411 ) -> Result<StepOutput<Self::IterationMetrics>> {
412 let truncation = if let Some(scale) = state.scratch.poisson_truncation_scale.take() {
413 Some(TruncationStatistics { scale })
414 } else {
415 self.compute_truncation_statistics(problem, state, batch)?
416 };
417 if self.parallel_workers > 1 && batch.indices.len() > 1 {
418 return self.parallel_step(problem, state, batch, iteration, truncation);
419 }
420 let model = &problem.model;
421 let shape = model.image_shape;
422 let image_len = checked_len_2d(shape)?;
423 let maximum_pupil_power = state
424 .pupil
425 .values
426 .as_slice()
427 .iter()
428 .map(|value| value.norm_sqr())
429 .fold(0.0, f64::max)
430 .max(self.epsilon);
431 let mut diagnostics = StepSummary::default();
432 let mut metrics = GradientDescentIterationMetrics {
433 truncation_enabled: truncation.is_some(),
434 ..GradientDescentIterationMetrics::default()
435 };
436 let mut active_frames = 0;
437 if self.recover_illumination {
438 prepare_illumination_accumulators(model, state)?;
439 }
440 state
441 .scratch
442 .object_gradient
443 .resize(state.object_spectrum.len(), Complex64::default());
444 state.scratch.object_gradient.fill(Complex64::default());
445 if self.recover_pupil {
446 state.scratch.pupil_gradient.fill(Complex64::default());
447 }
448
449 for &frame in &batch.indices {
450 let frame_weight = problem.measurements.frame_weight(frame)?;
451 if frame_weight == 0.0 {
452 diagnostics.push_frame(frame, 0.0, 0.0);
453 continue;
454 }
455 active_frames += 1;
456 let single_source = [(frame, 1.0)];
457 let sources = model
458 .multiplexing_matrix
459 .as_ref()
460 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
461
462 state.scratch.projected_field.fill(Complex64::default());
463 for &(source, source_weight) in sources {
464 let offset = state.effective_source_offset(model, source)?;
465 compute_source_field(problem, state, source, offset)?;
466 for (predicted, field) in state
467 .scratch
468 .projected_field
469 .iter_mut()
470 .zip(&state.scratch.field)
471 {
472 predicted.re += source_weight * field.norm_sqr();
473 }
474 }
475
476 let measured = problem.measurements.frame(frame)?;
477 let mask = problem.measurements.frame_mask(frame)?;
478 let gain = state
479 .frame_gains
480 .as_ref()
481 .map_or(1.0, |values| values[frame]);
482 if !gain.is_finite() || gain <= 0.0 {
483 return Err(Error::InvalidModel(format!(
484 "state frame {frame} has invalid gain {gain}"
485 )));
486 }
487 let mut frame_loss = 0.0;
488 let mut valid_pixels = 0;
489 for pixel in 0..image_len {
490 if mask.is_some_and(|values| values[pixel] == 0) {
491 state.scratch.projected_field[pixel] = Complex64::default();
492 state.scratch.data_gradient_mask[pixel] = 0;
493 continue;
494 }
495 valid_pixels += 1;
496 let background = background_value(state, frame, pixel, image_len);
497 let intrinsic_prediction = state.scratch.projected_field[pixel].re.max(0.0);
498 let target_intensity = ((measured[pixel] - background) / gain).max(0.0);
499 frame_loss += point_loss(intrinsic_prediction, target_intensity, self.loss_type);
500 let retained = truncation.is_none_or(|statistics| {
501 truncation_accepts(intrinsic_prediction, target_intensity, statistics.scale)
502 });
503 state.scratch.data_gradient_mask[pixel] = u8::from(retained);
504 if truncation.is_some() {
505 metrics.eligible_weight += frame_weight;
506 if retained {
507 metrics.retained_weight += frame_weight;
508 }
509 }
510 state.scratch.projected_field[pixel] = Complex64::new(
511 if retained {
512 descent_factor(
513 intrinsic_prediction,
514 target_intensity,
515 self.loss_type,
516 self.epsilon,
517 )
518 } else {
519 0.0
520 },
521 intrinsic_prediction,
522 );
523 }
524 if valid_pixels == 0 {
525 return Err(Error::InvalidMeasurements(format!(
526 "frame {frame} has no unmasked pixels"
527 )));
528 }
529 let valid_pixel_scale = image_len as f64 / valid_pixels as f64;
532 diagnostics.push_frame(frame, frame_loss / valid_pixels as f64, frame_weight);
533
534 for &(source, source_weight) in sources {
535 let offset = state.effective_source_offset(model, source)?;
536 compute_source_field(problem, state, source, offset)?;
537 if self.recover_illumination {
538 for (reference, field) in state
539 .scratch
540 .calibration_reference
541 .iter_mut()
542 .zip(&state.scratch.field)
543 {
544 *reference = field.norm_sqr();
545 }
546 }
547 for pixel in 0..image_len {
548 state.scratch.field[pixel] *=
549 source_weight * state.scratch.projected_field[pixel].re;
550 }
551 state.backend.fft2(
552 &mut state.scratch.field,
553 shape,
554 FftDirection::Forward,
555 &mut state.scratch.column,
556 )?;
557 fftshift_copy(
558 &state.scratch.field,
559 &mut state.scratch.projected_spectrum,
560 shape,
561 );
562 if self.recover_pupil {
563 let maximum_object_power = state
564 .scratch
565 .patch
566 .iter()
567 .map(|value| value.norm_sqr())
568 .fold(0.0, f64::max)
569 .max(self.epsilon);
570 for pixel in 0..image_len {
571 state.scratch.pupil_gradient[pixel] += frame_weight
572 * valid_pixel_scale
573 * state.scratch.patch[pixel].conj()
574 * state.scratch.projected_spectrum[pixel]
575 / (maximum_object_power + self.epsilon);
576 }
577 }
578 for pixel in 0..image_len {
579 state.scratch.difference[pixel] = state.pupil.values.as_slice()[pixel].conj()
580 * state.scratch.projected_spectrum[pixel]
581 / (maximum_pupil_power + self.epsilon);
582 }
583 model.insert_patch_adjoint_slice_at_offset(
584 &mut state.scratch.object_gradient,
585 source,
586 &state.scratch.difference,
587 frame_weight * valid_pixel_scale,
588 offset,
589 )?;
590 if self.recover_illumination {
591 let gradient = illumination_gradient(
592 problem,
593 state,
594 source,
595 offset,
596 IlluminationGradientConfiguration {
597 source_weight,
598 valid_pixels,
599 distance: self.illumination_finite_difference,
600 loss_type: self.loss_type,
601 epsilon: self.epsilon,
602 },
603 )?;
604 state.scratch.illumination_gradient[source].0 +=
605 frame_weight * gradient.row.gradient;
606 state.scratch.illumination_gradient[source].1 +=
607 frame_weight * gradient.column.gradient;
608 state.scratch.illumination_curvature[source].0 +=
609 frame_weight * gradient.row.curvature;
610 state.scratch.illumination_curvature[source].1 +=
611 frame_weight * gradient.column.curvature;
612 state.scratch.illumination_weight[source] += frame_weight;
613 }
614 }
615 }
616 if active_frames > 0 {
617 let step = self.object_step / active_frames as f64;
618 for (object, &gradient) in state
619 .object_spectrum
620 .as_slice_mut()
621 .iter_mut()
622 .zip(&state.scratch.object_gradient)
623 {
624 *object -= step * gradient;
625 }
626 if self.recover_pupil {
627 let pupil_step = self.pupil_step / active_frames as f64;
628 for (pupil, &gradient) in state
629 .pupil
630 .values
631 .as_slice_mut()
632 .iter_mut()
633 .zip(&state.scratch.pupil_gradient)
634 {
635 *pupil -= pupil_step * gradient;
636 if !pupil.re.is_finite() || !pupil.im.is_finite() {
637 return Err(Error::Numerical(
638 "pupil update produced a non-finite value".into(),
639 ));
640 }
641 }
642 }
643 }
644 let batch_fraction = batch.indices.len() as f64 / model.frame_count() as f64;
645 if self.object_tv_weight > 0.0 {
646 apply_object_tv(
647 state,
648 model.reconstruction_shape,
649 batch_fraction * self.object_tv_weight,
650 self.object_tv_epsilon,
651 )?;
652 }
653 if self.recover_pupil && self.pupil_smoothing_weight > 0.0 {
654 apply_quadratic_smoothing_step(
655 state.pupil.values.as_slice_mut(),
656 model.image_shape,
657 batch_fraction * self.pupil_smoothing_weight,
658 &mut state.scratch.pupil_gradient,
659 )?;
660 }
661 if self.recover_pupil && self.constrain_pupil_support {
662 state.pupil.apply_support();
663 }
664 if self.recover_illumination {
665 self.apply_illumination_update(model, state)?;
666 }
667 Ok(StepOutput {
668 summary: diagnostics,
669 metrics,
670 })
671 }
672
673 fn iterations(&self) -> usize {
674 self.iterations
675 }
676
677 fn batch_size(&self) -> usize {
678 self.batch_size
679 }
680}
681
682struct ParallelWorkerResult {
683 position: usize,
684 object_delta: Vec<Complex64>,
685 pupil_delta: Vec<Complex64>,
686 illumination_gradient: Vec<(f64, f64)>,
687 illumination_curvature: Vec<(f64, f64)>,
688 illumination_weight: Vec<f64>,
689 diagnostics: StepSummary,
690 metrics: GradientDescentIterationMetrics,
691 active_frames: usize,
692}
693
694#[derive(Clone, Copy, Debug)]
695struct TruncationStatistics {
696 scale: f64,
698}
699
700impl GradientDescent {
701 fn compute_truncation_statistics<M: MeasurementRead>(
702 &self,
703 problem: &ReconstructionProblem<M>,
704 state: &mut ReconstructionState,
705 batch: &Batch,
706 ) -> Result<Option<TruncationStatistics>> {
707 let Some(threshold) = self.poisson_truncation_threshold else {
708 return Ok(None);
709 };
710 let model = &problem.model;
711 let image_len = checked_len_2d(model.image_shape)?;
712 let mut weighted_residual_sum = 0.0;
713 let mut pixel_weight_sum = 0.0;
714
715 for &frame in &batch.indices {
716 let frame_weight = problem.measurements.frame_weight(frame)?;
717 if frame_weight == 0.0 {
718 continue;
719 }
720 let single_source = [(frame, 1.0)];
721 let sources = model
722 .multiplexing_matrix
723 .as_ref()
724 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
725 state.scratch.projected_field.fill(Complex64::default());
726 for &(source, source_weight) in sources {
727 let offset = state.effective_source_offset(model, source)?;
728 compute_source_field(problem, state, source, offset)?;
729 for (predicted, field) in state
730 .scratch
731 .projected_field
732 .iter_mut()
733 .zip(&state.scratch.field)
734 {
735 predicted.re += source_weight * field.norm_sqr();
736 }
737 }
738
739 let measured = problem.measurements.frame(frame)?;
740 let mask = problem.measurements.frame_mask(frame)?;
741 let gain = state
742 .frame_gains
743 .as_ref()
744 .map_or(1.0, |values| values[frame]);
745 if !gain.is_finite() || gain <= 0.0 {
746 return Err(Error::InvalidModel(format!(
747 "state frame {frame} has invalid gain {gain}"
748 )));
749 }
750 let mut valid_pixels = 0;
751 for pixel in 0..image_len {
752 if mask.is_some_and(|values| values[pixel] == 0) {
753 continue;
754 }
755 valid_pixels += 1;
756 let background = background_value(state, frame, pixel, image_len);
757 let predicted = state.scratch.projected_field[pixel].re.max(0.0);
758 let target = ((measured[pixel] - background) / gain).max(0.0);
759 weighted_residual_sum += frame_weight * (target - predicted).abs();
760 pixel_weight_sum += frame_weight;
761 }
762 if valid_pixels == 0 {
763 return Err(Error::InvalidMeasurements(format!(
764 "frame {frame} has no unmasked pixels"
765 )));
766 }
767 }
768
769 let object_norm = state
770 .object_spectrum
771 .as_slice()
772 .iter()
773 .map(|value| value.norm_sqr())
774 .sum::<f64>()
775 .sqrt();
776 let object_rms = object_norm.max(self.epsilon.sqrt());
780 let mean_residual = if pixel_weight_sum > 0.0 {
781 weighted_residual_sum / pixel_weight_sum
782 } else {
783 0.0
784 };
785 let scale = threshold * mean_residual / object_rms;
786 if !scale.is_finite() || scale < 0.0 {
787 return Err(Error::Numerical(
788 "Poisson truncation statistic is non-finite".into(),
789 ));
790 }
791 Ok(Some(TruncationStatistics { scale }))
792 }
793
794 fn parallel_step<M: MeasurementRead>(
795 &self,
796 problem: &ReconstructionProblem<M>,
797 state: &mut ReconstructionState,
798 batch: &Batch,
799 iteration: usize,
800 truncation: Option<TruncationStatistics>,
801 ) -> Result<StepOutput<GradientDescentIterationMetrics>> {
802 let worker_count = self.parallel_workers.min(batch.indices.len());
803 if self.recover_illumination {
804 prepare_illumination_accumulators(&problem.model, state)?;
805 }
806 state.scratch.poisson_truncation_scale = truncation.map(|statistics| statistics.scale);
807 let base_state = state.clone();
808 state.scratch.poisson_truncation_scale = None;
809 let mut results = thread::scope(|scope| -> Result<Vec<ParallelWorkerResult>> {
810 let base_chunk_len = batch.indices.len() / worker_count;
811 let remainder = batch.indices.len() % worker_count;
812 let handles: Vec<_> = (0..worker_count)
813 .map(|worker| {
814 let start = worker * base_chunk_len + worker.min(remainder);
815 let length = base_chunk_len + usize::from(worker < remainder);
816 let frames = &batch.indices[start..start + length];
817 let base_state = &base_state;
818 scope.spawn(move || -> Result<ParallelWorkerResult> {
819 let mut local_state = base_state.clone();
820 let mut local_algorithm = self.clone();
821 local_algorithm.parallel_workers = 1;
822 local_algorithm.object_tv_weight = 0.0;
823 local_algorithm.pupil_smoothing_weight = 0.0;
824 local_algorithm.constrain_pupil_support = false;
825 let mut output = ParallelWorkerResult {
826 position: worker,
827 object_delta: vec![
828 Complex64::default();
829 base_state.object_spectrum.len()
830 ],
831 pupil_delta: if self.recover_pupil {
832 vec![Complex64::default(); base_state.pupil.values.len()]
833 } else {
834 Vec::new()
835 },
836 illumination_gradient: if self.recover_illumination {
837 vec![(0.0, 0.0); problem.model.source_count()]
838 } else {
839 Vec::new()
840 },
841 illumination_curvature: if self.recover_illumination {
842 vec![(0.0, 0.0); problem.model.source_count()]
843 } else {
844 Vec::new()
845 },
846 illumination_weight: if self.recover_illumination {
847 vec![0.0; problem.model.source_count()]
848 } else {
849 Vec::new()
850 },
851 diagnostics: StepSummary::default(),
852 metrics: GradientDescentIterationMetrics::default(),
853 active_frames: 0,
854 };
855 for &frame in frames {
856 local_state.scratch.poisson_truncation_scale =
857 truncation.map(|statistics| statistics.scale);
858 local_state
859 .object_spectrum
860 .as_slice_mut()
861 .copy_from_slice(base_state.object_spectrum.as_slice());
862 if self.recover_pupil {
863 local_state
864 .pupil
865 .values
866 .as_slice_mut()
867 .copy_from_slice(base_state.pupil.values.as_slice());
868 }
869 if self.recover_illumination {
870 local_state
871 .illumination_corrections
872 .clone_from(&base_state.illumination_corrections);
873 }
874 let diagnostics = local_algorithm.step(
875 problem,
876 &mut local_state,
877 &Batch::single(frame),
878 iteration,
879 )?;
880 if diagnostics.summary.weight_sum > 0.0 {
881 output.active_frames += 1;
882 for ((sum, &value), &initial) in output
883 .object_delta
884 .iter_mut()
885 .zip(local_state.object_spectrum.as_slice())
886 .zip(base_state.object_spectrum.as_slice())
887 {
888 *sum += value - initial;
889 }
890 if self.recover_pupil {
891 for ((sum, &value), &initial) in output
892 .pupil_delta
893 .iter_mut()
894 .zip(local_state.pupil.values.as_slice())
895 .zip(base_state.pupil.values.as_slice())
896 {
897 *sum += value - initial;
898 }
899 }
900 if self.recover_illumination {
901 for source in 0..problem.model.source_count() {
902 output.illumination_gradient[source].0 +=
903 local_state.scratch.illumination_gradient[source].0;
904 output.illumination_gradient[source].1 +=
905 local_state.scratch.illumination_gradient[source].1;
906 output.illumination_curvature[source].0 +=
907 local_state.scratch.illumination_curvature[source].0;
908 output.illumination_curvature[source].1 +=
909 local_state.scratch.illumination_curvature[source].1;
910 output.illumination_weight[source] +=
911 local_state.scratch.illumination_weight[source];
912 }
913 }
914 }
915 output.diagnostics.merge(diagnostics.summary);
916 output.metrics.merge(diagnostics.metrics);
917 }
918 Ok(output)
919 })
920 })
921 .collect();
922 let mut output = Vec::with_capacity(handles.len());
923 for handle in handles {
924 output.push(
925 handle.join().map_err(|_| {
926 Error::Numerical("parallel gradient worker panicked".into())
927 })??,
928 );
929 }
930 Ok(output)
931 })?;
932 results.sort_by_key(|result| result.position);
933
934 state
935 .scratch
936 .object_gradient
937 .resize(state.object_spectrum.len(), Complex64::default());
938 state.scratch.object_gradient.fill(Complex64::default());
939 if self.recover_pupil {
940 state.scratch.pupil_gradient.fill(Complex64::default());
941 }
942 if self.recover_illumination {
943 state.scratch.illumination_gradient.fill((0.0, 0.0));
944 state.scratch.illumination_curvature.fill((0.0, 0.0));
945 state.scratch.illumination_weight.fill(0.0);
946 }
947 let mut diagnostics = StepSummary::default();
948 let mut metrics = GradientDescentIterationMetrics::default();
949 let mut active_frames = 0;
950 for result in results {
951 active_frames += result.active_frames;
952 for (sum, &value) in state
953 .scratch
954 .object_gradient
955 .iter_mut()
956 .zip(&result.object_delta)
957 {
958 *sum += value;
959 }
960 if self.recover_pupil {
961 for (sum, &value) in state
962 .scratch
963 .pupil_gradient
964 .iter_mut()
965 .zip(&result.pupil_delta)
966 {
967 *sum += value;
968 }
969 }
970 if self.recover_illumination {
971 for source in 0..problem.model.source_count() {
972 state.scratch.illumination_gradient[source].0 +=
973 result.illumination_gradient[source].0;
974 state.scratch.illumination_gradient[source].1 +=
975 result.illumination_gradient[source].1;
976 state.scratch.illumination_curvature[source].0 +=
977 result.illumination_curvature[source].0;
978 state.scratch.illumination_curvature[source].1 +=
979 result.illumination_curvature[source].1;
980 state.scratch.illumination_weight[source] += result.illumination_weight[source];
981 }
982 }
983 diagnostics.merge(result.diagnostics);
984 metrics.merge(result.metrics);
985 }
986 if active_frames > 0 {
987 let normalization = active_frames as f64;
988 for (object, &sum) in state
989 .object_spectrum
990 .as_slice_mut()
991 .iter_mut()
992 .zip(&state.scratch.object_gradient)
993 {
994 *object += sum / normalization;
995 }
996 if self.recover_pupil {
997 for (pupil, &sum) in state
998 .pupil
999 .values
1000 .as_slice_mut()
1001 .iter_mut()
1002 .zip(&state.scratch.pupil_gradient)
1003 {
1004 *pupil += sum / normalization;
1005 if !pupil.re.is_finite() || !pupil.im.is_finite() {
1006 return Err(Error::Numerical(
1007 "parallel pupil update produced a non-finite value".into(),
1008 ));
1009 }
1010 }
1011 }
1012 }
1013
1014 let model = &problem.model;
1015 let batch_fraction = batch.indices.len() as f64 / model.frame_count() as f64;
1016 if self.object_tv_weight > 0.0 {
1017 apply_object_tv(
1018 state,
1019 model.reconstruction_shape,
1020 batch_fraction * self.object_tv_weight,
1021 self.object_tv_epsilon,
1022 )?;
1023 }
1024 if self.recover_pupil && self.pupil_smoothing_weight > 0.0 {
1025 apply_quadratic_smoothing_step(
1026 state.pupil.values.as_slice_mut(),
1027 model.image_shape,
1028 batch_fraction * self.pupil_smoothing_weight,
1029 &mut state.scratch.pupil_gradient,
1030 )?;
1031 }
1032 if self.recover_pupil && self.constrain_pupil_support {
1033 state.pupil.apply_support();
1034 }
1035 if self.recover_illumination {
1036 self.apply_illumination_update(model, state)?;
1037 }
1038 Ok(StepOutput {
1039 summary: diagnostics,
1040 metrics,
1041 })
1042 }
1043
1044 fn apply_illumination_update(
1045 &self,
1046 model: &crate::model::ImagePlaneModel,
1047 state: &mut ReconstructionState,
1048 ) -> Result<()> {
1049 let corrections = state.illumination_corrections.as_mut().ok_or_else(|| {
1050 Error::InvalidModel("illumination corrections were not initialized".into())
1051 })?;
1052 for (source, correction) in corrections.iter_mut().enumerate() {
1053 let weight = state.scratch.illumination_weight[source];
1054 if weight == 0.0 {
1055 continue;
1056 }
1057 let gradient = state.scratch.illumination_gradient[source];
1058 let curvature = state.scratch.illumination_curvature[source];
1059 if !gradient.0.is_finite()
1060 || !gradient.1.is_finite()
1061 || !curvature.0.is_finite()
1062 || !curvature.1.is_finite()
1063 || curvature.0 < 0.0
1064 || curvature.1 < 0.0
1065 {
1066 return Err(Error::Numerical(format!(
1067 "illumination gradient or curvature for source {source} is invalid"
1068 )));
1069 }
1070 let candidate = (
1071 (correction.0 - self.illumination_step * gradient.0 / (curvature.0 + self.epsilon))
1072 .clamp(
1073 -self.maximum_illumination_correction,
1074 self.maximum_illumination_correction,
1075 ),
1076 (correction.1 - self.illumination_step * gradient.1 / (curvature.1 + self.epsilon))
1077 .clamp(
1078 -self.maximum_illumination_correction,
1079 self.maximum_illumination_correction,
1080 ),
1081 );
1082 let base = model.source_offset(source)?;
1083 let effective = FourierOffset::new(base.row + candidate.0, base.column + candidate.1);
1084 if model.validate_source_offset(source, effective).is_ok() {
1085 *correction = candidate;
1086 }
1087 }
1088 Ok(())
1089 }
1090}
1091
1092fn prepare_illumination_accumulators(
1093 model: &crate::model::ImagePlaneModel,
1094 state: &mut ReconstructionState,
1095) -> Result<()> {
1096 match &state.illumination_corrections {
1097 None => {
1098 state.illumination_corrections = Some(vec![(0.0, 0.0); model.source_count()]);
1099 }
1100 Some(corrections) if corrections.len() != model.source_count() => {
1101 return Err(Error::InvalidModel(
1102 "illumination correction count does not match source count".into(),
1103 ));
1104 }
1105 Some(_) => {}
1106 }
1107 state
1108 .scratch
1109 .illumination_gradient
1110 .resize(model.source_count(), (0.0, 0.0));
1111 state.scratch.illumination_gradient.fill((0.0, 0.0));
1112 state
1113 .scratch
1114 .illumination_curvature
1115 .resize(model.source_count(), (0.0, 0.0));
1116 state.scratch.illumination_curvature.fill((0.0, 0.0));
1117 state
1118 .scratch
1119 .illumination_weight
1120 .resize(model.source_count(), 0.0);
1121 state.scratch.illumination_weight.fill(0.0);
1122 Ok(())
1123}
1124
1125fn apply_object_tv(
1126 state: &mut ReconstructionState,
1127 shape: (usize, usize),
1128 weight: f64,
1129 epsilon: f64,
1130) -> Result<()> {
1131 state
1132 .scratch
1133 .regularization_field
1134 .resize(state.object_spectrum.len(), Complex64::default());
1135 ifftshift_copy(
1136 state.object_spectrum.as_slice(),
1137 &mut state.scratch.regularization_field,
1138 shape,
1139 );
1140 state.backend.fft2(
1141 &mut state.scratch.regularization_field,
1142 shape,
1143 FftDirection::Inverse,
1144 &mut state.scratch.column,
1145 )?;
1146 apply_complex_tv_step(
1147 &mut state.scratch.regularization_field,
1148 shape,
1149 weight,
1150 epsilon,
1151 &mut state.scratch.object_gradient,
1152 )?;
1153 state.backend.fft2(
1154 &mut state.scratch.regularization_field,
1155 shape,
1156 FftDirection::Forward,
1157 &mut state.scratch.column,
1158 )?;
1159 fftshift_copy(
1160 &state.scratch.regularization_field,
1161 state.object_spectrum.as_slice_mut(),
1162 shape,
1163 );
1164 Ok(())
1165}
1166
1167fn descent_factor(predicted: f64, measured: f64, loss_type: LossType, epsilon: f64) -> f64 {
1168 match loss_type {
1169 LossType::AmplitudeMse => 1.0 - measured.max(0.0).sqrt() / predicted.max(epsilon).sqrt(),
1170 LossType::IntensityMse => 2.0 * (predicted - measured),
1171 LossType::PoissonNegativeLogLikelihood => 1.0 - measured.max(0.0) / predicted.max(epsilon),
1172 LossType::HuberAmplitude => {
1173 let predicted_amplitude = predicted.max(epsilon).sqrt();
1174 let residual = predicted_amplitude - measured.max(0.0).sqrt();
1175 residual.clamp(-1.0, 1.0) / (2.0 * predicted_amplitude)
1176 }
1177 }
1178}
1179
1180fn truncation_accepts(predicted: f64, measured: f64, scale: f64) -> bool {
1181 (measured - predicted).abs() <= scale * predicted.max(0.0).sqrt()
1182}
1183
1184fn compute_source_field<M: MeasurementRead>(
1185 problem: &ReconstructionProblem<M>,
1186 state: &mut ReconstructionState,
1187 source: usize,
1188 offset: FourierOffset,
1189) -> Result<()> {
1190 let model = &problem.model;
1191 let shape = model.image_shape;
1192 model.extract_patch_at_offset(
1193 state.object_spectrum.view(),
1194 source,
1195 offset,
1196 &mut state.scratch.patch,
1197 )?;
1198 for pixel in 0..state.scratch.patch.len() {
1199 state.scratch.exit_spectrum[pixel] =
1200 state.scratch.patch[pixel] * state.pupil.values.as_slice()[pixel];
1201 }
1202 ifftshift_copy(
1203 &state.scratch.exit_spectrum,
1204 &mut state.scratch.field,
1205 shape,
1206 );
1207 state.backend.fft2(
1208 &mut state.scratch.field,
1209 shape,
1210 FftDirection::Inverse,
1211 &mut state.scratch.column,
1212 )
1213}
1214
1215#[derive(Clone, Copy)]
1216struct IlluminationGradientConfiguration {
1217 source_weight: f64,
1218 valid_pixels: usize,
1219 distance: f64,
1220 loss_type: LossType,
1221 epsilon: f64,
1222}
1223
1224#[derive(Clone, Copy, Default)]
1225struct AxisGradient {
1226 gradient: f64,
1227 curvature: f64,
1228}
1229
1230#[derive(Clone, Copy, Default)]
1231struct IlluminationGradient {
1232 row: AxisGradient,
1233 column: AxisGradient,
1234}
1235
1236fn illumination_gradient<M: MeasurementRead>(
1237 problem: &ReconstructionProblem<M>,
1238 state: &mut ReconstructionState,
1239 source: usize,
1240 offset: FourierOffset,
1241 configuration: IlluminationGradientConfiguration,
1242) -> Result<IlluminationGradient> {
1243 let row = illumination_axis_gradient(
1244 problem,
1245 state,
1246 source,
1247 offset,
1248 FourierOffset::new(configuration.distance, 0.0),
1249 configuration,
1250 )?;
1251 let column = illumination_axis_gradient(
1252 problem,
1253 state,
1254 source,
1255 offset,
1256 FourierOffset::new(0.0, configuration.distance),
1257 configuration,
1258 )?;
1259 let pixels = configuration.valid_pixels as f64;
1260 Ok(IlluminationGradient {
1261 row: AxisGradient {
1262 gradient: row.gradient / pixels,
1263 curvature: row.curvature / pixels,
1264 },
1265 column: AxisGradient {
1266 gradient: column.gradient / pixels,
1267 curvature: column.curvature / pixels,
1268 },
1269 })
1270}
1271
1272fn illumination_axis_gradient<M: MeasurementRead>(
1273 problem: &ReconstructionProblem<M>,
1274 state: &mut ReconstructionState,
1275 source: usize,
1276 offset: FourierOffset,
1277 displacement: FourierOffset,
1278 configuration: IlluminationGradientConfiguration,
1279) -> Result<AxisGradient> {
1280 let model = &problem.model;
1281 let distance = displacement.row.abs() + displacement.column.abs();
1282 let plus = FourierOffset::new(
1283 offset.row + displacement.row,
1284 offset.column + displacement.column,
1285 );
1286 let minus = FourierOffset::new(
1287 offset.row - displacement.row,
1288 offset.column - displacement.column,
1289 );
1290 let plus_valid = model.validate_source_offset(source, plus).is_ok();
1291 let minus_valid = model.validate_source_offset(source, minus).is_ok();
1292 if !plus_valid && !minus_valid {
1293 return Ok(AxisGradient::default());
1294 }
1295 if plus_valid {
1296 compute_source_field(problem, state, source, plus)?;
1297 for (candidate, field) in state
1298 .scratch
1299 .difference
1300 .iter_mut()
1301 .zip(&state.scratch.field)
1302 {
1303 candidate.re = field.norm_sqr();
1304 }
1305 }
1306 if minus_valid {
1307 compute_source_field(problem, state, source, minus)?;
1308 }
1309 let mut gradient = 0.0;
1310 let mut curvature = 0.0;
1311 for pixel in 0..state.scratch.field.len() {
1312 if state.scratch.data_gradient_mask[pixel] == 0 {
1313 continue;
1314 }
1315 let derivative = match (plus_valid, minus_valid) {
1316 (true, true) => {
1317 (state.scratch.difference[pixel].re - state.scratch.field[pixel].norm_sqr())
1318 / (2.0 * distance)
1319 }
1320 (true, false) => {
1321 (state.scratch.difference[pixel].re - state.scratch.calibration_reference[pixel])
1322 / distance
1323 }
1324 (false, true) => {
1325 (state.scratch.calibration_reference[pixel] - state.scratch.field[pixel].norm_sqr())
1326 / distance
1327 }
1328 (false, false) => 0.0,
1329 };
1330 let intensity_derivative = configuration.source_weight * derivative;
1331 let loss_derivative = state.scratch.projected_field[pixel].re;
1332 let predicted = state.scratch.projected_field[pixel].im;
1333 gradient += loss_derivative * intensity_derivative;
1334 curvature += descent_curvature(
1335 predicted,
1336 loss_derivative,
1337 configuration.loss_type,
1338 configuration.epsilon,
1339 ) * intensity_derivative
1340 * intensity_derivative;
1341 }
1342 Ok(AxisGradient {
1343 gradient,
1344 curvature,
1345 })
1346}
1347
1348fn descent_curvature(
1349 predicted: f64,
1350 descent_factor: f64,
1351 loss_type: LossType,
1352 epsilon: f64,
1353) -> f64 {
1354 let mean = predicted.max(epsilon);
1355 match loss_type {
1356 LossType::AmplitudeMse => 0.5 / mean,
1357 LossType::IntensityMse => 2.0,
1358 LossType::PoissonNegativeLogLikelihood => 1.0 / mean,
1359 LossType::HuberAmplitude => {
1360 let clipped_residual = descent_factor * 2.0 * mean.sqrt();
1361 if clipped_residual.abs() < 1.0 {
1362 0.25 / mean
1363 } else {
1364 epsilon
1365 }
1366 }
1367 }
1368}
1369
1370fn background_value(
1371 state: &ReconstructionState,
1372 frame: usize,
1373 pixel: usize,
1374 image_len: usize,
1375) -> f64 {
1376 state.background.as_ref().map_or(0.0, |values| {
1377 values[if values.len() == image_len {
1378 pixel
1379 } else {
1380 frame * image_len + pixel
1381 }]
1382 })
1383}
1384
1385#[cfg(test)]
1386mod tests {
1387 use super::truncation_accepts;
1388
1389 #[test]
1390 fn poisson_truncation_gate_includes_its_boundary() {
1391 assert!(truncation_accepts(4.0, 10.0, 3.0));
1392 assert!(!truncation_accepts(4.0, 10.000_001, 3.0));
1393 assert!(truncation_accepts(0.0, 0.0, 1.0));
1394 assert!(!truncation_accepts(0.0, 1.0, 1.0));
1395 }
1396}