Skip to main content

qsm_core/bgremove/
resharp.rs

1//! RESHARP background field removal
2//!
3//! Regularized SHARP — uses Tikhonov regularization instead of TSVD
4//! truncation for more robust deconvolution of the SMV-filtered field.
5//!
6//! Solves: argmin ||M·ifft(C·fft(x)) - M·ifft(C·fft(y))||² + λ||x||²
7//!
8//! Where C = (1 - S) is the high-pass (delta minus SMV) kernel in k-space,
9//! M is the eroded mask, and λ is the Tikhonov regularization parameter.
10//!
11//! Reference:
12//! Sun, H. and Wilman, A.H. (2013).
13//! "Background field removal using spherical mean value filtering and Tikhonov regularization."
14//! Magn Reson Med, 71(3):1151-1157. https://doi.org/10.1002/mrm.24765
15
16use num_complex::Complex64;
17use crate::Grid;
18use crate::fft::{fft3d, ifft3d, fft_real_kernel};
19use crate::kernels::smv::{smv_kernel, erode_mask_smv};
20use crate::solvers::cg_solve_with_progress;
21
22/// RESHARP algorithm parameters
23#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
24#[derive(Clone, Debug)]
25pub struct ResharpParams {
26    /// SMV kernel radius in mm
27    pub radius: f64,
28    /// Tikhonov regularization parameter
29    pub tik_reg: f64,
30    /// CG convergence tolerance
31    pub tol: f64,
32    /// Maximum CG iterations
33    pub max_iter: usize,
34}
35
36impl Default for ResharpParams {
37    fn default() -> Self {
38        Self {
39            radius: 6.0,
40            tik_reg: 1e-4,
41            tol: 1e-6,
42            max_iter: 30,
43        }
44    }
45}
46
47/// RESHARP background field removal
48///
49/// # Arguments
50/// * `field` - Unwrapped total field (nx * ny * nz)
51/// * `mask` - Binary mask (nx * ny * nz), 1 = inside ROI
52/// * `grid` - Volume dimensions and voxel sizes
53/// * `params` - RESHARP algorithm parameters
54/// * `progress` - Progress callback (current_iteration, max_iterations)
55///
56/// # Returns
57/// (local_field, eroded_mask)
58pub fn resharp(
59    field: &[f64],
60    mask: &[u8],
61    grid: &Grid,
62    params: &ResharpParams,
63    mut progress: impl FnMut(usize, usize),
64) -> (Vec<f64>, Vec<u8>) {
65    let (nx, ny, nz) = grid.dims;
66    let n_total = nx * ny * nz;
67
68    // Generate SMV kernel and FFT
69    let s_kernel = smv_kernel(grid, params.radius);
70    let s_fft = fft_real_kernel(&s_kernel, nx, ny, nz);
71
72    // DKER = 1 - S (high-pass / delta-kernel) in k-space
73    let dker: Vec<f64> = s_fft.iter().map(|&s| 1.0 - s).collect();
74
75    // Erode mask via SMV convolution
76    let eroded_mask = erode_mask_smv(mask, &s_fft, grid, 1.0 - 1e-7_f64.sqrt());
77    let eroded_mask_f64: Vec<f64> = eroded_mask.iter()
78        .map(|&m| m as f64)
79        .collect();
80
81    // Compute RHS: b = H'(H(field))
82    // H(x) = M * ifft(DKER * fft(x))
83    // H'(y) = ifft(DKER * fft(M * y))   [DKER is real so conj(DKER) = DKER]
84    //
85    // b = ifft(DKER * fft(M * ifft(DKER * fft(field))))
86    let b = apply_ht_h(&dker, &eroded_mask_f64, field, nx, ny, nz);
87
88    // Solve (H'H + λI)x = b via CG
89    let tik_reg = params.tik_reg;
90    let x0 = vec![0.0; n_total];
91    let x = cg_solve_with_progress(
92        |x_vec| {
93            let mut result = apply_ht_h(&dker, &eroded_mask_f64, x_vec, nx, ny, nz);
94            // Add Tikhonov term: + λx
95            for i in 0..n_total {
96                result[i] += tik_reg * x_vec[i];
97            }
98            result
99        },
100        &b,
101        &x0,
102        params.tol,
103        params.max_iter,
104        &mut progress,
105    );
106
107    // Apply eroded mask to result
108    let local_field: Vec<f64> = x.iter()
109        .enumerate()
110        .map(|(i, &v)| if eroded_mask[i] == 1 { v } else { 0.0 })
111        .collect();
112
113    (local_field, eroded_mask)
114}
115
116/// Apply H'H to a vector: ifft(DKER * fft(M * ifft(DKER * fft(x))))
117///
118/// H(x)  = M * ifft(DKER * fft(x))
119/// H'(y) = ifft(DKER * fft(M * y))
120fn apply_ht_h(
121    dker: &[f64],
122    mask: &[f64],
123    x: &[f64],
124    nx: usize, ny: usize, nz: usize,
125) -> Vec<f64> {
126    let n_total = nx * ny * nz;
127
128    // Step 1: H(x) = M * ifft(DKER * fft(x))
129    let mut tmp: Vec<Complex64> = x.iter()
130        .map(|&v| Complex64::new(v, 0.0))
131        .collect();
132    fft3d(&mut tmp, nx, ny, nz);
133
134    for i in 0..n_total {
135        tmp[i] *= dker[i];
136    }
137    ifft3d(&mut tmp, nx, ny, nz);
138
139    // Apply mask
140    for i in 0..n_total {
141        tmp[i] *= mask[i];
142    }
143
144    // Step 2: H'(Hx) = ifft(DKER * fft(M * Hx))
145    fft3d(&mut tmp, nx, ny, nz);
146
147    for i in 0..n_total {
148        tmp[i] *= dker[i];
149    }
150    ifft3d(&mut tmp, nx, ny, nz);
151
152    tmp.iter().map(|c| c.re).collect()
153}
154
155#[cfg(test)]
156mod tests {
157    use super::*;
158
159    #[test]
160    fn test_resharp_zero_field() {
161        let n = 16;
162        let field = vec![0.0; n * n * n];
163        let mask = vec![1u8; n * n * n];
164        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
165        let params = ResharpParams { radius: 2.0, tik_reg: 1e-4, tol: 1e-6, max_iter: 50 };
166
167        let (local, _) = resharp(&field, &mask, &grid, &params, |_, _| {});
168
169        for &val in local.iter() {
170            assert!(val.abs() < 1e-8, "Zero field should give zero local field, got {}", val);
171        }
172    }
173
174    #[test]
175    fn test_resharp_finite() {
176        let n = 16;
177        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.01).collect();
178        let mask = vec![1u8; n * n * n];
179        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
180        let params = ResharpParams { radius: 2.0, tik_reg: 1e-4, tol: 1e-6, max_iter: 50 };
181
182        let (local, eroded) = resharp(&field, &mask, &grid, &params, |_, _| {});
183
184        for (i, &val) in local.iter().enumerate() {
185            assert!(val.is_finite(), "Local field should be finite at index {}", i);
186        }
187
188        let eroded_count: usize = eroded.iter().map(|&m| m as usize).sum();
189        assert!(eroded_count > 0, "Eroded mask should have some voxels");
190    }
191}