Skip to main content

qsm_core/bgremove/
ismv.rs

1//! Iterative Spherical Mean Value (iSMV) background field removal
2//!
3//! Iterative approach that avoids mask erosion by iteratively
4//! correcting boundary values.
5//!
6//! Reference:
7//! Wen, Y., Zhou, D., Liu, T., Spincemaille, P., Wang, Y. (2014).
8//! "An iterative spherical mean value method for background field removal in MRI."
9//! Magnetic Resonance in Medicine, 72(4):1065-1071. https://doi.org/10.1002/mrm.24998
10//!
11//! Reference implementation: https://github.com/kamesy/QSM.jl
12
13use num_complex::Complex64;
14use crate::Grid;
15use crate::fft::{fft3d, ifft3d, fft_real_kernel};
16use crate::kernels::smv::{smv_kernel, erode_mask_smv};
17use crate::utils::vec_norm;
18
19/// iSMV algorithm parameters
20#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
21#[derive(Clone, Debug)]
22pub struct IsmvParams {
23    /// Convergence tolerance
24    pub tol: f64,
25    /// Maximum iterations
26    pub max_iter: usize,
27    /// SMV kernel radius in mm (default: 5.0)
28    pub radius: f64,
29}
30
31impl Default for IsmvParams {
32    fn default() -> Self {
33        Self { tol: 1e-6, max_iter: 50, radius: 5.0 }
34    }
35}
36
37/// iSMV background field removal
38///
39/// # Arguments
40/// * `field` - Total field (nx * ny * nz)
41/// * `mask` - Binary mask (nx * ny * nz), 1 = brain, 0 = background
42/// * `grid` - Volume dimensions and voxel sizes
43/// * `params` - iSMV parameters (tolerance, max iterations, radius in mm)
44/// * `progress` - Progress callback (iteration, max_iter)
45///
46/// # Returns
47/// Tuple of (local field, eroded mask)
48pub fn ismv(
49    field: &[f64],
50    mask: &[u8],
51    grid: &Grid,
52    params: &IsmvParams,
53    progress: impl FnMut(usize, usize),
54) -> (Vec<f64>, Vec<u8>) {
55    ismv_core(field, mask, grid, params.radius, params.tol, params.max_iter, progress)
56}
57
58/// iSMV with an explicit absolute kernel radius (mm).
59///
60/// Internal entry point; the public [`ismv`] wrapper derives the radius from
61/// [`IsmvParams`].
62pub(crate) fn ismv_core(
63    field: &[f64],
64    mask: &[u8],
65    grid: &Grid,
66    radius: f64,
67    tol: f64,
68    max_iter: usize,
69    mut progress: impl FnMut(usize, usize),
70) -> (Vec<f64>, Vec<u8>) {
71    let (nx, ny, nz) = grid.dims;
72    let n_total = nx * ny * nz;
73
74    // Generate SMV kernel
75    let smv = smv_kernel(grid, radius);
76    let smv_fft_real = fft_real_kernel(&smv, nx, ny, nz);
77
78    // Complex FFT needed for iteration loop
79    let mut smv_complex: Vec<Complex64> = smv.iter()
80        .map(|&x| Complex64::new(x, 0.0))
81        .collect();
82    fft3d(&mut smv_complex, nx, ny, nz);
83    let smv_fft = smv_complex;
84
85    // Convert mask to f64
86    let m0: Vec<f64> = mask.iter()
87        .map(|&m| if m != 0 { 1.0 } else { 0.0 })
88        .collect();
89
90    // Erode mask using SMV
91    let eroded_mask = erode_mask_f64(mask, &smv_fft_real, grid);
92
93    // Boundary mask: original mask minus eroded mask
94    let boundary: Vec<f64> = m0.iter()
95        .zip(eroded_mask.iter())
96        .map(|(&m, &e)| m - e)
97        .collect();
98
99    // Initialize: f = field
100    let mut f: Vec<f64> = field.to_vec();
101
102    // f0 = eroded_mask * field (for residual calculation)
103    let mut f0: Vec<f64> = field.iter()
104        .zip(eroded_mask.iter())
105        .map(|(&fi, &m)| fi * m)
106        .collect();
107
108    // Boundary correction: bc = boundary * field
109    let bc: Vec<f64> = field.iter()
110        .zip(boundary.iter())
111        .map(|(&fi, &b)| fi * b)
112        .collect();
113
114    // Initial residual norm
115    let mut nr = vec_norm(&f0);
116    let eps = tol * nr;
117
118    // iSMV iterations
119    for iter in 0..max_iter {
120        // Report progress
121        progress(iter + 1, max_iter);
122
123        if nr <= eps {
124            progress(iter + 1, iter + 1);
125            break;
126        }
127
128        // f = SMV(f)
129        let mut f_complex: Vec<Complex64> = f.iter()
130            .map(|&x| Complex64::new(x, 0.0))
131            .collect();
132
133        fft3d(&mut f_complex, nx, ny, nz);
134
135        for i in 0..n_total {
136            f_complex[i] *= smv_fft[i];
137        }
138
139        ifft3d(&mut f_complex, nx, ny, nz);
140
141        // f = eroded_mask * f + bc
142        for i in 0..n_total {
143            f[i] = eroded_mask[i] * f_complex[i].re + bc[i];
144        }
145
146        // Compute residual: ||f0 - f||
147        let mut residual_sq = 0.0;
148        for i in 0..n_total {
149            let diff = f0[i] - f[i];
150            residual_sq += diff * diff;
151            f0[i] = f[i];
152        }
153        nr = residual_sq.sqrt();
154    }
155
156    // Compute local field: m * (field - f)
157    let mut local_field = vec![0.0; n_total];
158    for i in 0..n_total {
159        if mask[i] != 0 {
160            local_field[i] = field[i] - f[i];
161        }
162    }
163
164    // Convert eroded mask to u8
165    let eroded_mask_u8: Vec<u8> = eroded_mask.iter()
166        .map(|&m| if m > 0.5 { 1 } else { 0 })
167        .collect();
168
169    (local_field, eroded_mask_u8)
170}
171
172/// Erode mask using SMV convolution (returns f64 for arithmetic compatibility)
173fn erode_mask_f64(mask: &[u8], smv_kernel_fft: &[f64], grid: &Grid) -> Vec<f64> {
174    let eroded = erode_mask_smv(mask, smv_kernel_fft, grid, 1.0 - 1e-10);
175    eroded.iter().map(|&m| m as f64).collect()
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181
182    #[test]
183    fn test_ismv_zero_field() {
184        let n = 8;
185        let field = vec![0.0; n * n * n];
186        let mask = vec![1u8; n * n * n];
187        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
188
189        let (local, eroded) = ismv_core(
190            &field, &mask, &grid,
191            2.0, 1e-3, 10, |_, _| {}
192        );
193
194        for &val in local.iter() {
195            assert!(val.abs() < 1e-10, "Zero field should give zero local field");
196        }
197
198        // Some voxels should be in eroded mask
199        let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
200        assert!(eroded_count > 0, "Eroded mask should have some voxels");
201    }
202
203    #[test]
204    fn test_ismv_finite() {
205        let n = 8;
206        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
207        let mask = vec![1u8; n * n * n];
208        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
209
210        let (local, _eroded) = ismv_core(
211            &field, &mask, &grid,
212            2.0, 1e-3, 20, |_, _| {}
213        );
214
215        for (i, &val) in local.iter().enumerate() {
216            assert!(val.is_finite(), "Local field should be finite at index {}", i);
217        }
218    }
219
220    #[test]
221    fn test_ismv_preserves_interior() {
222        // iSMV should preserve some of the mask interior
223        let n = 16;
224        let field = vec![0.1; n * n * n];
225
226        // Create a spherical mask
227        let mut mask = vec![0u8; n * n * n];
228        let center = n / 2;
229        let radius = n / 3;
230
231        for i in 0..n {
232            for j in 0..n {
233                for k in 0..n {
234                    let di = (i as i32) - (center as i32);
235                    let dj = (j as i32) - (center as i32);
236                    let dk = (k as i32) - (center as i32);
237                    if di*di + dj*dj + dk*dk <= (radius * radius) as i32 {
238                        mask[i * n * n + j * n + k] = 1;
239                    }
240                }
241            }
242        }
243
244        let mask_count: usize = mask.iter().map(|&m| m as usize).sum();
245        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
246
247        // Use small radius for less erosion
248        let (_, eroded) = ismv_core(
249            &field, &mask, &grid,
250            1.5, 1e-3, 50, |_, _| {}
251        );
252
253        let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
254
255        // Eroded mask should have fewer voxels than original
256        assert!(eroded_count <= mask_count, "Eroded mask should be smaller than original");
257        // Should preserve at least some interior voxels
258        assert!(eroded_count > 0, "Eroded mask should have some voxels");
259    }
260
261    #[test]
262    fn test_ismv_convergence() {
263        let n = 8;
264        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
265        let mask = vec![1u8; n * n * n];
266        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
267
268        // Run with more iterations and tight tolerance
269        let (local_many, _) = ismv_core(
270            &field, &mask, &grid,
271            2.0, 1e-6, 100, |_, _| {}
272        );
273
274        // Run with fewer iterations
275        let (local_few, _) = ismv_core(
276            &field, &mask, &grid,
277            2.0, 1e-6, 5, |_, _| {}
278        );
279
280        // Both should be finite
281        for (i, &val) in local_many.iter().enumerate() {
282            assert!(val.is_finite(), "iSMV many iters: finite at index {}", i);
283        }
284        for (i, &val) in local_few.iter().enumerate() {
285            assert!(val.is_finite(), "iSMV few iters: finite at index {}", i);
286        }
287
288        // More iterations should give different (hopefully more converged) result
289        // or the same if already converged
290        let diff_norm: f64 = local_many.iter()
291            .zip(local_few.iter())
292            .map(|(&a, &b)| (a - b).powi(2))
293            .sum::<f64>()
294            .sqrt();
295
296        // The difference should be finite (no NaN/Inf)
297        assert!(diff_norm.is_finite(), "Difference between runs should be finite");
298    }
299
300    #[test]
301    fn test_ismv_different_radius() {
302        let n = 8;
303        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
304        let mask = vec![1u8; n * n * n];
305        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
306
307        // Small radius
308        let (local_small, eroded_small) = ismv_core(
309            &field, &mask, &grid,
310            1.5, 1e-3, 20, |_, _| {}
311        );
312
313        // Larger radius
314        let (local_large, eroded_large) = ismv_core(
315            &field, &mask, &grid,
316            3.0, 1e-3, 20, |_, _| {}
317        );
318
319        // Both should produce finite results
320        for (i, &val) in local_small.iter().enumerate() {
321            assert!(val.is_finite(), "iSMV small radius: finite at index {}", i);
322        }
323        for (i, &val) in local_large.iter().enumerate() {
324            assert!(val.is_finite(), "iSMV large radius: finite at index {}", i);
325        }
326
327        // Larger radius should erode more
328        let small_count: usize = eroded_small.iter().map(|&m| m as usize).sum();
329        let large_count: usize = eroded_large.iter().map(|&m| m as usize).sum();
330        assert!(
331            large_count <= small_count,
332            "Larger radius should erode more: large={}, small={}",
333            large_count, small_count
334        );
335    }
336
337    #[test]
338    fn test_ismv_larger_volume() {
339        // Test with 16x16x16 volume and spherical mask
340        let n = 16;
341
342        // Create a field with a linear background
343        let mut field = vec![0.0; n * n * n];
344        for z in 0..n {
345            for y in 0..n {
346                for x in 0..n {
347                    field[x + y * n + z * n * n] = (z as f64) * 0.1;
348                }
349            }
350        }
351
352        // Spherical mask
353        let mut mask = vec![0u8; n * n * n];
354        let center = n / 2;
355        let radius = n / 3;
356        for z in 0..n {
357            for y in 0..n {
358                for x in 0..n {
359                    let dx = (x as i32) - (center as i32);
360                    let dy = (y as i32) - (center as i32);
361                    let dz = (z as i32) - (center as i32);
362                    if dx * dx + dy * dy + dz * dz <= (radius * radius) as i32 {
363                        mask[x + y * n + z * n * n] = 1;
364                    }
365                }
366            }
367        }
368
369        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
370        let (local, eroded) = ismv_core(
371            &field, &mask, &grid,
372            2.0, 1e-3, 50, |_, _| {}
373        );
374
375        assert_eq!(local.len(), n * n * n);
376        for &val in &local {
377            assert!(val.is_finite());
378        }
379
380        // Eroded mask should be non-empty
381        let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
382        assert!(eroded_count > 0, "Eroded mask should have some voxels");
383
384        // Outside original mask should be zero
385        for i in 0..n * n * n {
386            if mask[i] == 0 {
387                assert_eq!(local[i], 0.0, "Outside mask should be zero");
388            }
389        }
390    }
391
392    #[test]
393    fn test_ismv_with_progress() {
394        let n = 8;
395        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
396        let mask = vec![1u8; n * n * n];
397        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
398
399        let mut progress_calls = Vec::new();
400        let (local, _) = ismv_core(
401            &field, &mask, &grid,
402            2.0, 1e-3, 20,
403            |iter, max| { progress_calls.push((iter, max)); }
404        );
405
406        assert_eq!(local.len(), n * n * n);
407        assert!(!progress_calls.is_empty(), "Progress callback should be called");
408        for &val in &local {
409            assert!(val.is_finite());
410        }
411    }
412
413    #[test]
414    fn test_ismv_anisotropic_voxels() {
415        let n = 8;
416        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
417        let mask = vec![1u8; n * n * n];
418        let grid = Grid::new(n, n, n, 0.5, 1.0, 2.0);
419
420        // Anisotropic voxel sizes
421        let (local, eroded) = ismv_core(
422            &field, &mask, &grid,
423            3.0, 1e-3, 20, |_, _| {}
424        );
425
426        for &val in &local {
427            assert!(val.is_finite());
428        }
429        // Should still produce some eroded mask
430        let count: usize = eroded.iter().map(|&m| m as usize).sum();
431        assert!(count <= n * n * n);
432    }
433
434    #[test]
435    fn test_ismv_tight_convergence() {
436        // Test tight convergence to ensure the convergence check branch is hit
437        let n = 8;
438        let field = vec![0.5; n * n * n]; // Constant field
439        let mask = vec![1u8; n * n * n];
440        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
441
442        let (local, _) = ismv_core(
443            &field, &mask, &grid,
444            2.0, 1e-12, 200, |_, _| {} // Very tight tolerance, many iterations allowed
445        );
446
447        for &val in &local {
448            assert!(val.is_finite());
449        }
450    }
451
452    #[test]
453    fn test_ismv_with_background_mask() {
454        // Test where some voxels are outside the mask
455        let n = 8;
456        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
457
458        // Mask with zero border
459        let mut mask = vec![0u8; n * n * n];
460        for z in 1..(n - 1) {
461            for y in 1..(n - 1) {
462                for x in 1..(n - 1) {
463                    mask[x + y * n + z * n * n] = 1;
464                }
465            }
466        }
467
468        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
469        let (local, eroded) = ismv_core(
470            &field, &mask, &grid,
471            1.5, 1e-3, 30, |_, _| {}
472        );
473
474        for &val in &local {
475            assert!(val.is_finite());
476        }
477
478        // Outside mask should be zero
479        for i in 0..n * n * n {
480            if mask[i] == 0 {
481                assert_eq!(local[i], 0.0, "Outside mask should be zero at index {}", i);
482            }
483        }
484
485        // Eroded mask should be subset of original mask
486        for i in 0..n * n * n {
487            if eroded[i] != 0 {
488                assert_eq!(mask[i], 1, "Eroded voxel must be inside original mask");
489            }
490        }
491    }
492}