Skip to main content

fpm_rs/reconstruction/
schedule.rs

1use rand::{SeedableRng, rngs::StdRng, seq::SliceRandom};
2use serde::{Deserialize, Serialize};
3
4use crate::{Result, experiment::KVector, measurements::MeasurementRead, model::ImagePlaneModel};
5
6use super::ReconstructionProblem;
7
8/// Rule for ordering acquisition frames within each reconstruction iteration.
9#[derive(Clone, Debug, Default, Serialize, Deserialize)]
10pub enum FrameSchedule {
11    /// Ascending acquisition-frame index.
12    #[default]
13    Sequential,
14    /// Increasing transverse illumination magnitude, approximating brightfield first.
15    BrightfieldFirst,
16    /// Increasing polar radius then azimuth in transverse-wave-vector space.
17    SpiralOut,
18    /// Deterministic per-iteration pseudorandom permutation.
19    RandomShuffle {
20        /// Base random seed combined with the zero-based iteration number.
21        seed: u64,
22    },
23    /// Highest empirical shot-noise SNR first. This requires measurements and
24    /// is applied by [`Self::order_for_problem`]; [`Self::order`] retains
25    /// sequential order because a compiled model alone contains no signal data.
26    SnrWeighted,
27}
28
29impl FrameSchedule {
30    /// Returns a model-only frame order.
31    ///
32    /// [`Self::SnrWeighted`] cannot be estimated from model geometry, so this
33    /// method returns sequential order for that variant. Reconstruction runners
34    /// use [`Self::order_for_problem`] and therefore have access to measurements.
35    pub fn order(&self, model: &ImagePlaneModel, iteration: usize) -> Vec<usize> {
36        let mut order: Vec<usize> = (0..model.frame_count()).collect();
37        match self {
38            Self::Sequential | Self::SnrWeighted => {}
39            Self::BrightfieldFirst => order.sort_by(|&left, &right| {
40                let left_vector = frame_vector(model, left);
41                let right_vector = frame_vector(model, right);
42                left_vector
43                    .kx
44                    .hypot(left_vector.ky)
45                    .total_cmp(&right_vector.kx.hypot(right_vector.ky))
46            }),
47            Self::SpiralOut => order.sort_by(|&left, &right| {
48                let left_vector = frame_vector(model, left);
49                let right_vector = frame_vector(model, right);
50                left_vector
51                    .kx
52                    .hypot(left_vector.ky)
53                    .total_cmp(&right_vector.kx.hypot(right_vector.ky))
54                    .then_with(|| {
55                        left_vector
56                            .ky
57                            .atan2(left_vector.kx)
58                            .rem_euclid(std::f64::consts::TAU)
59                            .total_cmp(
60                                &right_vector
61                                    .ky
62                                    .atan2(right_vector.kx)
63                                    .rem_euclid(std::f64::consts::TAU),
64                            )
65                    })
66            }),
67            Self::RandomShuffle { seed } => {
68                let mut rng = StdRng::seed_from_u64(
69                    seed.wrapping_add((iteration as u64).wrapping_mul(0x9e3779b97f4a7c15)),
70                );
71                order.shuffle(&mut rng);
72            }
73        }
74        order
75    }
76
77    /// Returns a frame order using both the model and measured intensities.
78    ///
79    /// The SNR-weighted schedule ranks frames by
80    /// `frame_weight * sqrt(mean(max(measured - known_background, 0)))` over
81    /// unmasked pixels. This is an empirical Poisson shot-noise proxy, not a
82    /// camera-noise calibration. Zero-weight or fully masked frames are placed
83    /// last, and equal scores retain ascending frame-index order.
84    pub fn order_for_problem<M: MeasurementRead>(
85        &self,
86        problem: &ReconstructionProblem<M>,
87        iteration: usize,
88    ) -> Result<Vec<usize>> {
89        if !matches!(self, Self::SnrWeighted) {
90            return Ok(self.order(&problem.model, iteration));
91        }
92        let mut scored = Vec::with_capacity(problem.model.frame_count());
93        for frame in 0..problem.model.frame_count() {
94            scored.push((frame, frame_snr_score(problem, frame)?));
95        }
96        scored.sort_by(|&(left_frame, left_score), &(right_frame, right_score)| {
97            right_score
98                .total_cmp(&left_score)
99                .then_with(|| left_frame.cmp(&right_frame))
100        });
101        Ok(scored.into_iter().map(|(frame, _)| frame).collect())
102    }
103}
104
105fn frame_snr_score<M: MeasurementRead>(
106    problem: &ReconstructionProblem<M>,
107    frame: usize,
108) -> Result<f64> {
109    let weight = problem.measurements.frame_weight(frame)?;
110    if weight == 0.0 {
111        return Ok(f64::NEG_INFINITY);
112    }
113    let measured = problem.measurements.frame(frame)?;
114    let mask = problem.measurements.frame_mask(frame)?;
115    let mut signal_sum = 0.0;
116    let mut valid_pixels = 0;
117    for (pixel, &value) in measured.iter().enumerate() {
118        if mask.is_none_or(|values| values[pixel] != 0) {
119            let background = problem.model.background_value(frame, pixel)?;
120            signal_sum += (value - background).max(0.0);
121            valid_pixels += 1;
122        }
123    }
124    if valid_pixels == 0 {
125        Ok(f64::NEG_INFINITY)
126    } else {
127        Ok(weight * (signal_sum / valid_pixels as f64).sqrt())
128    }
129}
130
131fn frame_vector(model: &ImagePlaneModel, frame: usize) -> KVector {
132    if let Some(matrix) = &model.multiplexing_matrix {
133        let mut vector = KVector::default();
134        let mut total_weight = 0.0;
135        for &(source, weight) in &matrix[frame] {
136            vector.kx += weight * model.k_vectors[source].kx;
137            vector.ky += weight * model.k_vectors[source].ky;
138            total_weight += weight;
139        }
140        vector.kx /= total_weight;
141        vector.ky /= total_weight;
142        vector
143    } else {
144        model.k_vectors[frame]
145    }
146}