1use super::config::*;
8use super::phase_utils::{erode_mask, dilate_mask};
9
10pub 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
66pub 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], §ion.all_ops(), input_data, magnitude, metadata)
81}
82
83pub 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#[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
221pub 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(§ions[0], &input_data, magnitude, metadata);
247 }
248
249 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 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(§ions, &[], Some(&mag), &meta).unwrap();
296 assert_eq!(result.len(), n);
297
298 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]; 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(§ions, &[], Some(&mag), &meta).unwrap();
323 assert_eq!(result.len(), n);
324 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 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(§ions, &[], 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(§ions, &[], 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(§ions, &[], 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(§ions, &[], 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(§ions, &[&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(§ion(MaskOp::SignalErode(params.clone())), &[], Some(&mag), &meta).unwrap();
439 let plain = run_masking(§ion(MaskOp::Erode { iterations: 1 }), &[], Some(&mag), &meta).unwrap();
441 assert_eq!(result, plain);
442 assert!(run_masking(§ion(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(§ions, &[], None, &meta).is_err());
455 #[cfg(not(feature = "onnx"))]
456 assert!(run_masking(§ions, &[], Some(&vec![1.0; 512]), &meta).is_err());
457 }
458
459 #[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(§ion, &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 let ops = [MaskOp::SignalErode(Default::default())];
495 assert!(apply_mask_ops(vec![1u8; n], &ops, &input, None, &meta).is_err());
496 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}