Skip to main content

fpm_rs/algorithms/
gradient_descent.rs

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/// Truncation statistics emitted by one gradient-descent iteration.
25///
26/// The retained fraction is frame-weighted and includes only unmasked pixels
27/// from positive-weight frames. It is absent when Poisson truncation is
28/// disabled or when an iteration has no eligible pixels.
29#[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    /// Returns the frame-weighted fraction of eligible pixels retained by truncation.
38    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/// Wirtinger-style loss-gradient reconstruction for Fourier ptychography.
68///
69/// # Method
70///
71/// The solver differentiates a selected real-valued data loss through the
72/// complex FPM forward model, accumulates gradients from a mini-batch, and
73/// applies a pupil-power-preconditioned update to the shared object spectrum.
74/// This is the Fourier-ptychographic Wirtinger-flow viewpoint: phase retrieval
75/// is treated as direct optimization rather than alternating hard projections.
76/// Losses are evaluated after accounting for known linear gain and background,
77/// so detector count scaling does not change the intrinsic update scale.
78/// With `poisson_truncation_threshold` enabled, a pre-update mini-batch pass
79/// rejects signal-dependent intensity outliers from the gradient while the
80/// reported Poisson objective continues to include every valid pixel.
81///
82/// Optional extensions recover the pupil with an analogous normalized
83/// gradient, estimate illumination offsets with finite-difference derivatives
84/// and diagonal Gauss–Newton scaling, and regularize the complex object or
85/// pupil. Incoherent multiplexing, these calibration updates, and the selectable
86/// losses extend the reference formulation.
87///
88/// # Poisson truncation
89///
90/// For each mini-batch, the implementation computes the frame-weighted mean
91/// absolute residual `R` over unmasked pixels in positive-weight frames. A
92/// pixel with intrinsic target `y`, total predicted intensity `p`, and object
93/// RMS amplitude `z_rms` contributes exactly when
94/// `|y - p| <= alpha * R * sqrt(p) / max(z_rms, sqrt(epsilon))`. Known gain and
95/// background are removed before this test. One gate is shared by every mode
96/// of an incoherently multiplexed pixel and by object, pupil, and illumination
97/// gradients. The statistic is mini-batch-local rather than full-data, the
98/// object step is fixed rather than scheduled, and pupil and illumination
99/// recovery are implementation extensions beyond the cited object-only method.
100///
101/// # Gauge convention
102///
103/// With pupil recovery enabled, the compiled pupil fixes the reported joint
104/// object/pupil gauge at iteration boundaries. The projection matches supported
105/// pupil energy and piston, fixes the remaining object piston, and removes an
106/// affine pupil phase ramp only on axes with zero effective subpixel offsets.
107/// Reciprocal object corrections preserve predicted intensities. Fractional
108/// axes retain their affine phase because bilinear crop interpolation does not
109/// commute exactly with a discrete phase ramp. Canonicalization runs before
110/// iteration callbacks, checkpoints, and final results, including after resume.
111///
112/// # References
113///
114/// [L. Bian, J. Suo, G. Zheng, K. Guo, F. Chen, and Q. Dai, “Fourier
115/// ptychographic reconstruction using Wirtinger flow optimization”
116/// (2015)](https://doi.org/10.1364/OE.23.004856), *Optics Express* **23**(4),
117/// 4856–4866.
118///
119/// [L. Bian, J. Suo, J. Chung, X. Ou, C. Yang, F. Chen, and Q. Dai, “Fourier
120/// ptychographic reconstruction using Poisson maximum likelihood and truncated
121/// Wirtinger gradient” (2016)](https://doi.org/10.1038/srep27384), *Scientific
122/// Reports* **6**, 27384.
123///
124/// The blind object/pupil ambiguities follow [A. Fannjiang and P. Chen, “Blind
125/// ptychography: uniqueness and ambiguities” (2020)](https://doi.org/10.1088/1361-6420/ab6504),
126/// *Inverse Problems* **36**, 045005; this implementation uses a Fourier-domain
127/// object and retains affine phase on fractionally interpolated axes.
128#[derive(Clone, Debug)]
129pub struct GradientDescent {
130    /// Number of complete passes through the acquisition schedule.
131    pub iterations: usize,
132    /// Step size of the pupil-power-preconditioned object update.
133    pub object_step: f64,
134    /// Number of frame gradients averaged into one update.
135    pub batch_size: usize,
136    /// Positive numerical floor used by losses and preconditioners.
137    pub epsilon: f64,
138    /// Data-fidelity objective to differentiate and report.
139    pub loss_type: LossType,
140    /// Optional positive signal-dependent Poisson truncation coefficient.
141    ///
142    /// `None` uses the ordinary untruncated gradient. The cited TPWFP work used
143    /// `25`; truncation requires [`LossType::PoissonNegativeLogLikelihood`].
144    pub poisson_truncation_threshold: Option<f64>,
145    /// Whether to estimate a Fourier-grid offset for every illumination source.
146    pub recover_illumination: bool,
147    /// Step size of the diagonally scaled illumination-offset update.
148    pub illumination_step: f64,
149    /// Central finite-difference spacing in Fourier-grid pixels.
150    pub illumination_finite_difference: f64,
151    /// Maximum absolute row or column correction, in Fourier-grid pixels.
152    pub maximum_illumination_correction: f64,
153    /// Whether to update the complex pupil alongside the object.
154    pub recover_pupil: bool,
155    /// Step size of the normalized pupil update.
156    pub pupil_step: f64,
157    /// Whether to zero recovered pupil values outside the compiled aperture.
158    pub constrain_pupil_support: bool,
159    /// Weight of the isotropic total-variation step on the complex object;
160    /// `0` disables it.
161    pub object_tv_weight: f64,
162    /// Positive smoothing constant in the differentiable TV norm.
163    pub object_tv_epsilon: f64,
164    /// Weight of quadratic nearest-neighbor pupil smoothing; `0` disables it
165    /// and positive values require pupil recovery.
166    pub pupil_smoothing_weight: f64,
167    /// Maximum number of frame-gradient worker threads.
168    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    /// Sets the number of complete acquisition-schedule passes; validation requires non-zero.
197    pub fn iterations(mut self, iterations: usize) -> Self {
198        self.iterations = iterations;
199        self
200    }
201
202    /// Sets the finite positive step size of the preconditioned object update.
203    pub fn object_step(mut self, step: f64) -> Self {
204        self.object_step = step;
205        self
206    }
207
208    /// Sets the positive number of frame gradients averaged into one update.
209    pub fn batch_size(mut self, batch_size: usize) -> Self {
210        self.batch_size = batch_size;
211        self
212    }
213
214    /// Selects the differentiable data-fidelity objective used for updates and reporting.
215    pub fn loss_type(mut self, loss_type: LossType) -> Self {
216        self.loss_type = loss_type;
217        self
218    }
219
220    /// Enables signal-dependent Poisson-gradient truncation with `threshold`.
221    ///
222    /// The cited TPWFP experiments selected `25`. Validation requires a finite
223    /// positive value and Poisson negative log likelihood.
224    pub fn poisson_truncation_threshold(mut self, threshold: f64) -> Self {
225        self.poisson_truncation_threshold = Some(threshold);
226        self
227    }
228
229    /// Enables or disables per-source Fourier-grid offset recovery.
230    pub fn recover_illumination(mut self, recover: bool) -> Self {
231        self.recover_illumination = recover;
232        self
233    }
234
235    /// Sets the finite positive illumination-offset step size.
236    pub fn illumination_step(mut self, step: f64) -> Self {
237        self.illumination_step = step;
238        self
239    }
240
241    /// Sets positive central finite-difference spacing in Fourier-grid pixels.
242    pub fn illumination_finite_difference(mut self, distance: f64) -> Self {
243        self.illumination_finite_difference = distance;
244        self
245    }
246
247    /// Sets the non-negative maximum absolute row or column correction in grid pixels.
248    pub fn illumination_bounds(mut self, maximum_absolute_correction: f64) -> Self {
249        self.maximum_illumination_correction = maximum_absolute_correction;
250        self
251    }
252
253    /// Enables or disables simultaneous complex-pupil recovery.
254    pub fn recover_pupil(mut self, recover: bool) -> Self {
255        self.recover_pupil = recover;
256        self
257    }
258
259    /// Sets the finite positive normalized pupil-update step size.
260    pub fn pupil_step(mut self, step: f64) -> Self {
261        self.pupil_step = step;
262        self
263    }
264
265    /// Selects whether pupil values outside the compiled binary support are forced to zero.
266    pub fn constrain_pupil_support(mut self, constrain: bool) -> Self {
267        self.constrain_pupil_support = constrain;
268        self
269    }
270
271    /// Sets a non-negative complex-object isotropic total-variation weight.
272    pub fn object_tv(mut self, weight: f64) -> Self {
273        self.object_tv_weight = weight;
274        self
275    }
276
277    /// Sets the finite positive smoothing constant in the differentiable TV norm.
278    pub fn object_tv_epsilon(mut self, epsilon: f64) -> Self {
279        self.object_tv_epsilon = epsilon;
280        self
281    }
282
283    /// Sets a non-negative quadratic nearest-neighbor pupil-smoothing weight.
284    pub fn pupil_smoothing(mut self, weight: f64) -> Self {
285        self.pupil_smoothing_weight = weight;
286        self
287    }
288
289    /// Sets the maximum number of frame-gradient workers. Object, pupil, and
290    /// illumination-calibration contributions are reduced deterministically.
291    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            // The FFT adjoint contributes a 1/image_len normalization. Match
530            // the reported per-frame mean when masked pixels reduce its divisor.
531            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    /// Multiplying `sqrt(predicted)` yields the accepted residual bound.
697    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        // The backend normalizes the forward transform by 1/N, so Parseval's
777        // identity makes the spectrum L2 norm equal the object-domain RMS.
778        // Bian et al.'s MATLAB reference divides its unnormalized fft2 norm by N.
779        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}