qsm_core/bgremove/
sharp.rs1use num_complex::Complex64;
15use crate::Grid;
16use crate::fft::{fft3d, ifft3d, fft_real_kernel};
17use crate::kernels::smv::{smv_kernel, erode_mask_smv};
18
19#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
21#[derive(Clone, Debug)]
22pub struct SharpParams {
23 pub threshold: f64,
25 pub radius: f64,
27}
28
29impl Default for SharpParams {
30 fn default() -> Self {
31 Self {
32 threshold: 0.05,
33 radius: 6.0,
34 }
35 }
36}
37
38pub fn sharp(
52 field: &[f64],
53 mask: &[u8],
54 grid: &Grid,
55 params: &SharpParams,
56) -> (Vec<f64>, Vec<u8>) {
57 sharp_core(field, mask, grid, params.threshold, params.radius)
58}
59
60pub(crate) fn sharp_core(
65 field: &[f64],
66 mask: &[u8],
67 grid: &Grid,
68 threshold: f64,
69 radius: f64,
70) -> (Vec<f64>, Vec<u8>) {
71 let (nx, ny, nz) = grid.dims;
72 let n_total = nx * ny * nz;
73
74 let s_kernel = smv_kernel(grid, radius);
76 let s_fft = fft_real_kernel(&s_kernel, nx, ny, nz);
77
78 let eroded_mask = erode_mask_smv(mask, &s_fft, grid, 1.0 - 1e-7_f64.sqrt());
80
81 let mut field_complex: Vec<Complex64> = field.iter()
89 .map(|&x| Complex64::new(x, 0.0))
90 .collect();
91 fft3d(&mut field_complex, nx, ny, nz);
92
93 for i in 0..n_total {
95 field_complex[i] *= 1.0 - s_fft[i];
96 }
97
98 ifft3d(&mut field_complex, nx, ny, nz);
100
101 for i in 0..n_total {
103 if eroded_mask[i] == 0 {
104 field_complex[i] = Complex64::new(0.0, 0.0);
105 }
106 }
107
108 fft3d(&mut field_complex, nx, ny, nz);
110
111 for i in 0..n_total {
113 let one_minus_s = 1.0 - s_fft[i];
114 if one_minus_s.abs() < threshold {
115 field_complex[i] = Complex64::new(0.0, 0.0);
116 } else {
117 field_complex[i] /= one_minus_s;
118 }
119 }
120
121 ifft3d(&mut field_complex, nx, ny, nz);
123
124 let local_field: Vec<f64> = field_complex.iter()
126 .enumerate()
127 .map(|(i, c)| if eroded_mask[i] == 1 { c.re } else { 0.0 })
128 .collect();
129
130 (local_field, eroded_mask)
131}
132
133#[cfg(test)]
134mod tests {
135 use super::*;
136
137 #[test]
138 fn test_sharp_zero_field() {
139 let n = 16;
141 let field = vec![0.0; n * n * n];
142 let mask = vec![1u8; n * n * n];
143 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
144
145 let (local, _) = sharp(&field, &mask, &grid, &SharpParams { threshold: 0.05, radius: 2.0 });
147
148 for &val in local.iter() {
149 assert!(val.abs() < 1e-8, "Zero field should give zero local field, got {}", val);
150 }
151 }
152
153 #[test]
154 fn test_sharp_finite() {
155 let n = 16;
157 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.01).collect();
158 let mask = vec![1u8; n * n * n];
159 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
160
161 let (local, eroded) = sharp(&field, &mask, &grid, &SharpParams { threshold: 0.05, radius: 2.0 });
162
163 for (i, &val) in local.iter().enumerate() {
164 assert!(val.is_finite(), "Local field should be finite at index {}", i);
165 }
166
167 let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
169 assert!(eroded_count > 0, "Eroded mask should have some voxels");
170 }
171}