Skip to main content

fpm_rs/algorithms/
global_gauss_newton.rs

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/// Numerical work reported by one global Gauss–Newton object update.
18#[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    /// Returns the number of matrix-free conjugate-gradient iterations used.
29    pub fn conjugate_gradient_iterations(&self) -> usize {
30        self.conjugate_gradient_iterations
31    }
32
33    /// Returns the final linear residual norm divided by its initial norm.
34    pub fn linear_residual_ratio(&self) -> f64 {
35        self.linear_residual_ratio
36    }
37
38    /// Returns the number of full-data trial-objective evaluations.
39    pub fn line_search_evaluations(&self) -> usize {
40        self.line_search_evaluations
41    }
42
43    /// Returns the accepted multiplier applied to the Gauss–Newton direction.
44    pub fn accepted_step_scale(&self) -> f64 {
45        self.accepted_step_scale
46    }
47
48    /// Returns the Euclidean norm of `J^T r` at the pre-update object.
49    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/// Matrix-free damped Gauss–Newton reconstruction of a fixed-pupil FPM object.
92///
93/// # Method
94///
95/// The solver minimizes the full-stack, frame-weighted mean amplitude-MSE
96/// objective in intrinsic intensity units. Known detector gain and background
97/// are removed before forming residuals, masks omit pixels, and zero-weight
98/// frames contribute neither residuals nor derivatives. Incoherently
99/// multiplexed frames use one detector residual whose analytic derivative
100/// contains every contributing coherent source mode.
101///
102/// Each outer iteration linearizes the amplitude residual at the current
103/// centered object spectrum and solves
104/// `(J^T J + damping * diag(C)) d = -J^T r` with preconditioned conjugate
105/// gradients. `J` and `J^T` are applied analytically under the real inner
106/// product on complex arrays. `C` is a floored pupil-power Fourier-coverage
107/// approximation used for both damping and Jacobi preconditioning. Armijo
108/// backtracking accepts a step only when the same full-data objective
109/// decreases sufficiently.
110///
111/// Frames are processed for every gradient, normal-operator, and line-search
112/// evaluation. The implementation stores a fixed number of object-sized
113/// vectors and only the coherent fields of the current multiplexed frame; it
114/// never forms a Jacobian or Hessian and does not retain curvature across outer
115/// iterations. Consequently it supports lazy measurements and exact
116/// iteration-boundary checkpoint resume, but each iteration is substantially
117/// more expensive than a sequential projection or one gradient pass.
118///
119/// The pupil, frame response, generic source corrections, and physical model
120/// are treated as fixed during one step. The absence of persistent curvature
121/// permits use as the object phase of
122/// [`crate::algorithms::JointReconstruction`], whose subsequent physical phase
123/// recompiles the model before the next fresh linearization.
124///
125/// # Failure behavior
126///
127/// Configuration validation rejects non-positive limits or damping and
128/// tolerances outside their documented open intervals. A step also rejects an
129/// incomplete global batch or unrelated algorithm auxiliary state. Non-finite
130/// products, non-positive conjugate-gradient curvature, a non-descent
131/// direction, or an exhausted line search return an error without committing a
132/// trial object.
133///
134/// # Example
135///
136/// ```no_run
137/// use fpm_rs::{
138///     Result,
139///     algorithms::{GlobalGaussNewton, ReconstructionAlgorithm},
140///     measurements::MeasurementRead,
141///     reconstruction::ReconstructionProblem,
142/// };
143///
144/// # fn reconstruct<M: MeasurementRead>(problem: &ReconstructionProblem<M>) -> Result<()> {
145/// let result = GlobalGaussNewton::default()
146///     .iterations(10)
147///     .damping(1e-3)
148///     .maximum_cg_iterations(8)
149///     .run(problem)?;
150/// assert_eq!(result.trace.iterations.len(), 10);
151/// # Ok(())
152/// # }
153/// ```
154///
155/// # References
156///
157/// [L.-H. Yeh, J. Dong, J. Zhong, L. Tian, M. Chen, G. Tang,
158/// M. Soltanolkotabi, and L. Waller, “Experimental robustness of Fourier
159/// ptychography phase retrieval algorithms,” *Optics Express* **23**(26),
160/// 33214–33240 (2015)](https://doi.org/10.1364/OE.23.033214). That work forms
161/// exact CR-calculus Hessians for several objectives; this implementation keeps
162/// only the positive-semidefinite Gauss–Newton part of the amplitude residual
163/// and applies it without materializing a matrix.
164///
165/// Matrix-free second-order ptychographic optimization is also demonstrated by
166/// [S. Kandel, S. Maddali, Y. S. G. Nashed, S. O. Hruszkewycz, C. Jacobsen,
167/// and M. Allain, “Efficient ptychographic phase retrieval via a matrix-free
168/// Levenberg–Marquardt algorithm,” *Optics Express* **29**(15), 23019–23055
169/// (2021)](https://doi.org/10.1364/OE.422768). That work treats
170/// diffraction-plane ptychography with automatic differentiation, whereas
171/// this solver uses analytic products for the crate's image-plane FPM model.
172#[derive(Clone, Debug)]
173pub struct GlobalGaussNewton {
174    /// Number of complete global object updates.
175    pub iterations: usize,
176    /// Positive coverage-scaled diagonal damping coefficient.
177    pub damping: f64,
178    /// Maximum matrix-free conjugate-gradient iterations per outer update.
179    pub maximum_cg_iterations: usize,
180    /// Relative linear-residual tolerance for conjugate-gradient termination.
181    pub cg_relative_tolerance: f64,
182    /// Maximum full-data trial-objective evaluations per outer update.
183    pub maximum_line_search_steps: usize,
184    /// Multiplicative trial-step reduction in `(0, 1)`.
185    pub line_search_reduction: f64,
186    /// Armijo sufficient-decrease coefficient in `(0, 1)`.
187    pub line_search_sufficient_decrease: f64,
188    /// Positive floor for dark-field derivatives and Fourier coverage.
189    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    /// Sets the positive number of global object updates.
209    pub fn iterations(mut self, iterations: usize) -> Self {
210        self.iterations = iterations;
211        self
212    }
213
214    /// Sets the finite positive coverage-scaled damping coefficient.
215    pub fn damping(mut self, damping: f64) -> Self {
216        self.damping = damping;
217        self
218    }
219
220    /// Sets the positive maximum conjugate-gradient iteration count.
221    pub fn maximum_cg_iterations(mut self, iterations: usize) -> Self {
222        self.maximum_cg_iterations = iterations;
223        self
224    }
225
226    /// Sets the relative conjugate-gradient residual tolerance in `(0, 1)`.
227    pub fn cg_relative_tolerance(mut self, tolerance: f64) -> Self {
228        self.cg_relative_tolerance = tolerance;
229        self
230    }
231
232    /// Sets the positive maximum number of trial-objective evaluations.
233    pub fn maximum_line_search_steps(mut self, steps: usize) -> Self {
234        self.maximum_line_search_steps = steps;
235        self
236    }
237
238    /// Sets the multiplicative backtracking reduction in `(0, 1)`.
239    pub fn line_search_reduction(mut self, reduction: f64) -> Self {
240        self.line_search_reduction = reduction;
241        self
242    }
243
244    /// Sets the Armijo sufficient-decrease coefficient in `(0, 1)`.
245    pub fn line_search_sufficient_decrease(mut self, coefficient: f64) -> Self {
246        self.line_search_sufficient_decrease = coefficient;
247        self
248    }
249
250    /// Sets the finite positive numerical floor.
251    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}