1use ndarray::ArrayView2;
2use num_complex::Complex64;
3
4use crate::{
5 Result,
6 algorithms::{AlgorithmIterationMetrics, StepOutput, StepSummary},
7 array_layout::{StandardView2, checked_len_2d},
8 backend::FftDirection,
9 error::Error,
10 measurements::MeasurementRead,
11 model::{FourierOffset, ImagePlaneModel, fftshift_copy, ifftshift_copy},
12 reconstruction::{Batch, ReconstructionProblem, ReconstructionState},
13};
14
15use super::ReconstructionAlgorithm;
16
17#[derive(Clone, Copy, Debug, Default)]
19pub struct GlobalGaussNewtonIterationMetrics {
20 conjugate_gradient_iterations: usize,
21 linear_residual_ratio: f64,
22 line_search_evaluations: usize,
23 accepted_step_scale: f64,
24 gradient_norm: f64,
25}
26
27impl GlobalGaussNewtonIterationMetrics {
28 pub fn conjugate_gradient_iterations(&self) -> usize {
30 self.conjugate_gradient_iterations
31 }
32
33 pub fn linear_residual_ratio(&self) -> f64 {
35 self.linear_residual_ratio
36 }
37
38 pub fn line_search_evaluations(&self) -> usize {
40 self.line_search_evaluations
41 }
42
43 pub fn accepted_step_scale(&self) -> f64 {
45 self.accepted_step_scale
46 }
47
48 pub fn gradient_norm(&self) -> f64 {
50 self.gradient_norm
51 }
52}
53
54impl AlgorithmIterationMetrics for GlobalGaussNewtonIterationMetrics {
55 fn merge(&mut self, other: Self) {
56 self.conjugate_gradient_iterations += other.conjugate_gradient_iterations;
57 self.linear_residual_ratio = other.linear_residual_ratio;
58 self.line_search_evaluations += other.line_search_evaluations;
59 self.accepted_step_scale = other.accepted_step_scale;
60 self.gradient_norm = other.gradient_norm;
61 }
62
63 fn append_records(
64 &self,
65 iteration: usize,
66 output: &mut Vec<crate::reconstruction::AlgorithmMetricRecord>,
67 ) {
68 for (metric, value) in [
69 (
70 "conjugate_gradient_iterations",
71 self.conjugate_gradient_iterations as f64,
72 ),
73 ("linear_residual_ratio", self.linear_residual_ratio),
74 (
75 "line_search_evaluations",
76 self.line_search_evaluations as f64,
77 ),
78 ("accepted_step_scale", self.accepted_step_scale),
79 ("gradient_norm", self.gradient_norm),
80 ] {
81 output.push(crate::reconstruction::AlgorithmMetricRecord {
82 iteration,
83 namespace: "global_gauss_newton".into(),
84 metric: metric.into(),
85 value,
86 });
87 }
88 }
89}
90
91#[derive(Clone, Debug)]
173pub struct GlobalGaussNewton {
174 pub iterations: usize,
176 pub damping: f64,
178 pub maximum_cg_iterations: usize,
180 pub cg_relative_tolerance: f64,
182 pub maximum_line_search_steps: usize,
184 pub line_search_reduction: f64,
186 pub line_search_sufficient_decrease: f64,
188 pub epsilon: f64,
190}
191
192impl Default for GlobalGaussNewton {
193 fn default() -> Self {
194 Self {
195 iterations: 20,
196 damping: 1e-3,
197 maximum_cg_iterations: 12,
198 cg_relative_tolerance: 1e-3,
199 maximum_line_search_steps: 8,
200 line_search_reduction: 0.5,
201 line_search_sufficient_decrease: 1e-4,
202 epsilon: 1e-10,
203 }
204 }
205}
206
207impl GlobalGaussNewton {
208 pub fn iterations(mut self, iterations: usize) -> Self {
210 self.iterations = iterations;
211 self
212 }
213
214 pub fn damping(mut self, damping: f64) -> Self {
216 self.damping = damping;
217 self
218 }
219
220 pub fn maximum_cg_iterations(mut self, iterations: usize) -> Self {
222 self.maximum_cg_iterations = iterations;
223 self
224 }
225
226 pub fn cg_relative_tolerance(mut self, tolerance: f64) -> Self {
228 self.cg_relative_tolerance = tolerance;
229 self
230 }
231
232 pub fn maximum_line_search_steps(mut self, steps: usize) -> Self {
234 self.maximum_line_search_steps = steps;
235 self
236 }
237
238 pub fn line_search_reduction(mut self, reduction: f64) -> Self {
240 self.line_search_reduction = reduction;
241 self
242 }
243
244 pub fn line_search_sufficient_decrease(mut self, coefficient: f64) -> Self {
246 self.line_search_sufficient_decrease = coefficient;
247 self
248 }
249
250 pub fn epsilon(mut self, epsilon: f64) -> Self {
252 self.epsilon = epsilon;
253 self
254 }
255}
256
257impl ReconstructionAlgorithm for GlobalGaussNewton {
258 type IterationMetrics = GlobalGaussNewtonIterationMetrics;
259
260 fn validate(&self) -> Result<()> {
261 if self.iterations == 0 {
262 return Err(Error::InvalidParameter {
263 name: "iterations",
264 reason: "must be greater than zero".into(),
265 });
266 }
267 if !self.damping.is_finite() || self.damping <= 0.0 {
268 return Err(Error::InvalidParameter {
269 name: "damping",
270 reason: "must be finite and positive".into(),
271 });
272 }
273 if self.maximum_cg_iterations == 0 {
274 return Err(Error::InvalidParameter {
275 name: "maximum_cg_iterations",
276 reason: "must be greater than zero".into(),
277 });
278 }
279 if !self.cg_relative_tolerance.is_finite()
280 || self.cg_relative_tolerance <= 0.0
281 || self.cg_relative_tolerance >= 1.0
282 {
283 return Err(Error::InvalidParameter {
284 name: "cg_relative_tolerance",
285 reason: "must be finite and in (0, 1)".into(),
286 });
287 }
288 if self.maximum_line_search_steps == 0 {
289 return Err(Error::InvalidParameter {
290 name: "maximum_line_search_steps",
291 reason: "must be greater than zero".into(),
292 });
293 }
294 if !self.line_search_reduction.is_finite()
295 || self.line_search_reduction <= 0.0
296 || self.line_search_reduction >= 1.0
297 {
298 return Err(Error::InvalidParameter {
299 name: "line_search_reduction",
300 reason: "must be finite and in (0, 1)".into(),
301 });
302 }
303 if !self.line_search_sufficient_decrease.is_finite()
304 || self.line_search_sufficient_decrease <= 0.0
305 || self.line_search_sufficient_decrease >= 1.0
306 {
307 return Err(Error::InvalidParameter {
308 name: "line_search_sufficient_decrease",
309 reason: "must be finite and in (0, 1)".into(),
310 });
311 }
312 if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
313 return Err(Error::InvalidParameter {
314 name: "epsilon",
315 reason: "must be finite and positive".into(),
316 });
317 }
318 Ok(())
319 }
320
321 fn step<M: MeasurementRead>(
322 &mut self,
323 problem: &ReconstructionProblem<M>,
324 state: &mut ReconstructionState,
325 batch: &Batch,
326 _iteration: usize,
327 ) -> Result<StepOutput<Self::IterationMetrics>> {
328 validate_global_batch(problem.model.frame_count(), batch)?;
329 if state.algorithm_auxiliary.is_some() {
330 return Err(Error::InvalidModel(
331 "global Gauss–Newton cannot interpret state owned by another algorithm".into(),
332 ));
333 }
334
335 let object = state.object_spectrum.as_slice().to_vec();
336 let mut workspace = GaussNewtonWorkspace::new(&problem.model)?;
337 let coverage = coverage_diagonal(problem, state, self.epsilon, &mut workspace)?;
338 let (objective, current_summary, gradient) =
339 objective_and_gradient(problem, state, &object, self.epsilon, &mut workspace)?;
340 let gradient_norm = real_norm(&gradient);
341 if !gradient_norm.is_finite() {
342 return Err(Error::Numerical(
343 "global Gauss–Newton gradient norm is non-finite".into(),
344 ));
345 }
346 if gradient_norm <= self.epsilon {
347 return Ok(StepOutput {
348 summary: current_summary,
349 metrics: GlobalGaussNewtonIterationMetrics {
350 conjugate_gradient_iterations: 0,
351 linear_residual_ratio: 0.0,
352 line_search_evaluations: 0,
353 accepted_step_scale: 0.0,
354 gradient_norm,
355 },
356 });
357 }
358
359 let (direction, cg_iterations, linear_residual_ratio) = solve_direction(
360 problem,
361 state,
362 &object,
363 &gradient,
364 &coverage,
365 self,
366 &mut workspace,
367 )?;
368 let directional_derivative = real_dot(&gradient, &direction);
369 if !directional_derivative.is_finite() || directional_derivative >= 0.0 {
370 return Err(Error::Numerical(
371 "global Gauss–Newton produced a non-descent direction".into(),
372 ));
373 }
374
375 let mut step_scale = 1.0;
376 let mut accepted = None;
377 let mut evaluations = 0;
378 let mut candidate = vec![Complex64::default(); object.len()];
379 for _ in 0..self.maximum_line_search_steps {
380 evaluations += 1;
381 let mut finite = true;
382 for index in 0..object.len() {
383 candidate[index] = object[index] + step_scale * direction[index];
384 finite &= candidate[index].re.is_finite() && candidate[index].im.is_finite();
385 }
386 if finite {
387 match objective_only(problem, state, &candidate, &mut workspace) {
388 Ok((trial_objective, trial_summary)) => {
389 let armijo_bound = objective
390 + 2.0
391 * self.line_search_sufficient_decrease
392 * step_scale
393 * directional_derivative;
394 if trial_objective <= armijo_bound {
395 accepted = Some(trial_summary);
396 break;
397 }
398 }
399 Err(Error::Numerical(_)) => {}
400 Err(error) => return Err(error),
401 }
402 }
403 step_scale *= self.line_search_reduction;
404 }
405 let summary = accepted.ok_or_else(|| {
406 Error::Numerical(format!(
407 "global Gauss–Newton line search failed after {} evaluations",
408 self.maximum_line_search_steps
409 ))
410 })?;
411 state.object_spectrum.as_slice_mut().copy_from_slice(&candidate);
412 state.object_real_space_cache = None;
413
414 Ok(StepOutput {
415 summary,
416 metrics: GlobalGaussNewtonIterationMetrics {
417 conjugate_gradient_iterations: cg_iterations,
418 linear_residual_ratio,
419 line_search_evaluations: evaluations,
420 accepted_step_scale: step_scale,
421 gradient_norm,
422 },
423 })
424 }
425
426 fn iterations(&self) -> usize {
427 self.iterations
428 }
429
430 fn batch_size(&self) -> usize {
431 usize::MAX
432 }
433}
434
435struct GaussNewtonWorkspace {
436 patch: Vec<Complex64>,
437 centered: Vec<Complex64>,
438 field: Vec<Complex64>,
439 detector: Vec<Complex64>,
440 mode_fields: Vec<Complex64>,
441 predicted: Vec<f64>,
442 denominator: Vec<f64>,
443 residual: Vec<f64>,
444 directional: Vec<f64>,
445 column: Vec<Complex64>,
446}
447
448impl GaussNewtonWorkspace {
449 fn new(model: &ImagePlaneModel) -> Result<Self> {
450 let low_len = checked_len_2d(model.image_shape)?;
451 Ok(Self {
452 patch: vec![Complex64::default(); low_len],
453 centered: vec![Complex64::default(); low_len],
454 field: vec![Complex64::default(); low_len],
455 detector: vec![Complex64::default(); low_len],
456 mode_fields: Vec::new(),
457 predicted: vec![0.0; low_len],
458 denominator: vec![0.0; low_len],
459 residual: vec![0.0; low_len],
460 directional: vec![0.0; low_len],
461 column: vec![
462 Complex64::default();
463 model.image_shape.0.max(model.reconstruction_shape.0)
464 ],
465 })
466 }
467
468 fn resize_modes(&mut self, modes: usize, low_len: usize) -> Result<()> {
469 let length = modes
470 .checked_mul(low_len)
471 .ok_or_else(|| Error::InvalidShape("multiplexed field storage overflows".into()))?;
472 self.mode_fields.resize(length, Complex64::default());
473 Ok(())
474 }
475}
476
477fn validate_global_batch(frame_count: usize, batch: &Batch) -> Result<()> {
478 if batch.indices.len() != frame_count {
479 return Err(Error::InvalidParameter {
480 name: "batch",
481 reason: format!(
482 "global Gauss–Newton requires all {frame_count} frames in one step"
483 ),
484 });
485 }
486 let mut seen = vec![false; frame_count];
487 for &frame in &batch.indices {
488 if frame >= frame_count || seen[frame] {
489 return Err(Error::InvalidParameter {
490 name: "batch",
491 reason: "must contain every frame exactly once".into(),
492 });
493 }
494 seen[frame] = true;
495 }
496 Ok(())
497}
498
499fn positive_weight_sum<M: MeasurementRead>(problem: &ReconstructionProblem<M>) -> Result<f64> {
500 let mut total = 0.0;
501 for frame in 0..problem.model.frame_count() {
502 total += problem.measurements.frame_weight(frame)?;
503 }
504 if !total.is_finite() || total <= 0.0 {
505 return Err(Error::InvalidMeasurements(
506 "global Gauss–Newton requires positive finite frame weight".into(),
507 ));
508 }
509 Ok(total)
510}
511
512fn valid_pixel_count<M: MeasurementRead>(
513 problem: &ReconstructionProblem<M>,
514 frame: usize,
515) -> Result<usize> {
516 let mask = problem.measurements.frame_mask(frame)?;
517 let count = mask.map_or(problem.measurements.frame_len(), |values| {
518 values.iter().filter(|&&value| value != 0).count()
519 });
520 if count == 0 {
521 return Err(Error::InvalidMeasurements(format!(
522 "positive-weight frame {frame} has no unmasked pixels"
523 )));
524 }
525 Ok(count)
526}
527
528fn coverage_diagonal<M: MeasurementRead>(
529 problem: &ReconstructionProblem<M>,
530 state: &ReconstructionState,
531 epsilon: f64,
532 workspace: &mut GaussNewtonWorkspace,
533) -> Result<Vec<f64>> {
534 let model = &problem.model;
535 let total_weight = positive_weight_sum(problem)?;
536 let mut coverage = vec![Complex64::default(); state.object_spectrum.len()];
537 for pixel in 0..workspace.centered.len() {
538 workspace.centered[pixel] =
539 Complex64::new(state.pupil.values.as_slice()[pixel].norm_sqr(), 0.0);
540 }
541 for frame in 0..model.frame_count() {
542 let frame_weight = problem.measurements.frame_weight(frame)?;
543 if frame_weight == 0.0 {
544 continue;
545 }
546 let valid_pixels = valid_pixel_count(problem, frame)? as f64;
547 let single_source = [(frame, 1.0)];
548 let sources = model
549 .multiplexing_matrix
550 .as_ref()
551 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
552 for &(source, source_weight) in sources {
553 let offset = state.effective_source_offset(model, source)?;
554 model.insert_patch_adjoint_slice_at_offset(
555 &mut coverage,
556 source,
557 &workspace.centered,
558 frame_weight * source_weight / (valid_pixels * total_weight),
559 offset,
560 )?;
561 }
562 }
563 let maximum = coverage
564 .iter()
565 .map(|value| value.re)
566 .fold(0.0_f64, f64::max);
567 if !maximum.is_finite() || maximum <= 0.0 {
568 return Err(Error::InvalidModel(
569 "global Gauss–Newton Fourier coverage is empty or non-finite".into(),
570 ));
571 }
572 Ok(coverage
573 .into_iter()
574 .map(|value| (value.re / maximum).max(epsilon))
575 .collect())
576}
577
578fn objective_and_gradient<M: MeasurementRead>(
579 problem: &ReconstructionProblem<M>,
580 state: &ReconstructionState,
581 object: &[Complex64],
582 epsilon: f64,
583 workspace: &mut GaussNewtonWorkspace,
584) -> Result<(f64, StepSummary, Vec<Complex64>)> {
585 let model = &problem.model;
586 let low_len = checked_len_2d(model.image_shape)?;
587 let total_weight = positive_weight_sum(problem)?;
588 let mut gradient = vec![Complex64::default(); object.len()];
589 let mut summary = StepSummary::default();
590 for frame in 0..model.frame_count() {
591 let frame_weight = problem.measurements.frame_weight(frame)?;
592 if frame_weight == 0.0 {
593 summary.push_frame(frame, 0.0, 0.0);
594 continue;
595 }
596 let valid_pixels = valid_pixel_count(problem, frame)?;
597 let single_source = [(frame, 1.0)];
598 let sources = model
599 .multiplexing_matrix
600 .as_ref()
601 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
602 predict_frame(problem, state, object, sources, workspace)?;
603
604 let measured = problem.measurements.frame(frame)?;
605 let mask = problem.measurements.frame_mask(frame)?;
606 let gain = frame_gain(state, frame)?;
607 let residual_scale =
608 (frame_weight / (valid_pixels as f64 * total_weight)).sqrt();
609 let mut frame_loss = 0.0;
610 for pixel in 0..low_len {
611 if mask.is_some_and(|values| values[pixel] == 0) {
612 workspace.residual[pixel] = 0.0;
613 workspace.denominator[pixel] = epsilon.sqrt();
614 continue;
615 }
616 let target = ((measured[pixel] - background_value(state, frame, pixel, low_len))
617 / gain)
618 .max(0.0);
619 let predicted_amplitude = workspace.predicted[pixel].max(0.0).sqrt();
620 let residual = predicted_amplitude - target.sqrt();
621 frame_loss += residual * residual;
622 workspace.residual[pixel] = residual_scale * residual;
623 workspace.denominator[pixel] = predicted_amplitude.max(epsilon.sqrt());
624 }
625 summary.push_frame(frame, frame_loss / valid_pixels as f64, frame_weight);
626
627 for (mode, &(source, source_weight)) in sources.iter().enumerate() {
628 let start = mode * low_len;
629 let mode_field = &workspace.mode_fields[start..start + low_len];
630 for pixel in 0..low_len {
631 workspace.detector[pixel] = if mask.is_some_and(|values| values[pixel] == 0) {
632 Complex64::default()
633 } else {
634 mode_field[pixel]
635 * (source_weight * residual_scale * workspace.residual[pixel]
636 / workspace.denominator[pixel])
637 };
638 }
639 let offset = state.effective_source_offset(model, source)?;
640 adjoint_source(
641 model,
642 state,
643 source,
644 offset,
645 &workspace.detector,
646 &mut gradient,
647 &mut workspace.field,
648 &mut workspace.centered,
649 &mut workspace.column,
650 )?;
651 }
652 }
653 let objective = summary.mean_objective().ok_or_else(|| {
654 Error::InvalidMeasurements("global objective has no positive frame weight".into())
655 })?;
656 if !objective.is_finite() || gradient.iter().any(|v| !complex_is_finite(*v)) {
657 return Err(Error::Numerical(
658 "global Gauss–Newton objective or gradient is non-finite".into(),
659 ));
660 }
661 Ok((objective, summary, gradient))
662}
663
664fn objective_only<M: MeasurementRead>(
665 problem: &ReconstructionProblem<M>,
666 state: &ReconstructionState,
667 object: &[Complex64],
668 workspace: &mut GaussNewtonWorkspace,
669) -> Result<(f64, StepSummary)> {
670 let model = &problem.model;
671 let low_len = checked_len_2d(model.image_shape)?;
672 let mut summary = StepSummary::default();
673 for frame in 0..model.frame_count() {
674 let frame_weight = problem.measurements.frame_weight(frame)?;
675 if frame_weight == 0.0 {
676 summary.push_frame(frame, 0.0, 0.0);
677 continue;
678 }
679 let valid_pixels = valid_pixel_count(problem, frame)?;
680 let single_source = [(frame, 1.0)];
681 let sources = model
682 .multiplexing_matrix
683 .as_ref()
684 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
685 predict_frame(problem, state, object, sources, workspace)?;
686 let measured = problem.measurements.frame(frame)?;
687 let mask = problem.measurements.frame_mask(frame)?;
688 let gain = frame_gain(state, frame)?;
689 let mut frame_loss = 0.0;
690 for pixel in 0..low_len {
691 if mask.is_some_and(|values| values[pixel] == 0) {
692 continue;
693 }
694 let target = ((measured[pixel] - background_value(state, frame, pixel, low_len))
695 / gain)
696 .max(0.0);
697 let residual = workspace.predicted[pixel].max(0.0).sqrt() - target.sqrt();
698 frame_loss += residual * residual;
699 }
700 summary.push_frame(frame, frame_loss / valid_pixels as f64, frame_weight);
701 }
702 let objective = summary.mean_objective().ok_or_else(|| {
703 Error::InvalidMeasurements("global objective has no positive frame weight".into())
704 })?;
705 if !objective.is_finite() {
706 return Err(Error::Numerical(
707 "global Gauss–Newton trial objective is non-finite".into(),
708 ));
709 }
710 Ok((objective, summary))
711}
712
713#[allow(clippy::too_many_arguments)]
714fn solve_direction<M: MeasurementRead>(
715 problem: &ReconstructionProblem<M>,
716 state: &ReconstructionState,
717 object: &[Complex64],
718 gradient: &[Complex64],
719 coverage: &[f64],
720 algorithm: &GlobalGaussNewton,
721 workspace: &mut GaussNewtonWorkspace,
722) -> Result<(Vec<Complex64>, usize, f64)> {
723 let length = object.len();
724 let mut solution = vec![Complex64::default(); length];
725 let mut residual: Vec<_> = gradient.iter().map(|value| -*value).collect();
726 let initial_norm = real_norm(&residual);
727 if initial_norm == 0.0 {
728 return Ok((solution, 0, 0.0));
729 }
730 let mut preconditioned = vec![Complex64::default(); length];
731 apply_preconditioner(
732 &residual,
733 coverage,
734 algorithm.damping,
735 &mut preconditioned,
736 );
737 let mut direction = preconditioned.clone();
738 let mut residual_product = real_dot(&residual, &preconditioned);
739 if !residual_product.is_finite() || residual_product <= 0.0 {
740 return Err(Error::Numerical(
741 "global Gauss–Newton preconditioned residual is not positive".into(),
742 ));
743 }
744 let mut ratio = 1.0;
745 let mut completed = 0;
746 for iteration in 0..algorithm.maximum_cg_iterations {
747 let operator_direction = apply_normal_operator(
748 problem,
749 state,
750 object,
751 &direction,
752 coverage,
753 algorithm.damping,
754 algorithm.epsilon,
755 workspace,
756 )?;
757 let curvature = real_dot(&direction, &operator_direction);
758 if !curvature.is_finite() || curvature <= 0.0 {
759 return Err(Error::Numerical(
760 "global Gauss–Newton conjugate-gradient curvature is not positive".into(),
761 ));
762 }
763 let step = residual_product / curvature;
764 if !step.is_finite() {
765 return Err(Error::Numerical(
766 "global Gauss–Newton conjugate-gradient step is non-finite".into(),
767 ));
768 }
769 for index in 0..length {
770 solution[index] += step * direction[index];
771 residual[index] -= step * operator_direction[index];
772 }
773 completed = iteration + 1;
774 ratio = real_norm(&residual) / initial_norm;
775 if !ratio.is_finite() {
776 return Err(Error::Numerical(
777 "global Gauss–Newton linear residual is non-finite".into(),
778 ));
779 }
780 if ratio <= algorithm.cg_relative_tolerance {
781 break;
782 }
783 apply_preconditioner(
784 &residual,
785 coverage,
786 algorithm.damping,
787 &mut preconditioned,
788 );
789 let next_product = real_dot(&residual, &preconditioned);
790 if !next_product.is_finite() || next_product <= 0.0 {
791 return Err(Error::Numerical(
792 "global Gauss–Newton conjugate-gradient residual broke down".into(),
793 ));
794 }
795 let beta = next_product / residual_product;
796 for index in 0..length {
797 direction[index] = preconditioned[index] + beta * direction[index];
798 }
799 residual_product = next_product;
800 }
801 if solution.iter().any(|value| !complex_is_finite(*value)) {
802 return Err(Error::Numerical(
803 "global Gauss–Newton direction is non-finite".into(),
804 ));
805 }
806 Ok((solution, completed, ratio))
807}
808
809fn apply_preconditioner(
810 input: &[Complex64],
811 coverage: &[f64],
812 damping: f64,
813 output: &mut [Complex64],
814) {
815 for index in 0..input.len() {
816 output[index] = input[index] / ((1.0 + damping) * coverage[index]);
817 }
818}
819
820#[allow(clippy::too_many_arguments)]
821fn apply_normal_operator<M: MeasurementRead>(
822 problem: &ReconstructionProblem<M>,
823 state: &ReconstructionState,
824 object: &[Complex64],
825 vector: &[Complex64],
826 coverage: &[f64],
827 damping: f64,
828 epsilon: f64,
829 workspace: &mut GaussNewtonWorkspace,
830) -> Result<Vec<Complex64>> {
831 let model = &problem.model;
832 let low_len = checked_len_2d(model.image_shape)?;
833 let total_weight = positive_weight_sum(problem)?;
834 let mut output = vec![Complex64::default(); object.len()];
835 for frame in 0..model.frame_count() {
836 let frame_weight = problem.measurements.frame_weight(frame)?;
837 if frame_weight == 0.0 {
838 continue;
839 }
840 let valid_pixels = valid_pixel_count(problem, frame)?;
841 let mask = problem.measurements.frame_mask(frame)?;
842 let residual_scale =
843 (frame_weight / (valid_pixels as f64 * total_weight)).sqrt();
844 let single_source = [(frame, 1.0)];
845 let sources = model
846 .multiplexing_matrix
847 .as_ref()
848 .map_or(single_source.as_slice(), |matrix| matrix[frame].as_slice());
849 predict_frame(problem, state, object, sources, workspace)?;
850 workspace.directional.fill(0.0);
851 for (mode, &(source, source_weight)) in sources.iter().enumerate() {
852 let offset = state.effective_source_offset(model, source)?;
853 forward_source(
854 model,
855 state,
856 vector,
857 source,
858 offset,
859 &mut workspace.patch,
860 &mut workspace.centered,
861 &mut workspace.field,
862 &mut workspace.column,
863 )?;
864 let start = mode * low_len;
865 let mode_field = &workspace.mode_fields[start..start + low_len];
866 for pixel in 0..low_len {
867 workspace.directional[pixel] += source_weight
868 * (mode_field[pixel].conj() * workspace.field[pixel]).re;
869 }
870 }
871 for pixel in 0..low_len {
872 let amplitude = workspace.predicted[pixel].max(0.0).sqrt();
873 workspace.denominator[pixel] = amplitude.max(epsilon.sqrt());
874 workspace.residual[pixel] = if mask.is_some_and(|values| values[pixel] == 0) {
875 0.0
876 } else {
877 residual_scale * workspace.directional[pixel] / workspace.denominator[pixel]
878 };
879 }
880 for (mode, &(source, source_weight)) in sources.iter().enumerate() {
881 let start = mode * low_len;
882 let mode_field = &workspace.mode_fields[start..start + low_len];
883 for pixel in 0..low_len {
884 workspace.detector[pixel] = if mask.is_some_and(|values| values[pixel] == 0) {
885 Complex64::default()
886 } else {
887 mode_field[pixel]
888 * (source_weight * residual_scale * workspace.residual[pixel]
889 / workspace.denominator[pixel])
890 };
891 }
892 let offset = state.effective_source_offset(model, source)?;
893 adjoint_source(
894 model,
895 state,
896 source,
897 offset,
898 &workspace.detector,
899 &mut output,
900 &mut workspace.field,
901 &mut workspace.centered,
902 &mut workspace.column,
903 )?;
904 }
905 }
906 for index in 0..output.len() {
907 output[index] += damping * coverage[index] * vector[index];
908 }
909 if output.iter().any(|value| !complex_is_finite(*value)) {
910 return Err(Error::Numerical(
911 "global Gauss–Newton normal-operator product is non-finite".into(),
912 ));
913 }
914 Ok(output)
915}
916
917fn predict_frame<M: MeasurementRead>(
918 problem: &ReconstructionProblem<M>,
919 state: &ReconstructionState,
920 object: &[Complex64],
921 sources: &[(usize, f64)],
922 workspace: &mut GaussNewtonWorkspace,
923) -> Result<()> {
924 let low_len = checked_len_2d(problem.model.image_shape)?;
925 workspace.resize_modes(sources.len(), low_len)?;
926 workspace.predicted.fill(0.0);
927 for (mode, &(source, source_weight)) in sources.iter().enumerate() {
928 let offset = state.effective_source_offset(&problem.model, source)?;
929 forward_source(
930 &problem.model,
931 state,
932 object,
933 source,
934 offset,
935 &mut workspace.patch,
936 &mut workspace.centered,
937 &mut workspace.field,
938 &mut workspace.column,
939 )?;
940 let start = mode * low_len;
941 workspace.mode_fields[start..start + low_len].copy_from_slice(&workspace.field);
942 for pixel in 0..low_len {
943 workspace.predicted[pixel] += source_weight * workspace.field[pixel].norm_sqr();
944 }
945 }
946 if workspace
947 .predicted
948 .iter()
949 .any(|value| !value.is_finite() || *value < 0.0)
950 {
951 return Err(Error::Numerical(
952 "global Gauss–Newton forward prediction is non-finite".into(),
953 ));
954 }
955 Ok(())
956}
957
958#[allow(clippy::too_many_arguments)]
959fn forward_source(
960 model: &ImagePlaneModel,
961 state: &ReconstructionState,
962 object: &[Complex64],
963 source: usize,
964 offset: FourierOffset,
965 patch: &mut [Complex64],
966 centered: &mut [Complex64],
967 field: &mut [Complex64],
968 column: &mut [Complex64],
969) -> Result<()> {
970 let view = StandardView2::try_from(ArrayView2::from_shape(
971 model.reconstruction_shape,
972 object,
973 )?)?;
974 model.extract_patch_at_offset(view, source, offset, patch)?;
975 for pixel in 0..patch.len() {
976 centered[pixel] = patch[pixel] * state.pupil.values.as_slice()[pixel];
977 }
978 ifftshift_copy(centered, field, model.image_shape);
979 state
980 .backend
981 .fft2(field, model.image_shape, FftDirection::Inverse, column)
982}
983
984#[allow(clippy::too_many_arguments)]
985fn adjoint_source(
986 model: &ImagePlaneModel,
987 state: &ReconstructionState,
988 source: usize,
989 offset: FourierOffset,
990 detector: &[Complex64],
991 destination: &mut [Complex64],
992 field: &mut [Complex64],
993 centered: &mut [Complex64],
994 column: &mut [Complex64],
995) -> Result<()> {
996 field.copy_from_slice(detector);
997 state
998 .backend
999 .fft2(field, model.image_shape, FftDirection::Forward, column)?;
1000 fftshift_copy(field, centered, model.image_shape);
1001 for pixel in 0..centered.len() {
1002 centered[pixel] *= state.pupil.values.as_slice()[pixel].conj();
1003 }
1004 model.insert_patch_adjoint_slice_at_offset(
1005 destination,
1006 source,
1007 centered,
1008 checked_len_2d(model.image_shape)? as f64,
1009 offset,
1010 )
1011}
1012
1013fn frame_gain(state: &ReconstructionState, frame: usize) -> Result<f64> {
1014 let gain = state.frame_gains.as_ref().map_or(1.0, |values| values[frame]);
1015 if !gain.is_finite() || gain <= 0.0 {
1016 return Err(Error::InvalidModel(format!(
1017 "state frame {frame} has invalid gain {gain}"
1018 )));
1019 }
1020 Ok(gain)
1021}
1022
1023fn background_value(
1024 state: &ReconstructionState,
1025 frame: usize,
1026 pixel: usize,
1027 image_len: usize,
1028) -> f64 {
1029 state.background.as_ref().map_or(0.0, |values| {
1030 values[if values.len() == image_len {
1031 pixel
1032 } else {
1033 frame * image_len + pixel
1034 }]
1035 })
1036}
1037
1038fn real_dot(left: &[Complex64], right: &[Complex64]) -> f64 {
1039 left.iter()
1040 .zip(right)
1041 .map(|(&left, &right)| (left.conj() * right).re)
1042 .sum()
1043}
1044
1045fn real_norm(values: &[Complex64]) -> f64 {
1046 values.iter().map(|value| value.norm_sqr()).sum::<f64>().sqrt()
1047}
1048
1049fn complex_is_finite(value: Complex64) -> bool {
1050 value.re.is_finite() && value.im.is_finite()
1051}
1052
1053#[cfg(test)]
1054mod tests {
1055 use super::*;
1056 use crate::simulation::presets::noiseless_mixed_fpm;
1057
1058 fn deterministic_vector(length: usize, phase: usize) -> Vec<Complex64> {
1059 (0..length)
1060 .map(|index| {
1061 let real = ((index + 3 * phase) % 17) as f64 - 8.0;
1062 let imaginary = ((5 * index + phase) % 19) as f64 - 9.0;
1063 Complex64::new(real / 17.0, imaginary / 19.0)
1064 })
1065 .collect()
1066 }
1067
1068 #[test]
1069 fn compiled_source_forward_and_adjoint_obey_the_real_dot_product() {
1070 let simulation = noiseless_mixed_fpm(17).unwrap();
1071 let problem =
1072 ReconstructionProblem::new(simulation.measurements, simulation.reconstruction_model)
1073 .unwrap();
1074 let state = ReconstructionState::initialize(&problem).unwrap();
1075 let mut workspace = GaussNewtonWorkspace::new(&problem.model).unwrap();
1076 let object = deterministic_vector(state.object_spectrum.len(), 1);
1077 let detector = deterministic_vector(workspace.field.len(), 2);
1078 let offset = state.effective_source_offset(&problem.model, 0).unwrap();
1079 forward_source(
1080 &problem.model,
1081 &state,
1082 &object,
1083 0,
1084 offset,
1085 &mut workspace.patch,
1086 &mut workspace.centered,
1087 &mut workspace.field,
1088 &mut workspace.column,
1089 )
1090 .unwrap();
1091 let field = workspace.field.clone();
1092 let mut adjoint = vec![Complex64::default(); object.len()];
1093 adjoint_source(
1094 &problem.model,
1095 &state,
1096 0,
1097 offset,
1098 &detector,
1099 &mut adjoint,
1100 &mut workspace.field,
1101 &mut workspace.centered,
1102 &mut workspace.column,
1103 )
1104 .unwrap();
1105
1106 let forward_dot = real_dot(&field, &detector);
1107 let adjoint_dot = real_dot(&object, &adjoint);
1108 let scale = forward_dot.abs().max(adjoint_dot.abs()).max(1.0);
1109 assert!((forward_dot - adjoint_dot).abs() <= 1e-11 * scale);
1110 }
1111
1112 #[test]
1113 fn analytic_global_gradient_matches_a_centered_objective_difference() {
1114 let simulation = noiseless_mixed_fpm(23).unwrap();
1115 let problem =
1116 ReconstructionProblem::new(simulation.measurements, simulation.reconstruction_model)
1117 .unwrap();
1118 let state = ReconstructionState::initialize(&problem).unwrap();
1119 let object = state.object_spectrum.as_slice().to_vec();
1120 let direction = deterministic_vector(object.len(), 3);
1121 let mut workspace = GaussNewtonWorkspace::new(&problem.model).unwrap();
1122 let (_, _, gradient) =
1123 objective_and_gradient(&problem, &state, &object, 1e-10, &mut workspace).unwrap();
1124 let step = 1e-6;
1125 let plus: Vec<_> = object
1126 .iter()
1127 .zip(&direction)
1128 .map(|(&value, &delta)| value + step * delta)
1129 .collect();
1130 let minus: Vec<_> = object
1131 .iter()
1132 .zip(&direction)
1133 .map(|(&value, &delta)| value - step * delta)
1134 .collect();
1135 let plus_objective = objective_only(&problem, &state, &plus, &mut workspace)
1136 .unwrap()
1137 .0;
1138 let minus_objective = objective_only(&problem, &state, &minus, &mut workspace)
1139 .unwrap()
1140 .0;
1141 let finite_difference = (plus_objective - minus_objective) / (2.0 * step);
1142 let analytic = 2.0 * real_dot(&gradient, &direction);
1143 let scale = finite_difference.abs().max(analytic.abs()).max(1.0);
1144 assert!((finite_difference - analytic).abs() <= 5e-5 * scale);
1145 }
1146
1147 #[test]
1148 fn damped_normal_operator_is_real_symmetric_and_positive() {
1149 let simulation = noiseless_mixed_fpm(31).unwrap();
1150 let problem =
1151 ReconstructionProblem::new(simulation.measurements, simulation.reconstruction_model)
1152 .unwrap();
1153 let state = ReconstructionState::initialize(&problem).unwrap();
1154 let object = state.object_spectrum.as_slice().to_vec();
1155 let left = deterministic_vector(object.len(), 4);
1156 let right = deterministic_vector(object.len(), 5);
1157 let coverage = vec![1.0; object.len()];
1158 let mut workspace = GaussNewtonWorkspace::new(&problem.model).unwrap();
1159 let normal_left = apply_normal_operator(
1160 &problem,
1161 &state,
1162 &object,
1163 &left,
1164 &coverage,
1165 1e-3,
1166 1e-10,
1167 &mut workspace,
1168 )
1169 .unwrap();
1170 let normal_right = apply_normal_operator(
1171 &problem,
1172 &state,
1173 &object,
1174 &right,
1175 &coverage,
1176 1e-3,
1177 1e-10,
1178 &mut workspace,
1179 )
1180 .unwrap();
1181
1182 let left_right = real_dot(&left, &normal_right);
1183 let right_left = real_dot(&normal_left, &right);
1184 let scale = left_right.abs().max(right_left.abs()).max(1.0);
1185 assert!((left_right - right_left).abs() <= 1e-10 * scale);
1186 assert!(real_dot(&left, &normal_left) > 0.0);
1187 assert!(real_dot(&right, &normal_right) > 0.0);
1188 }
1189}