Skip to main content

qsm_core/inversion/
tv.rs

1//! Total Variation (TV) regularized dipole inversion using ADMM
2//!
3//! Solves the L1-regularized inverse problem:
4//! min_x ||Dx - f||_2^2 + lambda||grad(x)||_1
5//!
6//! using Alternating Direction Method of Multipliers (ADMM).
7//!
8//! Reference:
9//! Bilgic, B., Fan, A.P., Polimeni, J.R., et al. (2014).
10//! "Fast quantitative susceptibility mapping with L1-regularization and automatic
11//! parameter selection." Magnetic Resonance in Medicine, 72(5):1444-1459.
12//! https://doi.org/10.1002/mrm.25029
13//!
14//! Reference implementation: https://github.com/kamesy/QSM.jl
15
16use crate::utils::{shrink, apply_mask_zero};
17use crate::Grid;
18use super::admm::{AdmmBuffers, admm_step, prepare_admm_spectral};
19
20/// TV-ADMM algorithm parameters
21#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
22#[derive(Clone, Debug)]
23pub struct TvParams {
24    /// Regularization parameter (typically 1e-3 to 1e-4)
25    pub lambda: f64,
26    /// ADMM penalty parameter (typically 100*lambda)
27    pub rho: f64,
28    /// Convergence tolerance
29    pub tol: f64,
30    /// Maximum iterations
31    pub max_iter: usize,
32}
33
34impl Default for TvParams {
35    fn default() -> Self {
36        Self {
37            lambda: 2e-4,
38            rho: 2e-2,
39            tol: 1e-3,
40            max_iter: 250,
41        }
42    }
43}
44
45/// TV-ADMM dipole inversion
46///
47/// Optimized implementation with:
48/// - Pre-allocated buffers (zero allocations per iteration)
49/// - In-place gradient/divergence operations
50/// - Buffer swapping instead of cloning
51/// - Fused z-subproblem and u-update
52///
53/// # Arguments
54/// * `local_field` - Local field values (nx * ny * nz)
55/// * `mask` - Binary mask (nx * ny * nz), 1 = inside ROI
56/// * `grid` - Volume grid (dimensions and voxel sizes)
57/// * `bdir` - B0 field direction
58/// * `params` - TV-ADMM parameters
59/// * `progress` - Progress callback `(iteration, max_iter)`
60///
61/// # Returns
62/// Susceptibility map
63pub fn tv_admm(
64    local_field: &[f64],
65    mask: &[u8],
66    grid: &Grid,
67    bdir: (f64, f64, f64),
68    params: &TvParams,
69    mut progress: impl FnMut(usize, usize),
70) -> Vec<f64> {
71    let n_total = grid.n_total();
72
73    // Pre-compute spectral operators
74    let (mut fft_ws, inv_a, f_hat) = prepare_admm_spectral(local_field, grid, bdir, params.rho);
75
76    // Pre-allocate working buffers and run ADMM iterations
77    let mut buf = AdmmBuffers::new(n_total);
78    let lambda_over_rho = params.lambda / params.rho;
79
80    for iter in 0..params.max_iter {
81        progress(iter + 1, params.max_iter);
82
83        let converged = admm_step(
84            &mut buf, &mut fft_ws, &f_hat, &inv_a, params.rho, grid, params.tol,
85            |vx, vy, vz, _| (shrink(vx, lambda_over_rho), shrink(vy, lambda_over_rho), shrink(vz, lambda_over_rho)),
86        );
87
88        if converged {
89            progress(iter + 1, iter + 1);
90            break;
91        }
92    }
93
94    apply_mask_zero(&mut buf.x, mask);
95
96    buf.x
97}
98
99#[cfg(test)]
100mod tests {
101    use super::*;
102    use crate::utils::gradient::fgrad;
103
104    #[test]
105    fn test_shrink() {
106        assert!((shrink(1.0, 0.5) - 0.5).abs() < 1e-10);
107        assert!((shrink(-1.0, 0.5) - (-0.5)).abs() < 1e-10);
108        assert!((shrink(0.3, 0.5) - 0.0).abs() < 1e-10);
109        assert!((shrink(-0.3, 0.5) - 0.0).abs() < 1e-10);
110    }
111
112    #[test]
113    fn test_tv_admm_zero_field() {
114        // Zero field should give zero susceptibility
115        let n = 8;
116        let field = vec![0.0; n * n * n];
117        let mask = vec![1u8; n * n * n];
118        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
119        let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-2, max_iter: 10 };
120
121        let chi = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
122
123        for &val in chi.iter() {
124            assert!(val.abs() < 1e-8, "Zero field should give zero chi, got {}", val);
125        }
126    }
127
128    #[test]
129    fn test_tv_admm_finite() {
130        // Result should be finite
131        let n = 8;
132        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
133        let mask = vec![1u8; n * n * n];
134        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
135        let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-2, max_iter: 10 };
136
137        let chi = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
138
139        for (i, &val) in chi.iter().enumerate() {
140            assert!(val.is_finite(), "Chi should be finite at index {}", i);
141        }
142    }
143
144    #[test]
145    fn test_tv_admm_smoother_than_tkd() {
146        // TV should produce smoother results than TKD
147        let n = 8;
148        // Create noisy field
149        let mut field = vec![0.0; n * n * n];
150        for i in 0..n*n*n {
151            field[i] = if i % 2 == 0 { 0.01 } else { -0.01 };  // Alternating
152        }
153        let mask = vec![1u8; n * n * n];
154        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
155        let params = TvParams { lambda: 1e-2, rho: 1.0, tol: 1e-2, max_iter: 50 };
156
157        let chi_tv = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
158
159        // Compute total variation (L1 norm of gradient)
160        let (gx, gy, gz) = fgrad(&chi_tv, &grid);
161        let tv: f64 = gx.iter().chain(gy.iter()).chain(gz.iter())
162            .map(|&g| g.abs())
163            .sum();
164
165        // TV result should have small total variation
166        // (exact value depends on parameters, but should be bounded)
167        assert!(tv.is_finite(), "TV should be finite");
168    }
169
170    /// Verify parallel and sequential TV-ADMM produce identical results.
171    #[cfg(feature = "parallel")]
172    #[test]
173    fn test_tv_parallel_matches_sequential() {
174        let n = 16;
175        let field: Vec<f64> = (0..n*n*n).map(|i| ((i as f64) * 0.7).sin() * 0.01).collect();
176        let mask = vec![1u8; n * n * n];
177        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
178        let params = TvParams { lambda: 1e-3, rho: 0.1, tol: 1e-3, max_iter: 50 };
179
180        // Sequential (1 thread)
181        let pool_1 = rayon::ThreadPoolBuilder::new().num_threads(1).build().unwrap();
182        let chi_seq = pool_1.install(|| {
183            tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {})
184        });
185
186        // Parallel (default threads)
187        let chi_par = tv_admm(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
188
189        // Compare
190        for (i, (s, p)) in chi_seq.iter().zip(chi_par.iter()).enumerate() {
191            assert!(
192                (s - p).abs() < 1e-10,
193                "TV mismatch at voxel {}: seq={} par={} diff={}",
194                i, s, p, (s - p).abs()
195            );
196        }
197    }
198}