Skip to main content

fpm_rs/datasets/
subset.rs

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/// Axis-aligned detector/object crop using zero-based pixel indices.
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub struct Rect {
19    /// First included row (`y`).
20    pub row: usize,
21    /// First included column (`x`).
22    pub column: usize,
23    /// Positive number of rows.
24    pub height: usize,
25    /// Positive number of columns.
26    pub width: usize,
27}
28
29impl Rect {
30    /// Creates a non-empty rectangle; containment is checked when a subset is built.
31    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/// Deterministic acquisition-frame selection for a dataset subset.
47#[derive(Clone, Debug, PartialEq, Eq)]
48pub enum FrameSelector {
49    /// Retain every acquisition frame in original order.
50    All,
51    /// Retain indices `0, step, 2*step, ...`; `step` must be positive.
52    EveryNth(usize),
53    /// Retain explicit unique acquisition-frame indices in the supplied order.
54    Indices(Vec<usize>),
55}
56
57/// Owned dataset values after deterministic frame selection and optional spatial crop.
58#[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    /// Borrows subsetted resident measurements.
70    pub fn measurements(&self) -> &MeasurementStack {
71        &self.measurements
72    }
73
74    /// Borrows models and descriptions updated for the subset dimensions and sources.
75    pub fn configuration(&self) -> &SimulationConfiguration {
76        &self.configuration
77    }
78
79    /// Borrows optional cropped high-resolution complex ground truth.
80    pub fn ground_truth_object(&self) -> Option<&Array2<Complex64>> {
81        self.ground_truth_object.as_ref()
82    }
83
84    /// Borrows optional cropped binary object-validity mask.
85    pub fn valid_object_mask(&self) -> Option<&Array2<u8>> {
86        self.valid_object_mask.as_ref()
87    }
88
89    /// Borrows original provenance plus deterministic subset annotations.
90    pub fn provenance(&self) -> &std::collections::BTreeMap<String, String> {
91        &self.provenance
92    }
93
94    /// Borrows unchanged semantic measurement units.
95    pub fn measurement_units(&self) -> Option<&str> {
96        self.measurement_units.as_deref()
97    }
98
99    /// Clones subset measurements and assumed model into a validated reconstruction problem.
100    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/// Borrowing builder for deterministic dataset frame and spatial subsets.
112#[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    /// Selects acquisition frames; validation occurs in [`Self::build`].
129    pub fn frames(mut self, selector: FrameSelector) -> Self {
130        self.frames = selector;
131        self
132    }
133
134    /// Selects acquisition indices `0, step, 2*step, ...`.
135    pub fn every_nth_frame(self, step: usize) -> Self {
136        self.frames(FrameSelector::EveryNth(step))
137    }
138
139    /// Selects a low-resolution detector rectangle; corresponding object/model grids are updated.
140    pub fn crop(mut self, crop: Rect) -> Self {
141        self.crop = Some(crop);
142        self
143    }
144
145    /// Creates and selects a detector crop from zero-based row/column and positive size.
146    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    /// Validates selection and crop bounds, then owns subsetted measurements, models,
157    /// optional ground truth/mask, and provenance.
158    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}