Skip to main content

qsm_core/inversion/
fansi.rs

1//! FANSI nonlinear TV / TGV dipole inversion.
2//!
3//! Nonlinear total-variation (nlTV) and nonlinear total-generalized-variation
4//! (nlTGV) QSM dipole inversion with a nonlinear (wrapped-phase) data-fidelity
5//! term, solved with ADMM plus an inner complex-argument Newton iteration.
6//!
7//! The nonlinear fidelity models the field data as `exp(i * D x)` and matches
8//! it to the wrapped local-field phase, which makes the reconstruction robust to
9//! phase-wrap / high-field regimes.
10//!
11//! References:
12//! - nlTV: Milovic, C., Bilgic, B., Zhao, B., et al. (2018).
13//!   "Fast nonlinear susceptibility inversion with variational regularization."
14//!   Magnetic Resonance in Medicine, 80(2):814-821.
15//!   <https://doi.org/10.1002/mrm.27073>
16//! - nlTGV: the total-generalized-variation variant of the above.
17//!
18//! Ported faithfully from the FANSI toolbox `nlTV.m` and `nlTGV.m`
19//! (<https://gitlab.com/cmilovic/FANSI-toolbox>).
20//!
21//! # Units / `phase_scale`
22//! The FANSI fidelity is a *phase* (radians) model: `sin(z2 - phase)`. If the
23//! input `local_field` is already radians, use `phase_scale = 1.0`. If it is
24//! ppm-scale, `phase_scale` should convert ppm -> radians (`2*pi*gamma*B0*TE`);
25//! the returned map is divided by `phase_scale` so it is expressed on the same
26//! scale as the input field.
27
28use crate::inversion::admm::prepare_fansi_spectral;
29use crate::utils::gradient::{bdiv_inplace, fgrad_inplace};
30use crate::utils::{apply_mask_zero, shrink};
31use crate::Grid;
32use num_complex::Complex64;
33use std::f64::consts::PI;
34
35/// FANSI nlTV/nlTGV parameters.
36#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
37#[derive(Clone, Debug)]
38pub struct FansiParams {
39    /// First-order (TV / TGV gradient) L1 penalty weight.
40    pub alpha1: f64,
41    /// Gradient-consistency ADMM weight.
42    pub mu1: f64,
43    /// Fidelity-consistency ADMM weight.
44    pub mu2: f64,
45    /// Second-order (symmetric-gradient) L1 penalty weight (nlTGV only).
46    pub alpha0: f64,
47    /// Second-order consistency ADMM weight (nlTGV only).
48    pub mu0: f64,
49    /// Number of outer ADMM iterations.
50    pub max_iter: usize,
51    /// Percent-update convergence stopping tolerance.
52    pub tol_update: f64,
53    /// Inner Newton convergence tolerance.
54    pub tol_delta: f64,
55    /// Working (phase) scale applied to the input local field; output is divided
56    /// by it. Use 1.0 for radians input, ppm->radians factor for ppm input.
57    pub phase_scale: f64,
58    /// Select nlTGV (`true`) or nlTV (`false`).
59    pub is_tgv: bool,
60}
61
62impl Default for FansiParams {
63    fn default() -> Self {
64        Self {
65            alpha1: 2e-4,
66            mu1: 2e-2,
67            mu2: 1.0,
68            alpha0: 4e-4,
69            mu0: 4e-2,
70            max_iter: 150,
71            tol_update: 0.1,
72            tol_delta: 1e-6,
73            phase_scale: 1.0,
74            is_tgv: false,
75        }
76    }
77}
78
79/// L2 norm of a real slice.
80fn norm2(v: &[f64]) -> f64 {
81    v.iter().map(|&a| a * a).sum::<f64>().sqrt()
82}
83
84/// FANSI nlTV/nlTGV dipole inversion.
85///
86/// # Arguments
87/// * `local_field` - Local field values (nx * ny * nz).
88/// * `mask` - Binary mask (nx * ny * nz), non-zero = inside ROI.
89/// * `grid` - Volume grid (dimensions and voxel sizes).
90/// * `bdir` - B0 field direction.
91/// * `params` - FANSI parameters (`is_tgv` selects nlTGV vs nlTV).
92/// * `progress` - Progress callback `(iteration, max_iter)`.
93///
94/// # Returns
95/// Estimated susceptibility map, masked to the ROI.
96pub fn fansi(
97    local_field: &[f64],
98    mask: &[u8],
99    grid: &Grid,
100    bdir: (f64, f64, f64),
101    params: &FansiParams,
102    progress: impl FnMut(usize, usize),
103) -> Vec<f64> {
104    if params.is_tgv {
105        nltgv(local_field, mask, grid, bdir, params, progress)
106    } else {
107        nltv(local_field, mask, grid, bdir, params, progress)
108    }
109}
110
111/// Nonlinear total-variation dipole inversion (FANSI `nlTV.m`).
112fn nltv(
113    local_field: &[f64],
114    mask: &[u8],
115    grid: &Grid,
116    bdir: (f64, f64, f64),
117    params: &FansiParams,
118    mut progress: impl FnMut(usize, usize),
119) -> Vec<f64> {
120    let n = grid.n_total();
121
122    let (mut fft_ws, k, ee2) = prepare_fansi_spectral(grid, bdir);
123
124    // Scaled phase and W = mask (0/1 weight).
125    let phase: Vec<f64> = local_field.iter().map(|&f| f * params.phase_scale).collect();
126    let w: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
127
128    let mu1 = params.mu1;
129    let mu2 = params.mu2;
130    let alpha_over_mu = params.alpha1 / mu1;
131
132    // ADMM variables.
133    let mut x = vec![0.0f64; n];
134    let mut x_prev = vec![0.0f64; n];
135
136    // Gradient-consistency split (real) and its multipliers.
137    let mut z_dx = vec![0.0f64; n];
138    let mut z_dy = vec![0.0f64; n];
139    let mut z_dz = vec![0.0f64; n];
140    let mut s_dx = vec![0.0f64; n];
141    let mut s_dy = vec![0.0f64; n];
142    let mut s_dz = vec![0.0f64; n];
143
144    // Fidelity-consistency auxiliary z2 = W .* phase ./ (W + mu2), and multiplier.
145    let mut z2 = vec![0.0f64; n];
146    for i in 0..n {
147        let den = w[i] + mu2;
148        z2[i] = if den != 0.0 { w[i] * phase[i] / den } else { 0.0 };
149    }
150    let mut s2 = vec![0.0f64; n];
151
152    // Reusable buffers.
153    let mut fdiv = vec![Complex64::new(0.0, 0.0); n];
154    let mut fd2 = vec![Complex64::new(0.0, 0.0); n];
155    let mut xhat = vec![Complex64::new(0.0, 0.0); n];
156    let mut fx = vec![Complex64::new(0.0, 0.0); n];
157
158    let mut gxc = vec![0.0f64; n];
159    let mut gyc = vec![0.0f64; n];
160    let mut gzc = vec![0.0f64; n];
161    let mut x_dx = vec![0.0f64; n];
162    let mut x_dy = vec![0.0f64; n];
163    let mut x_dz = vec![0.0f64; n];
164    let mut div = vec![0.0f64; n];
165    let mut dx = vec![0.0f64; n]; // Dx = real(ifft(k .* fft(x)))
166    let mut rhs_z2 = vec![0.0f64; n];
167    let mut diff = vec![0.0f64; n];
168    let mut update = vec![0.0f64; n];
169
170    for t in 0..params.max_iter {
171        progress(t + 1, params.max_iter);
172
173        // ---- x-subproblem -------------------------------------------------
174        // Gradient side: mu1 * sum(E_t .* fft(z_d - s_d)) == mu1 * fft(bdiv(z_d - s_d)).
175        for i in 0..n {
176            gxc[i] = z_dx[i] - s_dx[i];
177            gyc[i] = z_dy[i] - s_dy[i];
178            gzc[i] = z_dz[i] - s_dz[i];
179        }
180        bdiv_inplace(&mut div, &gxc, &gyc, &gzc, grid);
181        for i in 0..n {
182            fdiv[i] = Complex64::new(div[i], 0.0);
183        }
184        fft_ws.fft3d(&mut fdiv);
185
186        // Fidelity side: mu2 * conj(K) .* fft(z2 - s2) (K real -> conj = K).
187        for i in 0..n {
188            fd2[i] = Complex64::new(z2[i] - s2[i], 0.0);
189        }
190        fft_ws.fft3d(&mut fd2);
191
192        for i in 0..n {
193            // NOTE the minus on the gradient-consistency term: the adjoint of the
194            // crate's forward-difference `fgrad` is `-bdiv` (not `+bdiv`), so the
195            // spectral term mu1*sum(E_t.*F(z_d-s_d)) = -mu1*F(bdiv(z_d-s_d)).
196            // Matches QSM.rs's own TV-ADMM (`f_hat - rho*fft(bdiv...)`). Using +bdiv
197            // fails to cancel mu1*∇*∇ at the fixed point, doubling the effective
198            // regularization and damping the susceptibility amplitude.
199            let num = -mu1 * fdiv[i] + mu2 * k[i] * fd2[i];
200            let den = mu2 * k[i] * k[i] + mu1 * ee2[i];
201            // Guard the dipole null-space (DC and any singular bin): both the
202            // dipole kernel and the Laplacian vanish there, so susceptibility is
203            // undetermined up to a constant. Zero it (matches TV's inv_a guard);
204            // dividing FFT round-off by ~0 would otherwise create a huge DC pedestal.
205            xhat[i] = if den > 1e-20 { num / den } else { Complex64::new(0.0, 0.0) };
206        }
207        fft_ws.ifft3d(&mut xhat);
208        x_prev.copy_from_slice(&x);
209        for i in 0..n {
210            x[i] = xhat[i].re;
211        }
212
213        // ---- convergence check --------------------------------------------
214        let xnorm = norm2(&x);
215        if xnorm > 0.0 {
216            for i in 0..n {
217                diff[i] = x[i] - x_prev[i];
218            }
219            let x_update = 100.0 * norm2(&diff) / xnorm;
220            if x_update < params.tol_update || x_update.is_nan() {
221                progress(t + 1, t + 1);
222                break;
223            }
224        }
225
226        if t + 1 >= params.max_iter {
227            break;
228        }
229
230        // ---- gradient split update (TV shrink) ----------------------------
231        // Fx = fft(x) (reused for the dipole rhs below).
232        for i in 0..n {
233            fx[i] = Complex64::new(x[i], 0.0);
234        }
235        fft_ws.fft3d(&mut fx);
236
237        fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
238        for i in 0..n {
239            z_dx[i] = shrink(x_dx[i] + s_dx[i], alpha_over_mu);
240            z_dy[i] = shrink(x_dy[i] + s_dy[i], alpha_over_mu);
241            z_dz[i] = shrink(x_dz[i] + s_dz[i], alpha_over_mu);
242            s_dx[i] += x_dx[i] - z_dx[i];
243            s_dy[i] += x_dy[i] - z_dy[i];
244            s_dz[i] += x_dz[i] - z_dz[i];
245        }
246
247        // ---- fidelity auxiliary (z2) via Newton ---------------------------
248        // Dx = real(ifft(K .* Fx)) ; rhs_z2 = mu2 * (Dx + s2) ; z2 init = Dx + s2.
249        for i in 0..n {
250            xhat[i] = fx[i] * k[i];
251        }
252        fft_ws.ifft3d(&mut xhat);
253        for i in 0..n {
254            dx[i] = xhat[i].re;
255            rhs_z2[i] = mu2 * (dx[i] + s2[i]);
256            z2[i] = rhs_z2[i] / mu2;
257        }
258
259        let mut delta = f64::INFINITY;
260        let mut inn = 0usize;
261        while delta > params.tol_delta && inn < 10 {
262            inn += 1;
263            let norm_old = norm2(&z2);
264            for i in 0..n {
265                let a = z2[i] - phase[i];
266                let numer = w[i] * a.sin() + mu2 * z2[i] - rhs_z2[i];
267                let denom = w[i] * a.cos() + mu2;
268                update[i] = numer / denom;
269                z2[i] -= update[i];
270            }
271            delta = if norm_old > 0.0 {
272                norm2(&update) / norm_old
273            } else {
274                0.0
275            };
276        }
277
278        // ---- multiplier update --------------------------------------------
279        // s2 = s2 + Dx - z2.
280        for i in 0..n {
281            s2[i] += dx[i] - z2[i];
282        }
283    }
284
285    if params.phase_scale != 1.0 {
286        for v in &mut x {
287            *v /= params.phase_scale;
288        }
289    }
290
291    apply_mask_zero(&mut x, mask);
292    x
293}
294
295/// Multiply a complex spectral array in place by a complex multiplier array.
296#[inline]
297fn spectral_mul_assign(dst: &mut [Complex64], m: &[Complex64]) {
298    for (d, &mm) in dst.iter_mut().zip(m.iter()) {
299        *d *= mm;
300    }
301}
302
303/// Nonlinear total-generalized-variation dipole inversion (FANSI `nlTGV.m`).
304///
305/// Everything (operators + normal-equation cofactors) is built from the local,
306/// *unscaled* spectral gradient multipliers `E1,E2,E3` to stay self-consistent
307/// with the MATLAB Cramer-rule algebra.
308#[allow(clippy::too_many_lines)]
309fn nltgv(
310    local_field: &[f64],
311    mask: &[u8],
312    grid: &Grid,
313    bdir: (f64, f64, f64),
314    params: &FansiParams,
315    mut progress: impl FnMut(usize, usize),
316) -> Vec<f64> {
317    let n = grid.n_total();
318    let (nx, ny, nz) = (grid.nx(), grid.ny(), grid.nz());
319
320    // Reuse only the real dipole kernel K from the shared prep.
321    let (mut fft_ws, k, _ee2) = prepare_fansi_spectral(grid, bdir);
322
323    // Local unscaled spectral gradient multipliers, MATLAB order i+j*nx+k*nx*ny.
324    let mut e1 = vec![Complex64::new(0.0, 0.0); n];
325    let mut e2 = vec![Complex64::new(0.0, 0.0); n];
326    let mut e3 = vec![Complex64::new(0.0, 0.0); n];
327    let two_pi = 2.0 * PI;
328    for kk in 0..nz {
329        let ez = Complex64::new(0.0, 1.0) * (two_pi * (kk as f64) / (nz as f64));
330        let e3v = Complex64::new(1.0, 0.0) - ez.exp();
331        for jj in 0..ny {
332            let ey = Complex64::new(0.0, 1.0) * (two_pi * (jj as f64) / (ny as f64));
333            let e2v = Complex64::new(1.0, 0.0) - ey.exp();
334            for ii in 0..nx {
335                let ex = Complex64::new(0.0, 1.0) * (two_pi * (ii as f64) / (nx as f64));
336                let e1v = Complex64::new(1.0, 0.0) - ex.exp();
337                let idx = ii + jj * nx + kk * nx * ny;
338                e1[idx] = e1v;
339                e2[idx] = e2v;
340                e3[idx] = e3v;
341            }
342        }
343    }
344
345    let phase: Vec<f64> = local_field.iter().map(|&f| f * params.phase_scale).collect();
346    let w: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
347
348    let mu0 = params.mu0;
349    let mu1 = params.mu1;
350    let mu2 = params.mu2;
351
352    // Precompute per-voxel normal-equation cofactors (all complex).
353    let mut d11 = vec![Complex64::new(0.0, 0.0); n];
354    let mut d21 = d11.clone();
355    let mut d31 = d11.clone();
356    let mut d41 = d11.clone();
357    let mut d12 = d11.clone();
358    let mut d22 = d11.clone();
359    let mut d32 = d11.clone();
360    let mut d42 = d11.clone();
361    let mut d13 = d11.clone();
362    let mut d23 = d11.clone();
363    let mut d33 = d11.clone();
364    let mut d43 = d11.clone();
365    let mut d14 = d11.clone();
366    let mut d24 = d11.clone();
367    let mut d34 = d11.clone();
368    let mut d44 = d11.clone();
369    let mut det_ainv = d11.clone();
370
371    let half = 0.5;
372    for i in 0..n {
373        let e1i = e1[i];
374        let e2i = e2[i];
375        let e3i = e3[i];
376        let et1 = e1i.conj();
377        let et2 = e2i.conj();
378        let et3 = e3i.conj();
379
380        let e1te1 = et1 * e1i;
381        let e2te2 = et2 * e2i;
382        let e3te3 = et3 * e3i;
383        let mu0h_e1te2 = (mu0 * half) * et1 * e2i;
384        let mu0h_e1te3 = (mu0 * half) * et1 * e3i;
385        let mu0h_e2te3 = (mu0 * half) * et2 * e3i;
386
387        // a0 = mu2 * conj(K) .* K = mu2 * k * k (real).
388        let a0 = Complex64::new(mu2 * k[i] * k[i], 0.0);
389        let a1 = a0 + mu1 * (e1te1 + e2te2 + e3te3);
390        let a2 = Complex64::new(mu1, 0.0) + mu0 * (e1te1 + (e2te2 + e3te3) * half);
391        let a3 = Complex64::new(mu1, 0.0) + mu0 * (e1te1 * half + e2te2 + e3te3 * half);
392        let a4 = Complex64::new(mu1, 0.0) + mu0 * ((e1te1 + e2te2) * half + e3te3);
393        let a5 = -mu1 * e1i;
394        let a6 = -mu1 * e2i;
395        let a7 = mu0h_e1te2;
396        let a8 = -mu1 * e3i;
397        let a9 = mu0h_e1te3;
398        let a10 = mu0h_e2te3;
399        let a5t = a5.conj();
400        let a6t = a6.conj();
401        let a7t = a7.conj();
402        let a8t = a8.conj();
403        let a9t = a9.conj();
404        let a10t = a10.conj();
405
406        let c11 = a2 * a3 * a4 + a7t * a9 * a10t + a7 * a9t * a10
407            - a3 * a9 * a9t
408            - a2 * a10 * a10t
409            - a4 * a7 * a7t;
410        let c21 = a3 * a4 * a5t + a6t * a9 * a10t + a7 * a8t * a10
411            - a3 * a8t * a9
412            - a5t * a10 * a10t
413            - a4 * a6t * a7;
414        let c31 = a4 * a5t * a7t + a6t * a9 * a9t + a2 * a8t * a10
415            - a7t * a8t * a9
416            - a5t * a9t * a10
417            - a2 * a4 * a6t;
418        let c41 = a5t * a7t * a10t + a6t * a7 * a9t + a2 * a3 * a8t
419            - a7 * a7t * a8t
420            - a3 * a5t * a9t
421            - a2 * a6t * a10t;
422        let c12 = a3 * a4 * a5 + a7t * a8 * a10t + a6 * a9t * a10
423            - a3 * a8 * a9t
424            - a5 * a10 * a10t
425            - a4 * a6 * a7t;
426        let c22 = a1 * a3 * a4 + a6t * a8 * a10t + a6 * a8t * a10
427            - a3 * a8 * a8t
428            - a1 * a10 * a10t
429            - a4 * a6 * a6t;
430        let c32 = a1 * a4 * a7t + a6t * a8 * a9t + a5 * a8t * a10
431            - a7t * a8 * a8t
432            - a1 * a9t * a10
433            - a4 * a5 * a6t;
434        let c42 = a1 * a7t * a10t + a6 * a6t * a9t + a3 * a5 * a8t
435            - a6 * a7t * a8t
436            - a1 * a3 * a9t
437            - a5 * a6t * a10t;
438        let c13 = a4 * a5 * a7 + a2 * a8 * a10t + a6 * a9 * a9t
439            - a7 * a8 * a9t
440            - a5 * a9 * a10t
441            - a2 * a4 * a6;
442        let c23 = a1 * a4 * a7 + a5t * a8 * a10t + a6 * a8t * a9
443            - a7 * a8 * a8t
444            - a1 * a9 * a10t
445            - a4 * a5t * a6;
446        let c33 = a1 * a2 * a4 + a5t * a8 * a9t + a5 * a8t * a9
447            - a2 * a8 * a8t
448            - a1 * a9 * a9t
449            - a4 * a5 * a5t;
450        let c43 = a1 * a2 * a10t + a5t * a6 * a9t + a5 * a7 * a8t
451            - a2 * a6 * a8t
452            - a1 * a7 * a9t
453            - a5 * a5t * a10t;
454        let c14 = a5 * a7 * a10 + a2 * a3 * a8 + a6 * a7t * a9
455            - a7 * a7t * a8
456            - a3 * a5 * a9
457            - a2 * a6 * a10;
458        let c24 = a1 * a7 * a10 + a3 * a5t * a8 + a6 * a6t * a9
459            - a6t * a7 * a8
460            - a1 * a3 * a9
461            - a5t * a6 * a10;
462        let c34 = a1 * a2 * a10 + a5t * a7t * a8 + a5 * a6t * a9
463            - a2 * a6t * a8
464            - a1 * a7t * a9
465            - a5 * a5t * a10;
466        let c44 = a1 * a2 * a3 + a5t * a6 * a7t + a5 * a6t * a7
467            - a2 * a6 * a6t
468            - a1 * a7 * a7t
469            - a3 * a5 * a5t;
470
471        let det_a = a1 * c11 - a5 * c21 + a6 * c31 - a8 * c41;
472
473        d11[i] = c11;
474        d21[i] = c21;
475        d31[i] = c31;
476        d41[i] = c41;
477        d12[i] = c12;
478        d22[i] = c22;
479        d32[i] = c32;
480        d42[i] = c42;
481        d13[i] = c13;
482        d23[i] = c23;
483        d33[i] = c33;
484        d43[i] = c43;
485        d14[i] = c14;
486        d24[i] = c24;
487        d34[i] = c34;
488        d44[i] = c44;
489        // Guard the singular (DC null-space) bins: at k=0 every spectral
490        // operator vanishes so det_a -> 0. Zero the solve there instead of
491        // dividing FFT round-off by ~eps (which would create a huge DC pedestal).
492        det_ainv[i] = if det_a.norm() > 1e-20 {
493            Complex64::new(1.0, 0.0) / det_a
494        } else {
495            Complex64::new(0.0, 0.0)
496        };
497    }
498
499    // Conjugates of E used in RHS assembly.
500    let et1: Vec<Complex64> = e1.iter().map(|c| c.conj()).collect();
501    let et2: Vec<Complex64> = e2.iter().map(|c| c.conj()).collect();
502    let et3: Vec<Complex64> = e3.iter().map(|c| c.conj()).collect();
503
504    // ADMM state.
505    let mut x = vec![0.0f64; n];
506    let mut x_prev = vec![0.0f64; n];
507    let mut v1 = vec![0.0f64; n];
508    let mut v2 = vec![0.0f64; n];
509    let mut v3 = vec![0.0f64; n];
510
511    // First-order splits.
512    let mut z1_1 = vec![0.0f64; n];
513    let mut z1_2 = vec![0.0f64; n];
514    let mut z1_3 = vec![0.0f64; n];
515    let mut s1_1 = vec![0.0f64; n];
516    let mut s1_2 = vec![0.0f64; n];
517    let mut s1_3 = vec![0.0f64; n];
518
519    // Symmetric (second-order) splits.
520    let mut z0_1 = vec![0.0f64; n];
521    let mut z0_2 = vec![0.0f64; n];
522    let mut z0_3 = vec![0.0f64; n];
523    let mut z0_4 = vec![0.0f64; n];
524    let mut z0_5 = vec![0.0f64; n];
525    let mut z0_6 = vec![0.0f64; n];
526    let mut s0_1 = vec![0.0f64; n];
527    let mut s0_2 = vec![0.0f64; n];
528    let mut s0_3 = vec![0.0f64; n];
529    let mut s0_4 = vec![0.0f64; n];
530    let mut s0_5 = vec![0.0f64; n];
531    let mut s0_6 = vec![0.0f64; n];
532
533    // Fidelity auxiliary z2 = W .* phase ./ (W + mu2).
534    let mut z2 = vec![0.0f64; n];
535    for i in 0..n {
536        let den = w[i] + mu2;
537        z2[i] = if den != 0.0 { w[i] * phase[i] / den } else { 0.0 };
538    }
539    let mut s2 = vec![0.0f64; n];
540
541    let alpha1_over_mu1 = params.alpha1 / mu1;
542    let alpha0_over_mu0 = params.alpha0 / mu0;
543
544    // Reusable complex/real buffers.
545    let mut rhs1 = vec![Complex64::new(0.0, 0.0); n];
546    let mut rhs2 = rhs1.clone();
547    let mut rhs3 = rhs1.clone();
548    let mut rhs4 = rhs1.clone();
549    let mut t1 = rhs1.clone();
550    let mut t2 = rhs1.clone();
551    let mut t3 = rhs1.clone();
552    let mut fx = rhs1.clone();
553    let mut fv1 = rhs1.clone();
554    let mut fv2 = rhs1.clone();
555    let mut fv3 = rhs1.clone();
556    let mut cbuf = rhs1.clone();
557
558    let mut dx1 = vec![0.0f64; n];
559    let mut dx2 = vec![0.0f64; n];
560    let mut dx3 = vec![0.0f64; n];
561    let mut ev1 = vec![0.0f64; n];
562    let mut ev2 = vec![0.0f64; n];
563    let mut ev3 = vec![0.0f64; n];
564    let mut ev4 = vec![0.0f64; n];
565    let mut ev5 = vec![0.0f64; n];
566    let mut ev6 = vec![0.0f64; n];
567    let mut dx = vec![0.0f64; n];
568    let mut rhs_z2 = vec![0.0f64; n];
569    let mut diff = vec![0.0f64; n];
570    let mut update = vec![0.0f64; n];
571
572    // Scratch real buffer for building fft inputs.
573    let mut rbuf = vec![0.0f64; n];
574
575    // real(fft) of a real field `src` into complex buffer `dst`.
576    macro_rules! fft_real {
577        ($dst:expr, $src:expr) => {{
578            for i in 0..n {
579                $dst[i] = Complex64::new($src[i], 0.0);
580            }
581            fft_ws.fft3d(&mut $dst);
582        }};
583    }
584    // real(fft) of (a - b) elementwise into complex buffer `dst`.
585    macro_rules! fft_real_diff {
586        ($dst:expr, $a:expr, $b:expr) => {{
587            for i in 0..n {
588                rbuf[i] = $a[i] - $b[i];
589            }
590            fft_real!($dst, rbuf);
591        }};
592    }
593
594    for t in 0..params.max_iter {
595        progress(t + 1, params.max_iter);
596
597        // ---- assemble RHS (spectral) --------------------------------------
598        // rhs1 = mu2*conj(K)*fft(z2 - s2) + mu1*(Et*fft(z1 - s1))
599        for i in 0..n {
600            cbuf[i] = Complex64::new(z2[i] - s2[i], 0.0);
601        }
602        fft_ws.fft3d(&mut cbuf);
603        for i in 0..n {
604            rhs1[i] = mu2 * k[i] * cbuf[i];
605        }
606        fft_real_diff!(t1, z1_1, s1_1);
607        fft_real_diff!(t2, z1_2, s1_2);
608        fft_real_diff!(t3, z1_3, s1_3);
609        for i in 0..n {
610            rhs1[i] += mu1 * (et1[i] * t1[i] + et2[i] * t2[i] + et3[i] * t3[i]);
611        }
612
613        // rhs2 = -mu1*fft(z1_1 - s1_1) + mu0*(Et1*fft(z0_1 - s0_1) + Et2*fft(z0_4 - s0_4) + Et3*fft(z0_5 - s0_5))
614        // t1 already holds fft(z1_1 - s1_1).
615        for i in 0..n {
616            rhs2[i] = -mu1 * t1[i];
617            rhs3[i] = -mu1 * t2[i];
618            rhs4[i] = -mu1 * t3[i];
619        }
620        // rhs2 second-order part.
621        fft_real_diff!(t1, z0_1, s0_1);
622        fft_real_diff!(t2, z0_4, s0_4);
623        fft_real_diff!(t3, z0_5, s0_5);
624        for i in 0..n {
625            rhs2[i] += mu0 * (et1[i] * t1[i] + et2[i] * t2[i] + et3[i] * t3[i]);
626        }
627        // rhs3 second-order part: mu0*(Et2*fft(z0_2 - s0_2) + Et1*fft(z0_4 - s0_4) + Et3*fft(z0_6 - s0_6))
628        fft_real_diff!(t1, z0_2, s0_2);
629        fft_real_diff!(t2, z0_4, s0_4);
630        fft_real_diff!(t3, z0_6, s0_6);
631        for i in 0..n {
632            rhs3[i] += mu0 * (et2[i] * t1[i] + et1[i] * t2[i] + et3[i] * t3[i]);
633        }
634        // rhs4 second-order part: mu0*(Et3*fft(z0_3 - s0_3) + Et1*fft(z0_5 - s0_5) + Et2*fft(z0_6 - s0_6))
635        fft_real_diff!(t1, z0_3, s0_3);
636        fft_real_diff!(t2, z0_5, s0_5);
637        fft_real_diff!(t3, z0_6, s0_6);
638        for i in 0..n {
639            rhs4[i] += mu0 * (et3[i] * t1[i] + et1[i] * t2[i] + et2[i] * t3[i]);
640        }
641
642        // ---- Cramer solve -------------------------------------------------
643        for i in 0..n {
644            let r1 = rhs1[i];
645            let r2 = rhs2[i];
646            let r3 = rhs3[i];
647            let r4 = rhs4[i];
648            let da = det_ainv[i];
649            fx[i] = (r1 * d11[i] - r2 * d21[i] + r3 * d31[i] - r4 * d41[i]) * da;
650            fv1[i] = (-r1 * d12[i] + r2 * d22[i] - r3 * d32[i] + r4 * d42[i]) * da;
651            fv2[i] = (r1 * d13[i] - r2 * d23[i] + r3 * d33[i] - r4 * d43[i]) * da;
652            fv3[i] = (-r1 * d14[i] + r2 * d24[i] - r3 * d34[i] + r4 * d44[i]) * da;
653        }
654        // x, v via ifft (take real part). Use t1 as scratch to preserve the
655        // spectral fx/fv (re-FFT'd after convergence anyway).
656        t1.copy_from_slice(&fx);
657        fft_ws.ifft3d(&mut t1);
658        x_prev.copy_from_slice(&x);
659        for i in 0..n {
660            x[i] = t1[i].re;
661        }
662        t1.copy_from_slice(&fv1);
663        fft_ws.ifft3d(&mut t1);
664        for i in 0..n {
665            v1[i] = t1[i].re;
666        }
667        t1.copy_from_slice(&fv2);
668        fft_ws.ifft3d(&mut t1);
669        for i in 0..n {
670            v2[i] = t1[i].re;
671        }
672        t1.copy_from_slice(&fv3);
673        fft_ws.ifft3d(&mut t1);
674        for i in 0..n {
675            v3[i] = t1[i].re;
676        }
677
678        // ---- convergence --------------------------------------------------
679        let xnorm = norm2(&x);
680        if xnorm > 0.0 {
681            for i in 0..n {
682                diff[i] = x[i] - x_prev[i];
683            }
684            let x_update = 100.0 * norm2(&diff) / xnorm;
685            if x_update < params.tol_update || x_update.is_nan() {
686                progress(t + 1, t + 1);
687                break;
688            }
689        }
690
691        if t + 1 >= params.max_iter {
692            break;
693        }
694
695        // ---- re-FFT for stability -----------------------------------------
696        fft_real!(fx, x);
697        fft_real!(fv1, v1);
698        fft_real!(fv2, v2);
699        fft_real!(fv3, v3);
700
701        // Dx1 = real(ifft(E1.*Fx)), etc.
702        cbuf.copy_from_slice(&fx);
703        spectral_mul_assign(&mut cbuf, &e1);
704        fft_ws.ifft3d(&mut cbuf);
705        for i in 0..n {
706            dx1[i] = cbuf[i].re;
707        }
708        cbuf.copy_from_slice(&fx);
709        spectral_mul_assign(&mut cbuf, &e2);
710        fft_ws.ifft3d(&mut cbuf);
711        for i in 0..n {
712            dx2[i] = cbuf[i].re;
713        }
714        cbuf.copy_from_slice(&fx);
715        spectral_mul_assign(&mut cbuf, &e3);
716        fft_ws.ifft3d(&mut cbuf);
717        for i in 0..n {
718            dx3[i] = cbuf[i].re;
719        }
720
721        // E_v1 = real(ifft(E1.*Fv1)), E_v2 = ...E2.*Fv2, E_v3 = ...E3.*Fv3.
722        cbuf.copy_from_slice(&fv1);
723        spectral_mul_assign(&mut cbuf, &e1);
724        fft_ws.ifft3d(&mut cbuf);
725        for i in 0..n {
726            ev1[i] = cbuf[i].re;
727        }
728        cbuf.copy_from_slice(&fv2);
729        spectral_mul_assign(&mut cbuf, &e2);
730        fft_ws.ifft3d(&mut cbuf);
731        for i in 0..n {
732            ev2[i] = cbuf[i].re;
733        }
734        cbuf.copy_from_slice(&fv3);
735        spectral_mul_assign(&mut cbuf, &e3);
736        fft_ws.ifft3d(&mut cbuf);
737        for i in 0..n {
738            ev3[i] = cbuf[i].re;
739        }
740
741        // E_v4 = real(ifft(E1.*Fv2 + E2.*Fv1))/2, etc.
742        for i in 0..n {
743            cbuf[i] = e1[i] * fv2[i] + e2[i] * fv1[i];
744        }
745        fft_ws.ifft3d(&mut cbuf);
746        for i in 0..n {
747            ev4[i] = cbuf[i].re * 0.5;
748        }
749        for i in 0..n {
750            cbuf[i] = e1[i] * fv3[i] + e3[i] * fv1[i];
751        }
752        fft_ws.ifft3d(&mut cbuf);
753        for i in 0..n {
754            ev5[i] = cbuf[i].re * 0.5;
755        }
756        for i in 0..n {
757            cbuf[i] = e2[i] * fv3[i] + e3[i] * fv2[i];
758        }
759        fft_ws.ifft3d(&mut cbuf);
760        for i in 0..n {
761            ev6[i] = cbuf[i].re * 0.5;
762        }
763
764        // ---- symmetric-gradient split (shrink) ----------------------------
765        for i in 0..n {
766            z0_1[i] = shrink(ev1[i] + s0_1[i], alpha0_over_mu0);
767            z0_2[i] = shrink(ev2[i] + s0_2[i], alpha0_over_mu0);
768            z0_3[i] = shrink(ev3[i] + s0_3[i], alpha0_over_mu0);
769            z0_4[i] = shrink(ev4[i] + s0_4[i], alpha0_over_mu0);
770            z0_5[i] = shrink(ev5[i] + s0_5[i], alpha0_over_mu0);
771            z0_6[i] = shrink(ev6[i] + s0_6[i], alpha0_over_mu0);
772        }
773
774        // ---- first-order split (shrink) -----------------------------------
775        for i in 0..n {
776            z1_1[i] = shrink(dx1[i] - v1[i] + s1_1[i], alpha1_over_mu1);
777            z1_2[i] = shrink(dx2[i] - v2[i] + s1_2[i], alpha1_over_mu1);
778            z1_3[i] = shrink(dx3[i] - v3[i] + s1_3[i], alpha1_over_mu1);
779        }
780
781        // ---- fidelity auxiliary (z2) via Newton ---------------------------
782        // Dx = real(ifft(K.*Fx)) ; rhs_z2 = mu2*(Dx + s2) ; z2 init = Dx + s2.
783        for i in 0..n {
784            cbuf[i] = fx[i] * k[i];
785        }
786        fft_ws.ifft3d(&mut cbuf);
787        for i in 0..n {
788            dx[i] = cbuf[i].re;
789            rhs_z2[i] = mu2 * (dx[i] + s2[i]);
790            z2[i] = rhs_z2[i] / mu2;
791        }
792        let mut delta = f64::INFINITY;
793        let mut inn = 0usize;
794        while delta > params.tol_delta && inn < 50 {
795            inn += 1;
796            let norm_old = norm2(&z2);
797            for i in 0..n {
798                let a = z2[i] - phase[i];
799                let numer = w[i] * a.sin() + mu2 * z2[i] - rhs_z2[i];
800                let denom = w[i] * a.cos() + mu2;
801                update[i] = numer / denom;
802                z2[i] -= update[i];
803            }
804            delta = if norm_old > 0.0 {
805                norm2(&update) / norm_old
806            } else {
807                0.0
808            };
809        }
810
811        // ---- multiplier updates -------------------------------------------
812        for i in 0..n {
813            s0_1[i] += ev1[i] - z0_1[i];
814            s0_2[i] += ev2[i] - z0_2[i];
815            s0_3[i] += ev3[i] - z0_3[i];
816            s0_4[i] += ev4[i] - z0_4[i];
817            s0_5[i] += ev5[i] - z0_5[i];
818            s0_6[i] += ev6[i] - z0_6[i];
819            s1_1[i] += dx1[i] - v1[i] - z1_1[i];
820            s1_2[i] += dx2[i] - v2[i] - z1_2[i];
821            s1_3[i] += dx3[i] - v3[i] - z1_3[i];
822            s2[i] += dx[i] - z2[i];
823        }
824    }
825
826    if params.phase_scale != 1.0 {
827        for v in &mut x {
828            *v /= params.phase_scale;
829        }
830    }
831
832    apply_mask_zero(&mut x, mask);
833    x
834}
835
836#[cfg(test)]
837mod tests {
838    use super::*;
839
840    #[test]
841    fn test_fansi_nltv_zero_field() {
842        // Zero field should give (approximately) zero susceptibility.
843        let n = 8;
844        let field = vec![0.0; n * n * n];
845        let mask = vec![1u8; n * n * n];
846        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
847        let params = FansiParams {
848            max_iter: 10,
849            is_tgv: false,
850            ..FansiParams::default()
851        };
852
853        let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
854
855        for &val in chi.iter() {
856            assert!(val.abs() < 1e-6, "Zero field should give ~zero chi, got {}", val);
857        }
858    }
859
860    #[test]
861    fn test_fansi_nltv_finite() {
862        // A small ramp field should produce all-finite output.
863        let n = 8;
864        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
865        let mask = vec![1u8; n * n * n];
866        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
867        let params = FansiParams {
868            max_iter: 10,
869            is_tgv: false,
870            ..FansiParams::default()
871        };
872
873        let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
874
875        for (i, &val) in chi.iter().enumerate() {
876            assert!(val.is_finite(), "Chi should be finite at index {}", i);
877        }
878    }
879
880    #[test]
881    fn test_fansi_nltgv_zero_field() {
882        // Zero field should give (approximately) zero susceptibility.
883        let n = 8;
884        let field = vec![0.0; n * n * n];
885        let mask = vec![1u8; n * n * n];
886        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
887        let params = FansiParams {
888            max_iter: 10,
889            is_tgv: true,
890            ..FansiParams::default()
891        };
892
893        let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
894
895        for &val in chi.iter() {
896            assert!(val.abs() < 1e-6, "Zero field should give ~zero chi, got {}", val);
897        }
898    }
899
900    #[test]
901    fn test_fansi_nltgv_finite() {
902        // A small ramp field should produce all-finite output.
903        let n = 8;
904        let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
905        let mask = vec![1u8; n * n * n];
906        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
907        let params = FansiParams {
908            max_iter: 10,
909            is_tgv: true,
910            ..FansiParams::default()
911        };
912
913        let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), &params, |_, _| {});
914
915        for (i, &val) in chi.iter().enumerate() {
916            assert!(val.is_finite(), "Chi should be finite at index {}", i);
917        }
918    }
919}