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#[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 pub fn field(&self) -> &[Complex64] {
40 &self.field
41 }
42}
43
44pub struct ForwardModel<'a> {
52 model: &'a ImagePlaneModel,
53 backend: Arc<dyn Backend>,
54}
55
56impl<'a> ForwardModel<'a> {
57 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 pub fn with_backend(model: &'a ImagePlaneModel, backend: Arc<dyn Backend>) -> Result<Self> {
68 model.validate()?;
69 Ok(Self { model, backend })
70 }
71
72 pub fn model(&self) -> &ImagePlaneModel {
74 self.model
75 }
76
77 pub fn workspace(&self) -> Result<ForwardWorkspace> {
79 ForwardWorkspace::new(self.model)
80 }
81
82 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 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 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 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 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 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 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 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 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 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 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 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}