Skip to main content

fpm_rs/model/
forward.rs

1use ndarray::{Array2, ArrayView2, ArrayViewMut2};
2use num_complex::Complex64;
3use std::{sync::Arc, thread};
4
5use crate::{
6    Result,
7    algorithms::objective::LossType,
8    array_layout::{StandardView2, checked_len_2d},
9    backend::{Backend, CpuBackend, FftDirection},
10    error::Error,
11};
12
13use super::{ImagePlaneModel, Pupil};
14
15/// Reusable host buffers for allocation-free repeated forward evaluations.
16#[derive(Clone, Debug)]
17pub struct ForwardWorkspace {
18    patch: Vec<Complex64>,
19    centered_exit: Vec<Complex64>,
20    field: Vec<Complex64>,
21    column: Vec<Complex64>,
22}
23
24impl ForwardWorkspace {
25    fn new(model: &ImagePlaneModel) -> Result<Self> {
26        let image_len = checked_len_2d(model.image_shape)?;
27        Ok(Self {
28            patch: vec![Complex64::default(); image_len],
29            centered_exit: vec![Complex64::default(); image_len],
30            field: vec![Complex64::default(); image_len],
31            column: vec![
32                Complex64::default();
33                model.image_shape.0.max(model.reconstruction_shape.0)
34            ],
35        })
36    }
37
38    /// Borrows the most recently computed coherent detector field in row-major order.
39    pub fn field(&self) -> &[Complex64] {
40        &self.field
41    }
42}
43
44/// Allocation-friendly CPU implementation of the image-plane FPM forward model.
45///
46/// Public spectrum and two-dimensional destination views must have standard
47/// C-style row-major layout. The layout is validated once at each public
48/// computational boundary and nonstandard views are rejected without copying.
49/// Methods returning an [`Array2`] allocate their result; `_into` methods borrow
50/// caller-provided workspace and destinations.
51pub struct ForwardModel<'a> {
52    model: &'a ImagePlaneModel,
53    backend: Arc<dyn Backend>,
54}
55
56impl<'a> ForwardModel<'a> {
57    /// Validates `model` and creates a CPU-backed evaluator borrowing it.
58    pub fn new(model: &'a ImagePlaneModel) -> Result<Self> {
59        let backend: Arc<dyn Backend> = Arc::new(CpuBackend::new(
60            model.image_shape,
61            model.reconstruction_shape,
62        )?);
63        Self::with_backend(model, backend)
64    }
65
66    /// Validates `model` and creates an evaluator using the shared execution `backend`.
67    pub fn with_backend(model: &'a ImagePlaneModel, backend: Arc<dyn Backend>) -> Result<Self> {
68        model.validate()?;
69        Ok(Self { model, backend })
70    }
71
72    /// Returns the compiled model borrowed by this evaluator.
73    pub fn model(&self) -> &ImagePlaneModel {
74        self.model
75    }
76
77    /// Allocates reusable buffers sized for the model's low-resolution grid.
78    pub fn workspace(&self) -> Result<ForwardWorkspace> {
79        ForwardWorkspace::new(self.model)
80    }
81
82    /// Allocates and returns one source's low-resolution complex Fourier patch.
83    pub fn extract_patch(
84        &self,
85        object_spectrum: ArrayView2<'_, Complex64>,
86        source: usize,
87    ) -> Result<Array2<Complex64>> {
88        let image_len = checked_len_2d(self.model.image_shape)?;
89        let mut values = vec![Complex64::default(); image_len];
90        self.model
91            .extract_patch(object_spectrum, source, &mut values)?;
92        Ok(Array2::from_shape_vec(self.model.image_shape, values)?)
93    }
94
95    /// Adds `scale * update` through the adjoint crop operator into `object_spectrum`.
96    pub fn insert_patch_update(
97        &self,
98        object_spectrum: ArrayViewMut2<'_, Complex64>,
99        source: usize,
100        update: &[Complex64],
101        scale: f64,
102    ) -> Result<()> {
103        self.model
104            .insert_patch_adjoint(object_spectrum, source, update, scale)
105    }
106
107    /// Multiplies a row-major complex patch by the same-shaped sampled `pupil`.
108    pub fn apply_pupil(
109        &self,
110        patch: &[Complex64],
111        pupil: &Pupil,
112        destination: &mut [Complex64],
113    ) -> Result<()> {
114        let len = checked_len_2d(self.model.image_shape)?;
115        if patch.len() != len || destination.len() != len || pupil.values.len() != len {
116            return Err(Error::InvalidShape(
117                "patch, pupil, and destination lengths must match image shape".into(),
118            ));
119        }
120        for ((destination, &patch), &pupil) in destination
121            .iter_mut()
122            .zip(patch)
123            .zip(pupil.values.as_slice())
124        {
125            *destination = patch * pupil;
126        }
127        Ok(())
128    }
129
130    /// Coherent low-resolution field for one illumination source.
131    pub fn forward_source_field(
132        &self,
133        object_spectrum: ArrayView2<'_, Complex64>,
134        pupil: &Pupil,
135        source: usize,
136    ) -> Result<Array2<Complex64>> {
137        let object_spectrum = StandardView2::try_from(object_spectrum)?;
138        let mut workspace = self.workspace()?;
139        self.forward_source_field_standard_into(object_spectrum, pupil, source, &mut workspace)?;
140        Ok(Array2::from_shape_vec(
141            self.model.image_shape,
142            workspace.field,
143        )?)
144    }
145
146    pub(crate) fn forward_source_field_standard_into<'b>(
147        &self,
148        object_spectrum: StandardView2<'_, Complex64>,
149        pupil: &Pupil,
150        source: usize,
151        workspace: &'b mut ForwardWorkspace,
152    ) -> Result<&'b [Complex64]> {
153        self.validate_workspace(workspace)?;
154        self.model
155            .extract_patch_standard(object_spectrum, source, &mut workspace.patch)?;
156        self.apply_pupil(&workspace.patch, pupil, &mut workspace.centered_exit)?;
157        ifftshift_copy(
158            &workspace.centered_exit,
159            &mut workspace.field,
160            self.model.image_shape,
161        );
162        self.backend.fft2(
163            &mut workspace.field,
164            self.model.image_shape,
165            FftDirection::Inverse,
166            &mut workspace.column,
167        )?;
168        Ok(&workspace.field)
169    }
170
171    /// Coherent field for a non-multiplexed frame. A coded frame with exactly
172    /// one source is represented by a field scaled by the square-root weight.
173    pub fn forward_field(
174        &self,
175        object_spectrum: ArrayView2<'_, Complex64>,
176        pupil: &Pupil,
177        frame: usize,
178    ) -> Result<Array2<Complex64>> {
179        if frame >= self.model.frame_count() {
180            return Err(Error::FrameOutOfRange {
181                index: frame,
182                frames: self.model.frame_count(),
183            });
184        }
185        if let Some(matrix) = &self.model.multiplexing_matrix {
186            let row = &matrix[frame];
187            if row.len() != 1 {
188                return Err(Error::Unsupported(
189                    "a multiplexed frame has no single coherent field".into(),
190                ));
191            }
192            let (source, weight) = row[0];
193            let mut field = self.forward_source_field(object_spectrum, pupil, source)?;
194            let field_scale = weight.sqrt();
195            for value in &mut field {
196                *value *= field_scale;
197            }
198            Ok(field)
199        } else {
200            self.forward_source_field(object_spectrum, pupil, frame)
201        }
202    }
203
204    /// Allocates the predicted `(row, column)` intensity array for one acquisition frame.
205    ///
206    /// Coded frames are incoherent weighted sums of source intensities. Optional frame
207    /// gain and optical background are applied to the result.
208    pub fn forward_intensity(
209        &self,
210        object_spectrum: ArrayView2<'_, Complex64>,
211        pupil: &Pupil,
212        frame: usize,
213    ) -> Result<Array2<f64>> {
214        let object_spectrum = StandardView2::try_from(object_spectrum)?;
215        let mut workspace = self.workspace()?;
216        let mut values = vec![0.0; checked_len_2d(self.model.image_shape)?];
217        self.forward_intensity_standard_into(
218            object_spectrum,
219            pupil,
220            frame,
221            &mut workspace,
222            &mut values,
223        )?;
224        Ok(Array2::from_shape_vec(self.model.image_shape, values)?)
225    }
226
227    /// Writes one predicted frame to row-major `destination`, reusing `workspace`.
228    pub fn forward_intensity_into(
229        &self,
230        object_spectrum: ArrayView2<'_, Complex64>,
231        pupil: &Pupil,
232        frame: usize,
233        workspace: &mut ForwardWorkspace,
234        destination: &mut [f64],
235    ) -> Result<()> {
236        self.forward_intensity_standard_into(
237            StandardView2::try_from(object_spectrum)?,
238            pupil,
239            frame,
240            workspace,
241            destination,
242        )
243    }
244
245    pub(crate) fn forward_intensity_standard_into(
246        &self,
247        object_spectrum: StandardView2<'_, Complex64>,
248        pupil: &Pupil,
249        frame: usize,
250        workspace: &mut ForwardWorkspace,
251        destination: &mut [f64],
252    ) -> Result<()> {
253        if frame >= self.model.frame_count() {
254            return Err(Error::FrameOutOfRange {
255                index: frame,
256                frames: self.model.frame_count(),
257            });
258        }
259        let gain = self.frame_gain(frame)?;
260        let image_len = checked_len_2d(self.model.image_shape)?;
261        if destination.len() != image_len {
262            return Err(Error::LengthMismatch {
263                actual: destination.len(),
264                expected: image_len,
265                shape: self.model.image_shape,
266            });
267        }
268        destination.fill(0.0);
269        if let Some(matrix) = &self.model.multiplexing_matrix {
270            for &(source, weight) in &matrix[frame] {
271                let field = self.forward_source_field_standard_into(
272                    object_spectrum,
273                    pupil,
274                    source,
275                    workspace,
276                )?;
277                for (intensity, value) in destination.iter_mut().zip(field) {
278                    *intensity += weight * value.norm_sqr();
279                }
280            }
281        } else {
282            let field =
283                self.forward_source_field_standard_into(object_spectrum, pupil, frame, workspace)?;
284            for (intensity, value) in destination.iter_mut().zip(field) {
285                *intensity = value.norm_sqr();
286            }
287        }
288        for (pixel, value) in destination.iter_mut().enumerate() {
289            *value = gain * *value + self.background(frame, pixel, image_len);
290        }
291        Ok(())
292    }
293
294    /// Predicts every measured frame into a contiguous
295    /// `[frame][row][column]` destination using independent worker scratch.
296    ///
297    /// Frame order and floating-point results within each frame are unchanged
298    /// by `worker_count`. Values larger than the frame count are capped; zero is
299    /// rejected. A worker count of one executes directly without spawning.
300    pub fn forward_intensity_stack_into(
301        &self,
302        object_spectrum: ArrayView2<'_, Complex64>,
303        pupil: &Pupil,
304        destination: &mut [f64],
305        worker_count: usize,
306    ) -> Result<()> {
307        let object_spectrum = StandardView2::try_from(object_spectrum)?;
308        if worker_count == 0 {
309            return Err(Error::InvalidParameter {
310                name: "worker_count",
311                reason: "must be greater than zero".into(),
312            });
313        }
314        let image_len = checked_len_2d(self.model.image_shape)?;
315        let expected = image_len
316            .checked_mul(self.model.frame_count())
317            .ok_or_else(|| Error::InvalidShape("forward stack length overflows".into()))?;
318        if destination.len() != expected {
319            return Err(Error::LengthMismatch {
320                actual: destination.len(),
321                expected,
322                shape: (self.model.frame_count(), image_len),
323            });
324        }
325        let workers = worker_count.min(self.model.frame_count());
326        if workers == 1 {
327            let mut workspace = self.workspace()?;
328            for (frame, frame_destination) in destination.chunks_exact_mut(image_len).enumerate() {
329                self.forward_intensity_standard_into(
330                    object_spectrum,
331                    pupil,
332                    frame,
333                    &mut workspace,
334                    frame_destination,
335                )?;
336            }
337            return Ok(());
338        }
339
340        let frames_per_worker = self.model.frame_count().div_ceil(workers);
341        let values_per_worker = frames_per_worker * image_len;
342        thread::scope(|scope| {
343            let handles: Vec<_> = destination
344                .chunks_mut(values_per_worker)
345                .enumerate()
346                .map(|(chunk_index, output)| {
347                    let first_frame = chunk_index * frames_per_worker;
348                    scope.spawn(move || -> Result<()> {
349                        let mut workspace = self.workspace()?;
350                        for (local_frame, frame_destination) in
351                            output.chunks_exact_mut(image_len).enumerate()
352                        {
353                            self.forward_intensity_standard_into(
354                                object_spectrum,
355                                pupil,
356                                first_frame + local_frame,
357                                &mut workspace,
358                                frame_destination,
359                            )?;
360                        }
361                        Ok(())
362                    })
363                })
364                .collect();
365            for handle in handles {
366                handle
367                    .join()
368                    .map_err(|_| Error::Numerical("parallel forward worker panicked".into()))??;
369            }
370            Ok(())
371        })
372    }
373
374    /// Allocation-owning convenience wrapper for [`Self::forward_intensity_stack_into`].
375    pub fn forward_intensity_stack(
376        &self,
377        object_spectrum: ArrayView2<'_, Complex64>,
378        pupil: &Pupil,
379        worker_count: usize,
380    ) -> Result<Vec<f64>> {
381        let image_len = checked_len_2d(self.model.image_shape)?;
382        let stack_len = image_len
383            .checked_mul(self.model.frame_count())
384            .ok_or_else(|| Error::ShapeOverflow {
385                shape: vec![
386                    self.model.frame_count(),
387                    self.model.image_shape.0,
388                    self.model.image_shape.1,
389                ],
390            })?;
391        let mut destination = vec![0.0; stack_len];
392        self.forward_intensity_stack_into(object_spectrum, pupil, &mut destination, worker_count)?;
393        Ok(destination)
394    }
395
396    /// Returns row-major `predicted - measured` intensity residuals for one frame.
397    pub fn residual(
398        &self,
399        object_spectrum: ArrayView2<'_, Complex64>,
400        pupil: &Pupil,
401        frame: usize,
402        measured: &[f64],
403    ) -> Result<Vec<f64>> {
404        let predicted = self.forward_intensity(object_spectrum, pupil, frame)?;
405        if measured.len() != predicted.len() {
406            return Err(Error::InvalidShape(
407                "measurement does not match predicted frame".into(),
408            ));
409        }
410        Ok(predicted
411            .iter()
412            .zip(measured)
413            .map(|(&predicted, &measured)| predicted - measured)
414            .collect())
415    }
416
417    /// Replaces coherent-field amplitude with measured amplitude while retaining phase.
418    ///
419    /// Frame gain and background are inverted before taking the square root of measured
420    /// intensity. `epsilon` is a positive floor for dark predicted amplitudes.
421    pub fn amplitude_projection(
422        &self,
423        field: &[Complex64],
424        measured_intensity: &[f64],
425        frame: usize,
426        epsilon: f64,
427    ) -> Result<Vec<Complex64>> {
428        if field.len() != measured_intensity.len() {
429            return Err(Error::InvalidShape(
430                "field and measured image lengths differ".into(),
431            ));
432        }
433        if frame >= self.model.frame_count() {
434            return Err(Error::FrameOutOfRange {
435                index: frame,
436                frames: self.model.frame_count(),
437            });
438        }
439        if !epsilon.is_finite() || epsilon <= 0.0 {
440            return Err(Error::InvalidParameter {
441                name: "epsilon",
442                reason: "must be finite and positive".into(),
443            });
444        }
445        if self
446            .model
447            .multiplexing_matrix
448            .as_ref()
449            .is_some_and(|matrix| matrix[frame].len() != 1)
450        {
451            return Err(Error::Unsupported(
452                "amplitude projection is not defined for an incoherent multiplexed field".into(),
453            ));
454        }
455        let gain = self.frame_gain(frame)?;
456        let image_len = field.len();
457        Ok(field
458            .iter()
459            .zip(measured_intensity)
460            .enumerate()
461            .map(|(pixel, (&value, &measurement))| {
462                let corrected =
463                    ((measurement - self.background(frame, pixel, image_len)) / gain).max(0.0);
464                let target = corrected.sqrt();
465                let magnitude = value.norm();
466                if magnitude > epsilon {
467                    value * (target / magnitude)
468                } else {
469                    Complex64::new(target, 0.0)
470                }
471            })
472            .collect())
473    }
474
475    /// Evaluates `loss_type` between one predicted and measured intensity frame.
476    pub fn frame_loss(
477        &self,
478        object_spectrum: ArrayView2<'_, Complex64>,
479        pupil: &Pupil,
480        frame: usize,
481        measured: &[f64],
482        loss_type: LossType,
483    ) -> Result<f64> {
484        let predicted = self.forward_intensity(object_spectrum, pupil, frame)?;
485        let predicted = predicted.as_slice().ok_or_else(|| {
486            Error::InvalidModel("internally generated intensity was not standard layout".into())
487        })?;
488        crate::algorithms::objective::loss(predicted, measured, loss_type)
489    }
490
491    fn frame_gain(&self, frame: usize) -> Result<f64> {
492        let gain = self
493            .model
494            .frame_gains
495            .as_ref()
496            .map_or(1.0, |gains| gains[frame]);
497        if !gain.is_finite() || gain <= 0.0 {
498            return Err(Error::InvalidModel(format!(
499                "frame {frame} has invalid gain {gain}"
500            )));
501        }
502        Ok(gain)
503    }
504
505    fn background(&self, frame: usize, pixel: usize, image_len: usize) -> f64 {
506        self.model.background.as_ref().map_or(0.0, |values| {
507            values[if values.len() == image_len {
508                pixel
509            } else {
510                frame * image_len + pixel
511            }]
512        })
513    }
514
515    fn validate_workspace(&self, workspace: &ForwardWorkspace) -> Result<()> {
516        let image_len = checked_len_2d(self.model.image_shape)?;
517        if workspace.patch.len() != image_len
518            || workspace.centered_exit.len() != image_len
519            || workspace.field.len() != image_len
520            || workspace.column.len() < self.model.image_shape.0
521        {
522            return Err(Error::InvalidShape(
523                "forward workspace does not match the model image shape".into(),
524            ));
525        }
526        Ok(())
527    }
528}
529
530pub(crate) fn fftshift_copy(
531    source: &[Complex64],
532    destination: &mut [Complex64],
533    shape: (usize, usize),
534) {
535    shift_copy(source, destination, shape, shape.0 / 2, shape.1 / 2);
536}
537
538pub(crate) fn ifftshift_copy(
539    source: &[Complex64],
540    destination: &mut [Complex64],
541    shape: (usize, usize),
542) {
543    shift_copy(
544        source,
545        destination,
546        shape,
547        shape.0.div_ceil(2),
548        shape.1.div_ceil(2),
549    );
550}
551
552fn shift_copy(
553    source: &[Complex64],
554    destination: &mut [Complex64],
555    shape: (usize, usize),
556    row_shift: usize,
557    column_shift: usize,
558) {
559    debug_assert_eq!(source.len(), shape.0 * shape.1);
560    debug_assert_eq!(destination.len(), source.len());
561    for row in 0..shape.0 {
562        for column in 0..shape.1 {
563            let destination_row = (row + row_shift) % shape.0;
564            let destination_column = (column + column_shift) % shape.1;
565            destination[destination_row * shape.1 + destination_column] =
566                source[row * shape.1 + column];
567        }
568    }
569}