Skip to main content

qsm_core/inversion/
hdqsm.rs

1//! HD-QSM: Hybrid data-fidelity two-stage linear dipole inversion.
2//!
3//! A linear dipole-inversion method that runs in two stages. Stage 1 solves an
4//! L1 data-fidelity problem (robust to phase/model errors) and derives a
5//! spatially-varying discrepancy weighting. Stage 2 solves an L2 data-fidelity
6//! problem reweighted by that discrepancy map. Both stages use ADMM with an
7//! L1 total-variation regularizer.
8//!
9//! Because the whole method is linear and scale-consistent, feed the local field
10//! in ppm directly and the output susceptibility is in ppm as well (no phase
11//! scaling required).
12//!
13//! Reference:
14//! Lambert, M., Tejos, C., Langkammer, C., et al. (2022).
15//! "Hybrid data fidelity term approach for quantitative susceptibility mapping."
16//! Magnetic Resonance in Medicine, 87(6):3059-3072.
17//! https://doi.org/10.1002/mrm.29218
18//!
19//! Reference implementation: HDQSM.m, https://github.com/mglambert/HD-QSM (MIT).
20
21use num_complex::Complex64;
22use crate::inversion::admm::prepare_fansi_spectral;
23use crate::utils::gradient::{bdiv_inplace, fgrad_inplace};
24use crate::utils::{apply_mask_zero, shrink};
25use crate::Grid;
26
27/// HD-QSM algorithm parameters.
28#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
29#[derive(Clone, Debug)]
30pub struct HdQsmParams {
31    /// L2-stage TV weight (default 1e-4).
32    pub alpha_l2: f64,
33    /// L2-stage gradient-consistency ADMM weight (default 1e-2 = 100*alpha_l2).
34    ///
35    /// The TV soft-threshold is `alpha_l2 / mu1_l2`; the FANSI HDQSM.m example's
36    /// `mu1_l2 = 10*alpha_l2` locks that ratio at 0.1, which over-smooths on the
37    /// ppm-scaled QSM-CI phantom (the README warns the example is calibrated to a
38    /// different, normalized brain). We use `mu1_l2 = 100*alpha_l2` (ratio 0.01),
39    /// matching QSM.rs's proven TV-ADMM balance — corr jumps from ~0.48 to ~0.84.
40    pub mu1_l2: f64,
41    /// Fidelity consistency weight (default 1.0).
42    pub mu2: f64,
43    /// Stage-1 (L1) iterations (default 20).
44    pub max_iter_l1: usize,
45    /// Stage-2 (L2) iterations (default 80).
46    pub max_iter_l2: usize,
47    /// Stage-2 percent-update stopping tolerance (default 1.0).
48    pub tol_update: f64,
49    // Stage-1 alpha_l1 / mu1_l1 default to sqrt(alpha_l2) / sqrt(mu1_l2) internally.
50}
51
52impl Default for HdQsmParams {
53    fn default() -> Self {
54        Self {
55            alpha_l2: 1e-4,
56            mu1_l2: 1e-2,
57            mu2: 1.0,
58            max_iter_l1: 20,
59            max_iter_l2: 80,
60            tol_update: 1.0,
61        }
62    }
63}
64
65/// L2 norm of a slice.
66#[inline]
67fn norm2(v: &[f64]) -> f64 {
68    v.iter().map(|&x| x * x).sum::<f64>().sqrt()
69}
70
71/// Percent update `100 * ||x - x_prev|| / ||x||`.
72#[inline]
73fn percent_update(x: &[f64], x_prev: &[f64]) -> f64 {
74    let diff: f64 = x.iter().zip(x_prev).map(|(&a, &b)| (a - b) * (a - b)).sum::<f64>().sqrt();
75    let nx = norm2(x);
76    if nx > 0.0 { 100.0 * diff / nx } else { 0.0 }
77}
78
79/// Forward dipole convolution: `Dx = real(ifft(kernel .* fft(x)))`.
80///
81/// `cbuf` is scratch of length n_total; result written to `out`.
82fn apply_dipole(
83    fft_ws: &mut crate::fft::Fft3dWorkspace,
84    k: &[f64],
85    x: &[f64],
86    cbuf: &mut [Complex64],
87    out: &mut [f64],
88) {
89    for (c, &xv) in cbuf.iter_mut().zip(x.iter()) {
90        *c = Complex64::new(xv, 0.0);
91    }
92    fft_ws.fft3d(cbuf);
93    for (c, &kv) in cbuf.iter_mut().zip(k.iter()) {
94        *c *= kv;
95    }
96    fft_ws.ifft3d(cbuf);
97    for (o, c) in out.iter_mut().zip(cbuf.iter()) {
98        *o = c.re;
99    }
100}
101
102/// HD-QSM dipole inversion.
103///
104/// # Arguments
105/// * `local_field` - Local field values in ppm (nx * ny * nz).
106/// * `mask` - Binary ROI mask (nx * ny * nz), 1 = inside.
107/// * `grid` - Volume grid (dimensions and voxel sizes).
108/// * `bdir` - B0 field direction.
109/// * `params` - HD-QSM parameters.
110/// * `progress` - Progress callback `(global_iter, max_iter_l1 + max_iter_l2)`.
111///
112/// # Returns
113/// Susceptibility map in ppm.
114pub fn hdqsm(
115    local_field: &[f64],
116    mask: &[u8],
117    grid: &Grid,
118    bdir: (f64, f64, f64),
119    params: &HdQsmParams,
120    mut progress: impl FnMut(usize, usize),
121) -> Vec<f64> {
122    let n = grid.n_total();
123    let (mut fft_ws, k, ee2) = prepare_fansi_spectral(grid, bdir);
124    // K2 = abs(kernel).^2 = k^2 (kernel is real). mu*EE2 is added per stage.
125    let denom_base: Vec<f64> = k.iter()
126        .map(|&kk| 1e-30 + params.mu2 * kk * kk)
127        .collect();
128
129    // weight = mask (as f64 in [0,1]).
130    let weight_mask: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
131
132    let total_iter = params.max_iter_l1 + params.max_iter_l2;
133    let mut global_iter = 0usize;
134
135    // Working buffers.
136    let mut x = vec![0.0f64; n];
137    let mut x_prev = vec![0.0f64; n];
138
139    let mut z_dx = vec![0.0f64; n];
140    let mut z_dy = vec![0.0f64; n];
141    let mut z_dz = vec![0.0f64; n];
142    let mut s_dx = vec![0.0f64; n];
143    let mut s_dy = vec![0.0f64; n];
144    let mut s_dz = vec![0.0f64; n];
145    let mut x_dx = vec![0.0f64; n];
146    let mut x_dy = vec![0.0f64; n];
147    let mut x_dz = vec![0.0f64; n];
148
149    let mut z2 = vec![0.0f64; n];
150    let mut s2 = vec![0.0f64; n];
151    let mut dx = vec![0.0f64; n]; // Dx = real(ifft(K.*Fx))
152
153    let mut div = vec![0.0f64; n];
154    let mut gx = vec![0.0f64; n];
155    let mut gy = vec![0.0f64; n];
156    let mut gz = vec![0.0f64; n];
157
158    let mut cbuf = vec![Complex64::new(0.0, 0.0); n];
159    let mut fdiv = vec![Complex64::new(0.0, 0.0); n];
160    let mut fdt = vec![Complex64::new(0.0, 0.0); n];
161
162    // ---- Helper closures cannot capture &mut fft_ws + buffers simultaneously,
163    //      so the x-subproblem is inlined below. ----
164
165    // =========================================================================
166    // STAGE 1: L1 data fidelity.
167    // =========================================================================
168    {
169        let mu = params.mu1_l2.sqrt();
170        let alpha = params.alpha_l2.sqrt();
171        let mu2 = params.mu2;
172        let ll = alpha / mu;
173
174        // Denominator for stage 1: eps + mu2*K2 + mu*EE2.
175        let denom: Vec<f64> = (0..n).map(|i| denom_base[i] + mu * ee2[i]).collect();
176
177        // Wy = input.
178        let wy = local_field;
179
180        for t in 0..params.max_iter_l1 {
181            x_prev.copy_from_slice(&x);
182
183            // Dt_kspace_arg = z2 - s2 + Wy.
184            for i in 0..n {
185                gx[i] = z2[i] - s2[i] + wy[i]; // reuse gx as dt_arg scratch
186            }
187            // fft(dt_arg) -> fdt
188            for (c, &v) in fdt.iter_mut().zip(gx.iter()) {
189                *c = Complex64::new(v, 0.0);
190            }
191            fft_ws.fft3d(&mut fdt);
192
193            // fft(bdiv(z_d - s_d)) -> fdiv
194            for i in 0..n {
195                gx[i] = z_dx[i] - s_dx[i];
196                gy[i] = z_dy[i] - s_dy[i];
197                gz[i] = z_dz[i] - s_dz[i];
198            }
199            bdiv_inplace(&mut div, &gx, &gy, &gz, grid);
200            for (c, &v) in fdiv.iter_mut().zip(div.iter()) {
201                *c = Complex64::new(v, 0.0);
202            }
203            fft_ws.fft3d(&mut fdiv);
204
205            // num = mu*fdiv + mu2*conj(K)*fdt  (K real -> conj(K)=K); divide by denom.
206            // Guard the dipole null-space (DC/singular bins): denom -> ~0 there
207            // (dipole kernel and Laplacian both vanish). Zero it instead of
208            // dividing FFT round-off by ~0, which would create a huge DC pedestal.
209            for i in 0..n {
210                // Minus on the gradient term: adjoint of crate `fgrad` is `-bdiv` (matches
211                // QSM.rs TV-ADMM). `+bdiv` doubles the effective regularization. See fansi.rs.
212                let num = -fdiv[i] * mu + fdt[i] * (mu2 * k[i]);
213                cbuf[i] = if denom[i] > 1e-20 { num / denom[i] } else { Complex64::new(0.0, 0.0) };
214            }
215            fft_ws.ifft3d(&mut cbuf);
216            for i in 0..n {
217                x[i] = cbuf[i].re;
218            }
219
220            global_iter += 1;
221            progress(global_iter, total_iter);
222
223            if t < params.max_iter_l1 - 1 {
224                // fgrad(x) -> x_d
225                fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
226                // z_d = shrink(x_d + s_d, ll); s_d += x_d - z_d.
227                for i in 0..n {
228                    let ax = x_dx[i] + s_dx[i];
229                    let ay = x_dy[i] + s_dy[i];
230                    let az = x_dz[i] + s_dz[i];
231                    z_dx[i] = shrink(ax, ll);
232                    z_dy[i] = shrink(ay, ll);
233                    z_dz[i] = shrink(az, ll);
234                    s_dx[i] += x_dx[i] - z_dx[i];
235                    s_dy[i] += x_dy[i] - z_dy[i];
236                    s_dz[i] += x_dz[i] - z_dz[i];
237                }
238
239                // Dx = real(ifft(K.*Fx))
240                apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
241                // z2_inner = Dx + s2 - Wy; z2 = shrink(z2_inner, weight/mu2); s2 = z2_inner - z2.
242                for i in 0..n {
243                    let z2_inner = dx[i] + s2[i] - wy[i];
244                    z2[i] = shrink(z2_inner, weight_mask[i] / mu2);
245                    s2[i] = z2_inner - z2[i];
246                }
247            }
248        }
249    }
250
251    // Discrepancy factor: dphi = ifft(fft(input) - fft(x).*K); dphi = abs(dphi).*mask; dphi /= max.
252    // = |input - Dx| .* mask, normalized.  (fft(input) - K.*fft(x) -> ifft = input - Dx)
253    let mut dphi = vec![0.0f64; n];
254    {
255        // Dx = real(ifft(K.*Fx)) with current x.
256        apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
257        let mut maxv = 0.0f64;
258        for i in 0..n {
259            let v = (local_field[i] - dx[i]).abs() * (if mask[i] != 0 { 1.0 } else { 0.0 });
260            dphi[i] = v;
261            if v > maxv {
262                maxv = v;
263            }
264        }
265        if maxv > 0.0 {
266            for v in dphi.iter_mut() {
267                *v /= maxv;
268            }
269        }
270    }
271
272    // =========================================================================
273    // STAGE 2: L2 data fidelity, reweighted by discrepancy.
274    // =========================================================================
275    {
276        let mu = params.mu1_l2;
277        let alpha = params.alpha_l2;
278        let mu2 = params.mu2;
279        let ll = alpha / mu;
280
281        // Denominator for stage 2: eps + mu2*K2 + mu*EE2.
282        let denom: Vec<f64> = (0..n).map(|i| denom_base[i] + mu * ee2[i]).collect();
283
284        // weight = weight.*weight.*(1-dphi).*(1-dphi).
285        let weight: Vec<f64> = (0..n).map(|i| {
286            let w = weight_mask[i];
287            let d = 1.0 - dphi[i];
288            w * w * d * d
289        }).collect();
290        // Wy = weight.*input./(weight+mu2).
291        let wy: Vec<f64> = (0..n).map(|i| weight[i] * local_field[i] / (weight[i] + mu2)).collect();
292
293        // Reset dual/aux for stage 2 (x carried over).
294        for i in 0..n {
295            z_dx[i] = 0.0; z_dy[i] = 0.0; z_dz[i] = 0.0;
296            s_dx[i] = 0.0; s_dy[i] = 0.0; s_dz[i] = 0.0;
297            s2[i] = 0.0;
298        }
299
300        for _t in 0..params.max_iter_l2 {
301            x_prev.copy_from_slice(&x);
302
303            // z/s updates BEFORE the x-update (unlike stage 1).
304            fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
305            for i in 0..n {
306                let ax = x_dx[i] + s_dx[i];
307                let ay = x_dy[i] + s_dy[i];
308                let az = x_dz[i] + s_dz[i];
309                z_dx[i] = shrink(ax, ll);
310                z_dy[i] = shrink(ay, ll);
311                z_dz[i] = shrink(az, ll);
312                s_dx[i] += x_dx[i] - z_dx[i];
313                s_dy[i] += x_dy[i] - z_dy[i];
314                s_dz[i] += x_dz[i] - z_dz[i];
315            }
316
317            // Dx = real(ifft(K.*Fx))
318            apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
319            // z2 = Wy + mu2*(Dx + s2)./(weight + mu2); s2 = s2 + Dx - z2.
320            for i in 0..n {
321                z2[i] = wy[i] + mu2 * (dx[i] + s2[i]) / (weight[i] + mu2);
322                s2[i] = s2[i] + dx[i] - z2[i];
323            }
324
325            // x update. Dt_kspace_arg = z2 - s2.
326            for i in 0..n {
327                gx[i] = z2[i] - s2[i];
328            }
329            for (c, &v) in fdt.iter_mut().zip(gx.iter()) {
330                *c = Complex64::new(v, 0.0);
331            }
332            fft_ws.fft3d(&mut fdt);
333
334            for i in 0..n {
335                gx[i] = z_dx[i] - s_dx[i];
336                gy[i] = z_dy[i] - s_dy[i];
337                gz[i] = z_dz[i] - s_dz[i];
338            }
339            bdiv_inplace(&mut div, &gx, &gy, &gz, grid);
340            for (c, &v) in fdiv.iter_mut().zip(div.iter()) {
341                *c = Complex64::new(v, 0.0);
342            }
343            fft_ws.fft3d(&mut fdiv);
344
345            for i in 0..n {
346                // Minus on the gradient term: adjoint of crate `fgrad` is `-bdiv` (matches
347                // QSM.rs TV-ADMM). `+bdiv` doubles the effective regularization. See fansi.rs.
348                let num = -fdiv[i] * mu + fdt[i] * (mu2 * k[i]);
349                // Guard the dipole null-space (DC/singular bins) — see stage 1.
350                cbuf[i] = if denom[i] > 1e-20 { num / denom[i] } else { Complex64::new(0.0, 0.0) };
351            }
352            fft_ws.ifft3d(&mut cbuf);
353            for i in 0..n {
354                x[i] = cbuf[i].re;
355            }
356
357            global_iter += 1;
358            progress(global_iter, total_iter);
359
360            let upd = percent_update(&x, &x_prev);
361            if upd < params.tol_update {
362                break;
363            }
364        }
365    }
366
367    // Ensure progress reaches the total even if stage 2 stopped early.
368    if global_iter < total_iter {
369        progress(total_iter, total_iter);
370    }
371
372    apply_mask_zero(&mut x, mask);
373    x
374}
375
376#[cfg(test)]
377mod tests {
378    use super::*;
379
380    #[test]
381    fn test_hdqsm_zero_field() {
382        let n = 8;
383        let field = vec![0.0; n * n * n];
384        let mask = vec![1u8; n * n * n];
385        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
386        let params = HdQsmParams {
387            max_iter_l1: 5,
388            max_iter_l2: 10,
389            ..HdQsmParams::default()
390        };
391
392        let chi = hdqsm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
393
394        for &val in chi.iter() {
395            assert!(val.abs() < 1e-6, "Zero field should give zero chi, got {}", val);
396        }
397    }
398
399    #[test]
400    fn test_hdqsm_finite() {
401        let n = 8;
402        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
403        let mask = vec![1u8; n * n * n];
404        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
405        let params = HdQsmParams {
406            max_iter_l1: 5,
407            max_iter_l2: 10,
408            ..HdQsmParams::default()
409        };
410
411        let chi = hdqsm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
412
413        for (i, &val) in chi.iter().enumerate() {
414            assert!(val.is_finite(), "Chi should be finite at index {}", i);
415        }
416    }
417}