Skip to main content

qsm_core/pipeline/
masking.rs

1//! Masking stage
2//!
3//! Generates a binary brain mask from magnitude/phase data using a
4//! configurable sequence of operations (threshold, BET, morphological ops).
5//! Multiple mask sections can be OR'd together.
6
7use super::config::*;
8use super::phase_utils::{erode_mask, dilate_mask};
9
10/// Resolve masking input data based on the MaskingInput type.
11///
12/// # Arguments
13/// * `input` - Which data source to use
14/// * `phases` - Per-echo phase arrays
15/// * `magnitudes` - Per-echo magnitude arrays (already resolved: RSS for Magnitude,
16///   specific echo for First/Last, optionally homogeneity-corrected)
17/// * `metadata` - Scan metadata
18///
19/// # Returns
20/// The input data array to threshold/mask from
21pub fn resolve_masking_input(
22    input: MaskingInput,
23    phases: &[&[f64]],
24    magnitude: Option<&[f64]>,
25    metadata: &ScanMetadata,
26) -> Vec<f64> {
27    let (nx, ny, nz) = metadata.dims;
28    let n_voxels = nx * ny * nz;
29
30    match input {
31        MaskingInput::MagnitudeFirst | MaskingInput::Magnitude | MaskingInput::MagnitudeLast => {
32            magnitude.map(|m| m.to_vec()).unwrap_or_else(|| vec![0.0; n_voxels])
33        }
34        MaskingInput::PhaseQuality => {
35            if phases.is_empty() {
36                return vec![0.0; n_voxels];
37            }
38            let all_ones = vec![1u8; n_voxels];
39            let mag = magnitude.unwrap_or(&[]);
40            let mag_data: Vec<f64> = if mag.is_empty() {
41                vec![1.0; n_voxels]
42            } else {
43                mag.to_vec()
44            };
45
46            let grid = metadata.grid();
47            if phases.len() >= 2 && metadata.echo_times.len() >= 2 {
48                crate::unwrap::voxel_quality_romeo(
49                    phases[0], &mag_data,
50                    Some(phases[1]),
51                    metadata.echo_times[0], metadata.echo_times[1],
52                    &all_ones, &grid,
53                )
54            } else {
55                crate::unwrap::voxel_quality_romeo(
56                    phases[0], &mag_data,
57                    None,
58                    metadata.echo_times.first().copied().unwrap_or(0.02),
59                    0.0, &all_ones, &grid,
60                )
61            }
62        }
63    }
64}
65
66/// Build a mask from a single section (generator + refinements).
67///
68/// # Arguments
69/// * `section` - Mask section config (input type, generator, refinements)
70/// * `input_data` - Pre-resolved input data (from `resolve_masking_input`)
71/// * `magnitude` - Magnitude data for BET (optional)
72/// * `metadata` - Scan metadata
73pub fn build_mask_section(
74    section: &MaskSection,
75    input_data: &[f64],
76    magnitude: Option<&[f64]>,
77    metadata: &ScanMetadata,
78) -> Result<Vec<u8>, PipelineError> {
79    let n_voxels = metadata.dims.0 * metadata.dims.1 * metadata.dims.2;
80    apply_mask_ops(vec![1u8; n_voxels], &section.all_ops(), input_data, magnitude, metadata)
81}
82
83/// Apply mask operations to an existing mask, in order.
84///
85/// This is the one implementation of what every mask op *means*; [`build_mask_section`] is this
86/// function starting from an all-ones mask, and hosts that drive masking themselves (e.g. an
87/// interactive UI applying one refinement at a time) should call it rather than reimplement the
88/// operations, so a mask built step-by-step matches the one a `--mask` section would produce.
89///
90/// **Two different images.** `input_data` is what a generator looks at — thresholding may use a
91/// phase-quality map, for instance. `magnitude` is the magnitude image, and is what the ops that
92/// need real signal use: [`MaskOp::Bet`], [`MaskOp::HdBet`] and [`MaskOp::SignalErode`]. Passing
93/// the section input as `magnitude` is a bug: signal-gated erosion divides out a receive-coil bias
94/// estimate and gates on the in-mask median, which only means anything for a magnitude image.
95/// Those ops error when `magnitude` is `None`.
96///
97/// # Arguments
98/// * `mask` - Starting mask (0/1), length `nx*ny*nz`
99/// * `ops` - Operations to apply, in order
100/// * `input_data` - Image the generators threshold (may be a phase-quality map)
101/// * `magnitude` - Magnitude image, for BET / HD-BET / signal-gated erosion
102/// * `metadata` - Scan metadata (dims + voxel size)
103pub fn apply_mask_ops(
104    mask: Vec<u8>,
105    ops: &[MaskOp],
106    input_data: &[f64],
107    magnitude: Option<&[f64]>,
108    metadata: &ScanMetadata,
109) -> Result<Vec<u8>, PipelineError> {
110    let (nx, ny, nz) = metadata.dims;
111    let (vsx, vsy, vsz) = metadata.voxel_size;
112    let grid = metadata.grid();
113    let n_voxels = nx * ny * nz;
114    if mask.len() != n_voxels {
115        return Err(PipelineError::DimensionMismatch { expected: n_voxels, got: mask.len() });
116    }
117    let mut mask = mask;
118
119    for op in ops {
120        match op {
121            MaskOp::Threshold { method, value } => {
122                let threshold = match method {
123                    MaskThresholdMethod::Otsu => {
124                        crate::utils::otsu_threshold(input_data, 256)
125                    }
126                    MaskThresholdMethod::Fixed => value.unwrap_or(0.5),
127                    MaskThresholdMethod::Percentile => {
128                        let pct = value.unwrap_or(75.0) / 100.0;
129                        let mut sorted: Vec<f64> = input_data.iter()
130                            .filter(|v| v.is_finite() && **v > 0.0)
131                            .copied().collect();
132                        sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
133                        if sorted.is_empty() { 0.0 }
134                        else {
135                            let idx = ((sorted.len() as f64 * pct) as usize).min(sorted.len() - 1);
136                            sorted[idx]
137                        }
138                    }
139                };
140                mask = input_data.iter()
141                    .map(|&v| if v > threshold { 1u8 } else { 0u8 })
142                    .collect();
143            }
144            MaskOp::Bet { fractional_intensity } => {
145                let mag_data = magnitude.ok_or_else(|| {
146                    PipelineError::InvalidInput("BET requires magnitude data".into())
147                })?;
148                let bet_params = crate::bet::BetParams {
149                    fractional_intensity: *fractional_intensity,
150                    ..crate::bet::BetParams::default()
151                };
152                let grid = crate::Grid::new(nx, ny, nz, vsx, vsy, vsz);
153                mask = crate::bet::run_bet(mag_data, &grid, &bet_params, |_, _| {});
154            }
155            MaskOp::Erode { iterations } => {
156                mask = erode_mask(&mask, &grid, *iterations);
157            }
158            MaskOp::Dilate { iterations } => {
159                mask = dilate_mask(&mask, &grid, *iterations);
160            }
161            MaskOp::Close { radius } => {
162                mask = crate::utils::morphological_close(&mask, &grid, *radius as i32);
163            }
164            MaskOp::FillHoles { max_size } => {
165                let effective_size = if *max_size == 0 { n_voxels / 20 } else { *max_size };
166                mask = crate::utils::fill_holes(&mask, &grid, effective_size);
167            }
168            MaskOp::GaussianSmooth { sigma_mm } => {
169                let sigma = *sigma_mm;
170                let mask_f64: Vec<f64> = mask.iter().map(|&m| m as f64).collect();
171                let smoothed = crate::utils::gaussian_smooth_3d(
172                    &mask_f64,
173                    [sigma, sigma, sigma],
174                    None, None, 3,
175                    &grid,
176                );
177                mask = smoothed.iter().map(|&v| if v > 0.5 { 1u8 } else { 0u8 }).collect();
178            }
179            MaskOp::SignalErode(params) => {
180                let mag_data = magnitude.ok_or_else(|| {
181                    PipelineError::InvalidInput("signal-gated erosion requires magnitude data".into())
182                })?;
183                mask = crate::utils::signal_gated_erosion(&mask, mag_data, &grid, params);
184            }
185            MaskOp::HdBet(params) => {
186                let mag_data = magnitude.ok_or_else(|| {
187                    PipelineError::InvalidInput("HD-BET requires magnitude data".into())
188                })?;
189                mask = run_hd_bet(mag_data, &grid, params)?;
190            }
191        }
192    }
193
194    Ok(mask)
195}
196
197/// Source the HD-BET weights and run it. Requires the `onnx` feature; weights come from the
198/// model registry (local `$QSM_MODEL_DIR`/cache, or the `download` feature).
199#[cfg(feature = "onnx")]
200fn run_hd_bet(
201    magnitude: &[f64],
202    grid: &crate::Grid,
203    params: &crate::bet::HdBetParams,
204) -> Result<Vec<u8>, PipelineError> {
205    let bytes = crate::models::primary_weight("hd-bet").map_err(PipelineError::InvalidConfig)?;
206    crate::bet::hd_bet(magnitude, grid, &bytes, params, |_, _| {})
207        .map_err(|e| PipelineError::AlgorithmError(e.to_string()))
208}
209
210#[cfg(not(feature = "onnx"))]
211fn run_hd_bet(
212    _magnitude: &[f64],
213    _grid: &crate::Grid,
214    _params: &crate::bet::HdBetParams,
215) -> Result<Vec<u8>, PipelineError> {
216    Err(PipelineError::InvalidConfig(
217        "HD-BET requires building qsm-core with the 'onnx' feature".into(),
218    ))
219}
220
221/// Build a mask from multiple sections, OR'd together.
222///
223/// Each section specifies an input source, a generator (threshold/BET),
224/// and optional refinements (erode, dilate, close, fill holes, smooth).
225///
226/// # Arguments
227/// * `sections` - Mask section configs
228/// * `phases` - Per-echo phase arrays (for PhaseQuality input)
229/// * `magnitude` - Combined magnitude (for Magnitude/BET input)
230/// * `metadata` - Scan metadata
231pub fn run_masking(
232    sections: &[MaskSection],
233    phases: &[&[f64]],
234    magnitude: Option<&[f64]>,
235    metadata: &ScanMetadata,
236) -> Result<Vec<u8>, PipelineError> {
237    let (nx, ny, nz) = metadata.dims;
238    let n_voxels = nx * ny * nz;
239
240    if sections.is_empty() {
241        return Err(PipelineError::InvalidConfig("no mask sections configured".into()));
242    }
243
244    if sections.len() == 1 {
245        let input_data = resolve_masking_input(sections[0].input, phases, magnitude, metadata);
246        return build_mask_section(&sections[0], &input_data, magnitude, metadata);
247    }
248
249    // Multiple sections: run each, OR together
250    let mut final_mask = vec![0u8; n_voxels];
251    for section in sections {
252        let input_data = resolve_masking_input(section.input, phases, magnitude, metadata);
253        let section_mask = build_mask_section(section, &input_data, magnitude, metadata)?;
254        for j in 0..n_voxels {
255            final_mask[j] |= section_mask[j];
256        }
257    }
258
259    Ok(final_mask)
260}
261
262#[cfg(test)]
263mod tests {
264    use super::*;
265
266    fn test_metadata() -> ScanMetadata {
267        ScanMetadata {
268            dims: (8, 8, 8),
269            voxel_size: (1.0, 1.0, 1.0),
270            echo_times: vec![0.005, 0.010],
271            field_strength: 3.0,
272            b0_direction: (0.0, 0.0, 1.0),
273        }
274    }
275
276    #[test]
277    fn test_run_masking_otsu_threshold() {
278        let meta = test_metadata();
279        let n = 8 * 8 * 8;
280        // Half bright, half dark
281        let mut mag = vec![0.1; n];
282        for i in n / 2..n {
283            mag[i] = 10.0;
284        }
285
286        let sections = vec![MaskSection {
287            input: MaskingInput::Magnitude,
288            generator: MaskOp::Threshold {
289                method: MaskThresholdMethod::Otsu,
290                value: None,
291            },
292            refinements: vec![],
293        }];
294
295        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
296        assert_eq!(result.len(), n);
297
298        // Bright half should be masked in, dark half masked out
299        let bright_count: usize = result[n / 2..].iter().map(|&m| m as usize).sum();
300        let dark_count: usize = result[..n / 2].iter().map(|&m| m as usize).sum();
301        assert!(bright_count > dark_count, "Otsu should separate bright from dark");
302    }
303
304    #[test]
305    fn test_run_masking_with_refinements() {
306        let meta = test_metadata();
307        let n = 8 * 8 * 8;
308        let mag = vec![10.0; n]; // all bright → all masked in
309
310        let sections = vec![MaskSection {
311            input: MaskingInput::Magnitude,
312            generator: MaskOp::Threshold {
313                method: MaskThresholdMethod::Fixed,
314                value: Some(0.5),
315            },
316            refinements: vec![
317                MaskOp::Erode { iterations: 1 },
318                MaskOp::Dilate { iterations: 1 },
319            ],
320        }];
321
322        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
323        assert_eq!(result.len(), n);
324        // After erode+dilate, interior should still be masked
325        let (nx, ny, nz) = meta.dims;
326        let center = nx / 2 + (ny / 2) * nx + (nz / 2) * nx * ny;
327        assert_eq!(result[center], 1, "center should survive erode+dilate");
328    }
329
330    #[test]
331    fn test_run_masking_or_sections() {
332        let meta = test_metadata();
333        let n = 8 * 8 * 8;
334        let mag = vec![10.0; n];
335
336        // Section 1: mask only first half via fixed threshold
337        // Section 2: mask only second half
338        // OR → should get everything
339        let sections = vec![
340            MaskSection {
341                input: MaskingInput::Magnitude,
342                generator: MaskOp::Threshold {
343                    method: MaskThresholdMethod::Fixed,
344                    value: Some(0.5),
345                },
346                refinements: vec![],
347            },
348            MaskSection {
349                input: MaskingInput::Magnitude,
350                generator: MaskOp::Threshold {
351                    method: MaskThresholdMethod::Fixed,
352                    value: Some(0.5),
353                },
354                refinements: vec![],
355            },
356        ];
357
358        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
359        let count: usize = result.iter().map(|&m| m as usize).sum();
360        assert_eq!(count, n, "OR of identical sections should give full mask");
361    }
362
363    #[test]
364    fn test_masking_fixed_threshold() {
365        let meta = test_metadata();
366        let n = 8 * 8 * 8;
367        let mag = vec![10.0; n];
368        let sections = vec![MaskSection {
369            input: MaskingInput::Magnitude,
370            generator: MaskOp::Threshold { method: MaskThresholdMethod::Fixed, value: Some(5.0) },
371            refinements: vec![],
372        }];
373        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
374        let count: usize = result.iter().map(|&m| m as usize).sum();
375        assert_eq!(count, n, "all voxels above threshold=5");
376    }
377
378    #[test]
379    fn test_masking_percentile_threshold() {
380        let meta = test_metadata();
381        let n = 8 * 8 * 8;
382        let mag: Vec<f64> = (0..n).map(|i| i as f64).collect();
383        let sections = vec![MaskSection {
384            input: MaskingInput::Magnitude,
385            generator: MaskOp::Threshold { method: MaskThresholdMethod::Percentile, value: Some(50.0) },
386            refinements: vec![],
387        }];
388        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
389        let count: usize = result.iter().map(|&m| m as usize).sum();
390        assert!(count > 0 && count < n, "percentile should mask ~half");
391    }
392
393    #[test]
394    fn test_masking_close_and_fill_holes() {
395        let meta = test_metadata();
396        let n = 8 * 8 * 8;
397        let mag = vec![10.0; n];
398        let sections = vec![MaskSection {
399            input: MaskingInput::Magnitude,
400            generator: MaskOp::Threshold { method: MaskThresholdMethod::Fixed, value: Some(0.5) },
401            refinements: vec![
402                MaskOp::Close { radius: 1 },
403                MaskOp::FillHoles { max_size: 0 },
404                MaskOp::GaussianSmooth { sigma_mm: 1.0 },
405            ],
406        }];
407        let result = run_masking(&sections, &[], Some(&mag), &meta).unwrap();
408        assert_eq!(result.len(), n);
409    }
410
411    #[test]
412    fn test_masking_phase_quality_input() {
413        let meta = test_metadata();
414        let n = 8 * 8 * 8;
415        let phase1 = vec![0.5; n];
416        let phase2 = vec![1.0; n];
417        let mag = vec![1.0; n];
418        let sections = vec![MaskSection {
419            input: MaskingInput::PhaseQuality,
420            generator: MaskOp::Threshold { method: MaskThresholdMethod::Otsu, value: None },
421            refinements: vec![],
422        }];
423        let result = run_masking(&sections, &[&phase1, &phase2], Some(&mag), &meta).unwrap();
424        assert_eq!(result.len(), n);
425    }
426
427    #[test]
428    fn test_masking_signal_erode() {
429        let meta = test_metadata();
430        let n = 8 * 8 * 8;
431        let mag = vec![10.0; n];
432        let section = |refinement| vec![MaskSection {
433            input: MaskingInput::Magnitude,
434            generator: MaskOp::Threshold { method: MaskThresholdMethod::Fixed, value: Some(5.0) },
435            refinements: vec![refinement],
436        }];
437        let params = crate::utils::SignalErosionParams { min_component: 1, ..Default::default() };
438        let result = run_masking(&section(MaskOp::SignalErode(params.clone())), &[], Some(&mag), &meta).unwrap();
439        // Uniform signal: nothing is gated, so only the one global erosion applies.
440        let plain = run_masking(&section(MaskOp::Erode { iterations: 1 }), &[], Some(&mag), &meta).unwrap();
441        assert_eq!(result, plain);
442        // Needs magnitude.
443        assert!(run_masking(&section(MaskOp::SignalErode(params)), &[], None, &meta).is_err());
444    }
445
446    #[test]
447    fn test_masking_hd_bet_requires_magnitude() {
448        let meta = test_metadata();
449        let sections = vec![MaskSection {
450            input: MaskingInput::Magnitude,
451            generator: MaskOp::HdBet(Default::default()),
452            refinements: vec![],
453        }];
454        assert!(run_masking(&sections, &[], None, &meta).is_err());
455        #[cfg(not(feature = "onnx"))]
456        assert!(run_masking(&sections, &[], Some(&vec![1.0; 512]), &meta).is_err());
457    }
458
459    /// Applying ops one at a time (what an interactive host does) must equal building the whole
460    /// section in one call — same implementation, so a step-by-step mask matches the `--mask`
461    /// section a host prints alongside it.
462    #[test]
463    fn test_apply_mask_ops_matches_build_mask_section() {
464        let meta = test_metadata();
465        let n = 8 * 8 * 8;
466        let input: Vec<f64> = (0..n).map(|i| ((i * 37) % 19) as f64).collect();
467        let mag: Vec<f64> = (0..n).map(|i| 50.0 + ((i * 7) % 13) as f64).collect();
468        let section = MaskSection {
469            input: MaskingInput::Magnitude,
470            generator: MaskOp::Threshold { method: MaskThresholdMethod::Fixed, value: Some(5.0) },
471            refinements: vec![
472                MaskOp::Dilate { iterations: 1 },
473                MaskOp::FillHoles { max_size: 0 },
474                MaskOp::Erode { iterations: 1 },
475                MaskOp::SignalErode(crate::utils::SignalErosionParams { min_component: 1, ..Default::default() }),
476            ],
477        };
478        let whole = build_mask_section(&section, &input, Some(&mag), &meta).unwrap();
479
480        let mut step = vec![1u8; n];
481        for op in section.all_ops() {
482            step = apply_mask_ops(step, std::slice::from_ref(&op), &input, Some(&mag), &meta).unwrap();
483        }
484        assert_eq!(step, whole);
485    }
486
487    #[test]
488    fn test_apply_mask_ops_needs_magnitude_and_matching_length() {
489        let meta = test_metadata();
490        let n = 8 * 8 * 8;
491        let input = vec![10.0; n];
492        // Signal-gated erosion gates on the magnitude, so without one it is an error rather than
493        // silently falling back to the section input.
494        let ops = [MaskOp::SignalErode(Default::default())];
495        assert!(apply_mask_ops(vec![1u8; n], &ops, &input, None, &meta).is_err());
496        // A mask of the wrong size is rejected rather than mis-indexed.
497        assert!(apply_mask_ops(vec![1u8; n - 1], &[MaskOp::Erode { iterations: 1 }], &input, Some(&input), &meta).is_err());
498    }
499
500    #[test]
501    fn test_run_masking_empty_sections() {
502        let meta = test_metadata();
503        let result = run_masking(&[], &[], None, &meta);
504        assert!(result.is_err());
505    }
506}