1use ndarray::{Array2, Array3};
2
3use crate::{
4 Complex64, Result,
5 array_layout::checked_len_2d,
6 configuration::{ExperimentDescription, SimulationConfiguration},
7 error::Error,
8 experiment::Illumination,
9 measurements::MeasurementStack,
10 model::{ImagePlaneModel, ReconstructionShape},
11 reconstruction::ReconstructionProblem,
12};
13
14use super::Dataset;
15
16#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub struct Rect {
19 pub row: usize,
21 pub column: usize,
23 pub height: usize,
25 pub width: usize,
27}
28
29impl Rect {
30 pub fn new(row: usize, column: usize, height: usize, width: usize) -> Result<Self> {
32 if height == 0 || width == 0 {
33 return Err(Error::Dataset(
34 "dataset crop height and width must be non-zero".into(),
35 ));
36 }
37 Ok(Self {
38 row,
39 column,
40 height,
41 width,
42 })
43 }
44}
45
46#[derive(Clone, Debug, PartialEq, Eq)]
48pub enum FrameSelector {
49 All,
51 EveryNth(usize),
53 Indices(Vec<usize>),
55}
56
57#[derive(Clone, Debug)]
59pub struct DatasetSubset {
60 measurements: MeasurementStack,
61 configuration: SimulationConfiguration,
62 ground_truth_object: Option<Array2<Complex64>>,
63 valid_object_mask: Option<Array2<u8>>,
64 provenance: std::collections::BTreeMap<String, String>,
65 measurement_units: Option<String>,
66}
67
68impl DatasetSubset {
69 pub fn measurements(&self) -> &MeasurementStack {
71 &self.measurements
72 }
73
74 pub fn configuration(&self) -> &SimulationConfiguration {
76 &self.configuration
77 }
78
79 pub fn ground_truth_object(&self) -> Option<&Array2<Complex64>> {
81 self.ground_truth_object.as_ref()
82 }
83
84 pub fn valid_object_mask(&self) -> Option<&Array2<u8>> {
86 self.valid_object_mask.as_ref()
87 }
88
89 pub fn provenance(&self) -> &std::collections::BTreeMap<String, String> {
91 &self.provenance
92 }
93
94 pub fn measurement_units(&self) -> Option<&str> {
96 self.measurement_units.as_deref()
97 }
98
99 pub fn reconstruction_problem(&self) -> Result<ReconstructionProblem<MeasurementStack>> {
101 ReconstructionProblem::new(
102 self.measurements.clone(),
103 self.configuration
104 .compiled_models
105 .reconstruction_model
106 .clone(),
107 )
108 }
109}
110
111#[derive(Clone, Debug)]
113pub struct DatasetSubsetBuilder<'a> {
114 dataset: &'a Dataset,
115 frames: FrameSelector,
116 crop: Option<Rect>,
117}
118
119impl<'a> DatasetSubsetBuilder<'a> {
120 pub(crate) fn new(dataset: &'a Dataset) -> Self {
121 Self {
122 dataset,
123 frames: FrameSelector::All,
124 crop: None,
125 }
126 }
127
128 pub fn frames(mut self, selector: FrameSelector) -> Self {
130 self.frames = selector;
131 self
132 }
133
134 pub fn every_nth_frame(self, step: usize) -> Self {
136 self.frames(FrameSelector::EveryNth(step))
137 }
138
139 pub fn crop(mut self, crop: Rect) -> Self {
141 self.crop = Some(crop);
142 self
143 }
144
145 pub fn crop_pixels(
147 self,
148 row: usize,
149 column: usize,
150 height: usize,
151 width: usize,
152 ) -> Result<Self> {
153 Ok(self.crop(Rect::new(row, column, height, width)?))
154 }
155
156 pub fn build(self) -> Result<DatasetSubset> {
159 let source = self.dataset.measurements();
160 let model = &self
161 .dataset
162 .configuration()
163 .compiled_models
164 .reconstruction_model;
165 let indices = selected_indices(&self.frames, source.frame_count())?;
166 let source_shape = source.image_shape();
167 let crop = self.crop.unwrap_or(Rect {
168 row: 0,
169 column: 0,
170 height: source_shape.0,
171 width: source_shape.1,
172 });
173 validate_crop(crop, source_shape)?;
174 let crop_len = checked_len_2d((crop.height, crop.width))?;
175 let data_len = crop_len
176 .checked_mul(indices.len())
177 .ok_or_else(|| Error::ShapeOverflow {
178 shape: vec![indices.len(), crop.height, crop.width],
179 })?;
180
181 let mut data = Vec::with_capacity(data_len);
182 let mut metadata = Vec::with_capacity(indices.len());
183 for (new_index, &source_index) in indices.iter().enumerate() {
184 crop_frame(source.frame(source_index)?, source_shape, crop, &mut data);
185 let mut frame_metadata = source.frame_metadata()[source_index].clone();
186 frame_metadata.original_frame_index =
187 Some(frame_metadata.original_frame_index.unwrap_or(source_index));
188 frame_metadata.original_illumination_index = frame_metadata
189 .original_illumination_index
190 .or(frame_metadata.illumination_index);
191 frame_metadata.frame_index = new_index;
192 frame_metadata.illumination_index = (!model.is_multiplexed()).then_some(new_index);
193 metadata.push(frame_metadata);
194 }
195 let mut measurements =
196 MeasurementStack::from_vec(data, (crop.height, crop.width), metadata)?;
197 if let Some(values) = crop_optional_shared(source.dark_frame_slice(), source_shape, crop)? {
198 measurements = measurements
199 .with_dark_frame(Array2::from_shape_vec((crop.height, crop.width), values)?)?;
200 }
201 if let Some(values) = crop_optional_shared(source.flat_field_slice(), source_shape, crop)? {
202 measurements = measurements
203 .with_flat_field(Array2::from_shape_vec((crop.height, crop.width), values)?)?;
204 }
205 if let Some(values) = crop_optional_frames(
206 source.background_slice(),
207 source_shape,
208 source.frame_count(),
209 &indices,
210 crop,
211 )? {
212 if values.len() == crop_len {
213 measurements = measurements
214 .with_background(Array2::from_shape_vec((crop.height, crop.width), values)?)?;
215 } else {
216 measurements = measurements.with_per_frame_background(Array3::from_shape_vec(
217 (indices.len(), crop.height, crop.width),
218 values,
219 )?)?;
220 }
221 }
222 if let Some(values) = crop_optional_masks(
223 source.masks_slice(),
224 source_shape,
225 source.frame_count(),
226 &indices,
227 crop,
228 )? {
229 if values.len() == crop_len {
230 measurements = measurements
231 .with_masks(Array2::from_shape_vec((crop.height, crop.width), values)?)?;
232 } else {
233 measurements = measurements.with_per_frame_masks(Array3::from_shape_vec(
234 (indices.len(), crop.height, crop.width),
235 values,
236 )?)?;
237 }
238 }
239 measurements = measurements.with_preprocessing(source.preprocessing().clone())?;
240
241 let scale_y = model.reconstruction_shape.0 as f64 / model.image_shape.0 as f64;
242 let scale_x = model.reconstruction_shape.1 as f64 / model.image_shape.1 as f64;
243 let reconstruction_shape = (
244 (crop.height as f64 * scale_y).round() as usize,
245 (crop.width as f64 * scale_x).round() as usize,
246 );
247 let object_crop = scaled_object_crop(crop, scale_y, scale_x)?;
248 let ground_truth_object = self
249 .dataset
250 .ground_truth_object()
251 .map(|truth| crop_array(truth, object_crop))
252 .transpose()?;
253 let valid_object_mask = self
254 .dataset
255 .valid_object_mask()
256 .map(|mask| crop_array(mask, object_crop))
257 .transpose()?;
258 let source_configuration = self.dataset.configuration();
259 let true_experiment = subset_experiment(
260 &source_configuration.true_experiment,
261 &source_configuration.compiled_models.true_model,
262 &indices,
263 crop,
264 )?;
265 let reconstruction_experiment = subset_experiment(
266 &source_configuration.reconstruction_experiment,
267 &source_configuration.compiled_models.reconstruction_model,
268 &indices,
269 crop,
270 )?;
271 let configuration = SimulationConfiguration::new(
272 true_experiment,
273 reconstruction_experiment,
274 (crop.height, crop.width),
275 ReconstructionShape::Exact(reconstruction_shape),
276 )?
277 .with_random_seed(source_configuration.random_seed);
278 Ok(DatasetSubset {
279 measurements,
280 configuration,
281 ground_truth_object,
282 valid_object_mask,
283 provenance: self.dataset.provenance().clone(),
284 measurement_units: self.dataset.measurement_units().map(str::to_owned),
285 })
286 }
287}
288
289fn subset_experiment(
290 description: &ExperimentDescription,
291 model: &ImagePlaneModel,
292 indices: &[usize],
293 crop: Rect,
294) -> Result<ExperimentDescription> {
295 let (k_vectors, frame_weights) = match &model.multiplexing_matrix {
296 Some(matrix) => (
297 model.k_vectors.clone(),
298 Some(
299 indices
300 .iter()
301 .map(|&index| matrix[index].clone())
302 .collect::<Vec<_>>(),
303 ),
304 ),
305 None => (
306 indices
307 .iter()
308 .map(|&index| model.k_vectors[index])
309 .collect(),
310 None,
311 ),
312 };
313 let frame_gains: Vec<f64> = model.frame_gains.as_ref().map_or_else(
314 || vec![1.0; indices.len()],
315 |gains| indices.iter().map(|&index| gains[index]).collect(),
316 );
317 let acquisition = crate::experiment::AcquisitionPlan::from_sparse(match frame_weights {
318 Some(rows) => rows
319 .into_iter()
320 .zip(frame_gains)
321 .map(|(row, gain)| {
322 crate::experiment::IlluminationFrame::new(
323 row.into_iter()
324 .map(|(source, intensity_weight)| {
325 crate::experiment::SourceContribution::new(source, intensity_weight)
326 })
327 .collect(),
328 gain,
329 )
330 })
331 .collect(),
332 None => (0..k_vectors.len())
333 .zip(frame_gains)
334 .map(|(source, gain)| {
335 crate::experiment::IlluminationFrame::new(
336 vec![crate::experiment::SourceContribution::new(source, 1.0)],
337 gain,
338 )
339 })
340 .collect(),
341 })?;
342 let illumination = Illumination::new(
343 crate::experiment::KVectorList::new(k_vectors).into(),
344 crate::experiment::SourceCalibration::unity(),
345 acquisition,
346 );
347 let mut subset = ExperimentDescription::new(description.optics.clone(), illumination);
348 subset.optical_background = crop_optional_frames(
349 model.background.as_deref(),
350 model.image_shape,
351 model.frame_count(),
352 indices,
353 crop,
354 )?;
355 subset.validate()?;
356 Ok(subset)
357}
358
359fn selected_indices(selector: &FrameSelector, frame_count: usize) -> Result<Vec<usize>> {
360 let indices = match selector {
361 FrameSelector::All => (0..frame_count).collect(),
362 FrameSelector::EveryNth(0) => {
363 return Err(Error::Dataset(
364 "frame subset step must be greater than zero".into(),
365 ));
366 }
367 FrameSelector::EveryNth(step) => (0..frame_count).step_by(*step).collect(),
368 FrameSelector::Indices(indices) => indices.clone(),
369 };
370 if indices.is_empty() {
371 return Err(Error::Dataset(
372 "dataset frame subset must contain at least one frame".into(),
373 ));
374 }
375 let mut seen = vec![false; frame_count];
376 for &index in &indices {
377 if index >= frame_count {
378 return Err(Error::FrameOutOfRange {
379 index,
380 frames: frame_count,
381 });
382 }
383 if seen[index] {
384 return Err(Error::Dataset(format!(
385 "dataset frame subset contains duplicate index {index}"
386 )));
387 }
388 seen[index] = true;
389 }
390 Ok(indices)
391}
392
393fn validate_crop(crop: Rect, shape: (usize, usize)) -> Result<()> {
394 if crop
395 .row
396 .checked_add(crop.height)
397 .is_none_or(|end| end > shape.0)
398 || crop
399 .column
400 .checked_add(crop.width)
401 .is_none_or(|end| end > shape.1)
402 {
403 return Err(Error::Dataset(format!(
404 "crop {crop:?} is outside measurement shape {shape:?}"
405 )));
406 }
407 Ok(())
408}
409
410fn crop_frame<T: Copy>(source: &[T], shape: (usize, usize), crop: Rect, output: &mut Vec<T>) {
411 for row in crop.row..crop.row + crop.height {
412 let start = row * shape.1 + crop.column;
413 output.extend_from_slice(&source[start..start + crop.width]);
414 }
415}
416
417fn scaled_object_crop(crop: Rect, scale_y: f64, scale_x: f64) -> Result<Rect> {
418 let values = [
419 crop.row as f64 * scale_y,
420 crop.column as f64 * scale_x,
421 crop.height as f64 * scale_y,
422 crop.width as f64 * scale_x,
423 ];
424 if values
425 .iter()
426 .any(|value| (value - value.round()).abs() > 1e-9)
427 {
428 return Err(Error::Dataset(
429 "measurement crop does not map to integer reconstruction pixels".into(),
430 ));
431 }
432 Rect::new(
433 values[0].round() as usize,
434 values[1].round() as usize,
435 values[2].round() as usize,
436 values[3].round() as usize,
437 )
438}
439
440fn crop_array<T: Copy>(source: &Array2<T>, crop: Rect) -> Result<Array2<T>> {
441 validate_crop(crop, source.dim())?;
442 let mut output = Vec::with_capacity(checked_len_2d((crop.height, crop.width))?);
443 for row in crop.row..crop.row + crop.height {
444 for column in crop.column..crop.column + crop.width {
445 output.push(source[(row, column)]);
446 }
447 }
448 Ok(Array2::from_shape_vec((crop.height, crop.width), output)?)
449}
450
451fn crop_optional_shared(
452 source: Option<&[f64]>,
453 shape: (usize, usize),
454 crop: Rect,
455) -> Result<Option<Vec<f64>>> {
456 source
457 .map(|source| {
458 let mut output = Vec::with_capacity(checked_len_2d((crop.height, crop.width))?);
459 crop_frame(source, shape, crop, &mut output);
460 Ok(output)
461 })
462 .transpose()
463}
464
465fn crop_optional_frames(
466 source: Option<&[f64]>,
467 shape: (usize, usize),
468 frame_count: usize,
469 indices: &[usize],
470 crop: Rect,
471) -> Result<Option<Vec<f64>>> {
472 let Some(source) = source else {
473 return Ok(None);
474 };
475 let frame_len = checked_len_2d(shape)?;
476 if source.len() == frame_len {
477 return crop_optional_shared(Some(source), shape, crop);
478 }
479 let stack_len = frame_len
480 .checked_mul(frame_count)
481 .ok_or_else(|| Error::ShapeOverflow {
482 shape: vec![frame_count, shape.0, shape.1],
483 })?;
484 if source.len() != stack_len {
485 return Err(Error::InvalidMeasurements(
486 "source background length is inconsistent".into(),
487 ));
488 }
489 let output_len = checked_len_2d((crop.height, crop.width))?
490 .checked_mul(indices.len())
491 .ok_or_else(|| Error::ShapeOverflow {
492 shape: vec![indices.len(), crop.height, crop.width],
493 })?;
494 let mut output = Vec::with_capacity(output_len);
495 for &index in indices {
496 let start = index
497 .checked_mul(frame_len)
498 .ok_or_else(|| Error::ShapeOverflow {
499 shape: vec![index, shape.0, shape.1],
500 })?;
501 let end = start
502 .checked_add(frame_len)
503 .ok_or_else(|| Error::ShapeOverflow {
504 shape: vec![index.saturating_add(1), shape.0, shape.1],
505 })?;
506 crop_frame(
507 source.get(start..end).ok_or_else(|| {
508 Error::InvalidMeasurements("source background frame is out of range".into())
509 })?,
510 shape,
511 crop,
512 &mut output,
513 );
514 }
515 Ok(Some(output))
516}
517
518fn crop_optional_masks(
519 source: Option<&[u8]>,
520 shape: (usize, usize),
521 frame_count: usize,
522 indices: &[usize],
523 crop: Rect,
524) -> Result<Option<Vec<u8>>> {
525 let Some(source) = source else {
526 return Ok(None);
527 };
528 let frame_len = checked_len_2d(shape)?;
529 if source.len() == frame_len {
530 let mut output = Vec::with_capacity(checked_len_2d((crop.height, crop.width))?);
531 crop_frame(source, shape, crop, &mut output);
532 return Ok(Some(output));
533 }
534 let stack_len = frame_len
535 .checked_mul(frame_count)
536 .ok_or_else(|| Error::ShapeOverflow {
537 shape: vec![frame_count, shape.0, shape.1],
538 })?;
539 if source.len() != stack_len {
540 return Err(Error::InvalidMeasurements(
541 "source mask length is inconsistent".into(),
542 ));
543 }
544 let output_len = checked_len_2d((crop.height, crop.width))?
545 .checked_mul(indices.len())
546 .ok_or_else(|| Error::ShapeOverflow {
547 shape: vec![indices.len(), crop.height, crop.width],
548 })?;
549 let mut output = Vec::with_capacity(output_len);
550 for &index in indices {
551 let start = index
552 .checked_mul(frame_len)
553 .ok_or_else(|| Error::ShapeOverflow {
554 shape: vec![index, shape.0, shape.1],
555 })?;
556 let end = start
557 .checked_add(frame_len)
558 .ok_or_else(|| Error::ShapeOverflow {
559 shape: vec![index.saturating_add(1), shape.0, shape.1],
560 })?;
561 crop_frame(
562 source.get(start..end).ok_or_else(|| {
563 Error::InvalidMeasurements("source mask frame is out of range".into())
564 })?,
565 shape,
566 crop,
567 &mut output,
568 );
569 }
570 Ok(Some(output))
571}