1use ndarray::Array2;
2use num_complex::Complex64;
3use rand::{SeedableRng, rngs::StdRng};
4use rand_distr::{Distribution, Normal};
5
6use crate::{
7 Result,
8 array_layout::checked_len_2d,
9 backend::{Backend, CpuBackend, FftDirection},
10 error::Error,
11 measurements::{FrameMetadata, MeasurementStack},
12 model::{ForwardModel, ImagePlaneModel, fftshift_copy},
13};
14
15use super::{
16 CameraModel, IlluminationAcquisitionErrors, SimulationResult, SyntheticObject,
17 result::SimulationParameters,
18};
19
20pub struct Simulator {
26 true_model: ImagePlaneModel,
27 reconstruction_model: Option<ImagePlaneModel>,
28 object: Option<SyntheticObject>,
29 camera: Option<CameraModel>,
30 illumination_acquisition_errors: Option<IlluminationAcquisitionErrors>,
31 seed: u64,
32 ideal: bool,
33}
34
35impl Simulator {
36 pub fn new(true_model: ImagePlaneModel) -> Self {
41 Self {
42 true_model,
43 reconstruction_model: None,
44 object: None,
45 camera: None,
46 illumination_acquisition_errors: None,
47 seed: 0,
48 ideal: false,
49 }
50 }
51
52 pub fn ideal(model: ImagePlaneModel) -> Self {
54 Self {
55 ideal: true,
56 ..Self::new(model)
57 }
58 }
59
60 pub fn object(mut self, object: SyntheticObject) -> Self {
62 self.object = Some(object);
63 self
64 }
65
66 pub fn camera(mut self, camera: CameraModel) -> Self {
68 self.camera = Some(camera);
69 self.ideal = false;
70 self
71 }
72
73 pub fn illumination_acquisition_errors(
75 mut self,
76 errors: IlluminationAcquisitionErrors,
77 ) -> Self {
78 self.illumination_acquisition_errors = Some(errors);
79 self.ideal = false;
80 self
81 }
82
83 pub fn reconstruction_model(mut self, model: ImagePlaneModel) -> Self {
90 self.reconstruction_model = Some(model);
91 self
92 }
93
94 pub fn seed(mut self, seed: u64) -> Self {
96 self.seed = seed;
97 self
98 }
99
100 pub fn simulate(self) -> Result<SimulationResult> {
108 self.true_model.validate()?;
109 let object = self.object.ok_or(Error::InvalidParameter {
110 name: "object",
111 reason: "a ground-truth object is required".into(),
112 })?;
113 if object.shape() != self.true_model.reconstruction_shape {
114 return Err(Error::InvalidShape(format!(
115 "object shape {:?} differs from reconstruction shape {:?}",
116 object.shape(),
117 self.true_model.reconstruction_shape
118 )));
119 }
120 let mut reconstruction_model = self
121 .reconstruction_model
122 .unwrap_or_else(|| self.true_model.clone());
123 reconstruction_model.validate()?;
124 if reconstruction_model.image_shape != self.true_model.image_shape
125 || reconstruction_model.reconstruction_shape != self.true_model.reconstruction_shape
126 || reconstruction_model.frame_count() != self.true_model.frame_count()
127 {
128 return Err(Error::InvalidModel(
129 "true and reconstruction models must have matching image, object, and frame dimensions"
130 .into(),
131 ));
132 }
133 if let Some(camera) = &self.camera {
134 reconstruction_model = camera.compile_reconstruction_model(reconstruction_model)?;
135 }
136 let mut true_model = self.true_model;
137 let mut rng = StdRng::seed_from_u64(self.seed);
138 let missing_frames = self
139 .illumination_acquisition_errors
140 .as_ref()
141 .map_or_else(Vec::new, |errors| errors.missing_frames.clone());
142 if let Some(errors) = &self.illumination_acquisition_errors {
143 apply_illumination_acquisition_errors(&mut true_model, errors, &mut rng)?;
144 }
145 true_model.validate()?;
146
147 let object_spectrum = object_spectrum(&object, &true_model)?;
148 let forward = ForwardModel::new(&true_model)?;
149 let image_len = checked_len_2d(true_model.image_shape)?;
150 let worker_count = std::thread::available_parallelism().map_or(1, |count| count.get());
151 let stack_len = image_len
152 .checked_mul(true_model.frame_count())
153 .ok_or_else(|| Error::ShapeOverflow {
154 shape: vec![
155 true_model.frame_count(),
156 true_model.image_shape.0,
157 true_model.image_shape.1,
158 ],
159 })?;
160 let mut data = vec![0.0; stack_len];
161 forward.forward_intensity_stack_into(
162 object_spectrum.view(),
163 &true_model.pupil,
164 &mut data,
165 worker_count,
166 )?;
167 for (frame, frame_data) in data.chunks_exact_mut(image_len).enumerate() {
170 if missing_frames.contains(&frame) {
171 for (pixel, value) in frame_data.iter_mut().enumerate() {
174 *value = true_model.background_value(frame, pixel)?;
175 }
176 }
177 if let Some(camera) = &self.camera {
178 camera.measure_frame(frame_data, &mut rng)?;
179 }
180 }
181 let metadata = (0..true_model.frame_count())
182 .map(|frame| {
183 let mut metadata = FrameMetadata::new(frame);
184 if missing_frames.contains(&frame) {
185 metadata.weight = 0.0;
186 }
187 metadata
188 })
189 .collect();
190 let measurements = MeasurementStack::from_vec(data, true_model.image_shape, metadata)?;
191 Ok(SimulationResult {
192 measurements,
193 ground_truth_object: object.field.into_inner(),
194 true_model: true_model.clone(),
195 reconstruction_model,
196 camera: self.camera,
197 illumination_acquisition_errors: self.illumination_acquisition_errors,
198 parameters: SimulationParameters {
199 ideal: self.ideal,
200 frame_count: true_model.frame_count(),
201 image_shape: true_model.image_shape,
202 missing_frames,
203 },
204 random_seed: self.seed,
205 })
206 }
207}
208
209fn object_spectrum(object: &SyntheticObject, model: &ImagePlaneModel) -> Result<Array2<Complex64>> {
210 let backend = CpuBackend::new(model.image_shape, model.reconstruction_shape)?;
211 let mut values = object.field.as_slice().to_vec();
212 let mut column =
213 vec![Complex64::default(); model.image_shape.0.max(model.reconstruction_shape.0)];
214 backend.fft2(
215 &mut values,
216 model.reconstruction_shape,
217 FftDirection::Forward,
218 &mut column,
219 )?;
220 let mut centered = vec![Complex64::default(); values.len()];
221 fftshift_copy(&values, &mut centered, model.reconstruction_shape);
222 Ok(Array2::from_shape_vec(
223 model.reconstruction_shape,
224 centered,
225 )?)
226}
227
228fn apply_illumination_acquisition_errors(
229 model: &mut ImagePlaneModel,
230 errors: &IlluminationAcquisitionErrors,
231 rng: &mut StdRng,
232) -> Result<()> {
233 if !errors.frame_gain_relative_std.is_finite() || errors.frame_gain_relative_std < 0.0 {
234 return Err(Error::InvalidParameter {
235 name: "frame_gain_relative_std",
236 reason: "must be finite and non-negative".into(),
237 });
238 }
239 let frame_count = model.frame_count();
240 let source_count = model.source_count();
241 let multiplexed = model.is_multiplexed();
242 let mut sorted_missing = errors.missing_frames.clone();
243 sorted_missing.sort_unstable();
244 if sorted_missing.iter().any(|&frame| frame >= frame_count)
245 || sorted_missing.windows(2).any(|pair| pair[0] == pair[1])
246 {
247 return Err(Error::InvalidParameter {
248 name: "missing_frames",
249 reason: format!("must contain unique frame indices below {frame_count}"),
250 });
251 }
252 if let Some(permutation) = &errors.source_permutation {
253 let mut sorted = permutation.clone();
254 sorted.sort_unstable();
255 if permutation.len() != source_count
256 || sorted
257 .iter()
258 .enumerate()
259 .any(|(expected, &actual)| expected != actual)
260 {
261 return Err(Error::InvalidParameter {
262 name: "source_permutation",
263 reason: format!("must be a permutation of 0..{source_count}"),
264 });
265 }
266 }
267 if let Some(permutation) = &errors.source_permutation {
268 model.k_vectors = permutation
269 .iter()
270 .map(|&index| model.k_vectors[index])
271 .collect();
272 model.crop_indices.crops = permutation
273 .iter()
274 .map(|&index| model.crop_indices.crops[index])
275 .collect();
276 model.subpixel_offsets = model
277 .subpixel_offsets
278 .as_ref()
279 .map(|offsets| permutation.iter().map(|&index| offsets[index]).collect());
280 if !multiplexed && let Some(gains) = &model.frame_gains {
281 model.frame_gains = Some(permutation.iter().map(|&index| gains[index]).collect());
282 }
283 let image_len = checked_len_2d(model.image_shape)?;
284 let stack_len = image_len
285 .checked_mul(frame_count)
286 .ok_or_else(|| Error::ShapeOverflow {
287 shape: vec![frame_count, model.image_shape.0, model.image_shape.1],
288 })?;
289 if !multiplexed
290 && let Some(background) = &model.background
291 && background.len() == stack_len
292 {
293 let mut reordered = Vec::with_capacity(background.len());
294 for &index in permutation {
295 let start = index
296 .checked_mul(image_len)
297 .ok_or_else(|| Error::ShapeOverflow {
298 shape: vec![index, model.image_shape.0, model.image_shape.1],
299 })?;
300 let end = start
301 .checked_add(image_len)
302 .ok_or_else(|| Error::ShapeOverflow {
303 shape: vec![
304 index.saturating_add(1),
305 model.image_shape.0,
306 model.image_shape.1,
307 ],
308 })?;
309 reordered.extend_from_slice(background.get(start..end).ok_or_else(|| {
310 Error::InvalidModel("background source permutation is out of range".into())
311 })?);
312 }
313 model.background = Some(reordered);
314 }
315 }
316 if errors.frame_gain_relative_std > 0.0 {
317 let distribution = Normal::new(1.0, errors.frame_gain_relative_std)
318 .map_err(|error| Error::Numerical(error.to_string()))?;
319 let gains = model
320 .frame_gains
321 .get_or_insert_with(|| vec![1.0; frame_count]);
322 for gain in gains {
323 *gain *= distribution.sample(rng).max(0.01);
324 }
325 }
326 Ok(())
327}