Skip to main content

qsm_core/inversion/
medi.rs

1//! MEDI (Morphology Enabled Dipole Inversion) L1 regularization
2//!
3//! Gauss-Newton optimization with L1 TV regularization and
4//! morphology-based gradient weighting from magnitude images.
5//!
6//! Features:
7//! - **Per-direction gradient masks** (mx, my, mz) matching the original MEDI formulation
8//! - Adaptive edge detection with configurable percentage threshold (default: 30% edges)
9//! - SNR-based data weighting using noise standard deviation maps
10//! - Optional SMV (Spherical Mean Value) preprocessing
11//! - Optional merit-based outlier adjustment (MERIT)
12//! - **Optimized with f32 single precision for WASM performance**
13//! - **Buffer reuse to minimize allocations**
14//! - **Standard CG convergence (relative tolerance)**
15//! - **Linear extrapolation boundary conditions** matching MATLAB's gradf
16//!
17//! Reference:
18//! Liu, T., Liu, J., de Rochefort, L., Spincemaille, P., Khalidov, I., Ledoux, J.R.,
19//! Wang, Y. (2011). "Morphology enabled dipole inversion (MEDI) from a single-angle
20//! acquisition: comparison with COSMOS in human brain imaging."
21//! Magnetic Resonance in Medicine, 66(3):777-783. https://doi.org/10.1002/mrm.22816
22//!
23//! Liu, J., Liu, T., de Rochefort, L., Ledoux, J., Khalidov, I., Chen, W., Tsiouris, A.J.,
24//! Wisnieff, C., Spincemaille, P., Prince, M.R., Wang, Y. (2012).
25//! "Morphology enabled dipole inversion for quantitative susceptibility mapping using
26//! structural consistency between the magnitude image and the susceptibility map."
27//! NeuroImage, 59(3):2560-2568.
28//!
29//! Reference implementation: https://github.com/huawu02/MEDI_toolbox
30
31use num_complex::Complex32;
32#[cfg(feature = "parallel")]
33use rayon::prelude::*;
34use crate::fft::Fft3dWorkspaceF32;
35use crate::kernels::dipole::dipole_kernel_f32;
36use crate::kernels::smv::smv_kernel_f32;
37use crate::utils::simd_ops::{
38    dot_product_f32, norm_squared_f32, axpy_f32, xpby_f32,
39    apply_gradient_weights_f32, compute_p_weights_f32, combine_terms_f32, negate_f32,
40};
41use crate::Grid;
42// Note: Uses fgrad_periodic_inplace_f32 / bdiv_periodic_inplace_f32 (periodic BCs)
43// for the MEDI inner loop matching MATLAB's gradfp_mex / gradfp_adj_mex.
44// Uses fgrad_linext_inplace_f32 (linear extrapolation BCs) only for gradient mask
45// computation, matching MATLAB's gradf_mex.
46
47/// MEDI algorithm parameters
48#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
49#[derive(Clone, Debug)]
50pub struct MediParams {
51    /// Regularization weight
52    pub lambda: f64,
53    /// Enable MERIT (outlier adjustment)
54    pub merit: bool,
55    /// Enable SMV preprocessing
56    pub smv: bool,
57    /// SMV radius in mm
58    pub smv_radius: f64,
59    /// Data weighting mode (1 = SNR)
60    pub data_weighting: i32,
61    /// Fraction of voxels considered edges (0.0-1.0)
62    pub percentage: f64,
63    /// CG convergence tolerance
64    pub cg_tol: f64,
65    /// Maximum CG iterations
66    pub cg_max_iter: usize,
67    /// Maximum outer iterations
68    pub max_iter: usize,
69    /// Outer convergence tolerance
70    pub tol: f64,
71}
72
73impl Default for MediParams {
74    fn default() -> Self {
75        Self {
76            // Tuned for the field in RADIANS (the scale `medi` expects — the data
77            // term uses exp(i·field)). At 7T/4ms a typical local field is ~0.1 rad
78            // std; 1e-3 matches Cornell MEDI's smoothness there. A smaller λ that
79            // "looks fine" is usually a symptom of feeding a ppm-scale field.
80            lambda: 1e-3,
81            merit: false,
82            smv: true,
83            smv_radius: 5.0,
84            data_weighting: 1,
85            percentage: 0.3,
86            cg_tol: 0.01,
87            cg_max_iter: 10,
88            max_iter: 30,
89            tol: 0.1,
90        }
91    }
92}
93
94/// Workspace for MEDI operations - holds all reusable buffers (f32 version)
95/// Uses single precision for ~2x speedup on WASM
96pub struct MediWorkspace {
97    pub n_total: usize,
98    pub nx: usize,
99    pub ny: usize,
100    pub nz: usize,
101    pub vsx: f32,
102    pub vsy: f32,
103    pub vsz: f32,
104
105    // FFT workspace with cached plans (f32)
106    pub fft_ws: Fft3dWorkspaceF32,
107
108    // Gradient buffers (3 components)
109    pub gx: Vec<f32>,
110    pub gy: Vec<f32>,
111    pub gz: Vec<f32>,
112
113    // Weighted gradient buffers
114    pub reg_x: Vec<f32>,
115    pub reg_y: Vec<f32>,
116    pub reg_z: Vec<f32>,
117
118    // Divergence buffer
119    pub div_buf: Vec<f32>,
120
121    // Complex buffer for FFT operations
122    pub complex_buf: Vec<Complex32>,
123    pub complex_buf2: Vec<Complex32>,
124
125    // Real buffer for dipole result
126    pub dipole_buf: Vec<f32>,
127
128    // CG solver buffers
129    pub cg_r: Vec<f32>,
130    pub cg_p: Vec<f32>,
131    pub cg_ap: Vec<f32>,
132}
133
134impl MediWorkspace {
135    /// Create a new MEDI workspace for the given grid
136    pub fn new(grid: &Grid) -> Self {
137        let (nx, ny, nz) = grid.dims;
138        let (vsx, vsy, vsz) = (grid.vsx() as f32, grid.vsy() as f32, grid.vsz() as f32);
139        let n_total = grid.n_total();
140
141        Self {
142            n_total,
143            nx, ny, nz,
144            vsx, vsy, vsz,
145            fft_ws: Fft3dWorkspaceF32::new(nx, ny, nz),
146            gx: vec![0.0; n_total],
147            gy: vec![0.0; n_total],
148            gz: vec![0.0; n_total],
149            reg_x: vec![0.0; n_total],
150            reg_y: vec![0.0; n_total],
151            reg_z: vec![0.0; n_total],
152            div_buf: vec![0.0; n_total],
153            complex_buf: vec![Complex32::new(0.0, 0.0); n_total],
154            complex_buf2: vec![Complex32::new(0.0, 0.0); n_total],
155            dipole_buf: vec![0.0; n_total],
156            cg_r: vec![0.0; n_total],
157            cg_p: vec![0.0; n_total],
158            cg_ap: vec![0.0; n_total],
159        }
160    }
161}
162
163/// Apply dipole convolution: out = real(ifft(D * fft(x)))
164#[inline]
165pub(crate) fn apply_dipole_conv(
166    fft_ws: &mut Fft3dWorkspaceF32,
167    x: &[f32],
168    d_kernel: &[f32],
169    out: &mut [f32],
170    complex_buf: &mut [Complex32],
171) {
172    fft_ws.apply_dipole_inplace(x, d_kernel, out, complex_buf);
173}
174
175/// MEDI operator buffers - separate struct to allow split borrowing
176pub(crate) struct MediOpBuffers<'a> {
177    pub gx: &'a mut [f32],
178    pub gy: &'a mut [f32],
179    pub gz: &'a mut [f32],
180    pub reg_x: &'a mut [f32],
181    pub reg_y: &'a mut [f32],
182    pub reg_z: &'a mut [f32],
183    pub div_buf: &'a mut [f32],
184    pub dipole_buf: &'a mut [f32],
185    pub complex_buf: &'a mut [Complex32],
186    pub complex_buf2: &'a mut [Complex32],
187}
188
189/// Apply MEDI operator in-place: out = fidelity(dx) + lambda*reg(dx)
190/// This is the hot path - called many times per Gauss-Newton iteration
191/// Uses per-direction gradient masks (mx, my, mz) matching MATLAB MEDI
192/// SIMD-accelerated for element-wise operations
193#[inline]
194pub(crate) fn apply_medi_operator_core(
195    fft_ws: &mut Fft3dWorkspaceF32,
196    bufs: &mut MediOpBuffers,
197    n: usize,
198    nx: usize, ny: usize, nz: usize,
199    vsx: f32, vsy: f32, vsz: f32,
200    dx: &[f32],
201    w: &[Complex32],
202    d_kernel: &[f32],
203    mx: &[f32],  // Per-direction gradient mask for x
204    my: &[f32],  // Per-direction gradient mask for y
205    mz: &[f32],  // Per-direction gradient mask for z
206    vr: &[f32],
207    lambda: f32,
208    out: &mut [f32],
209) {
210    // 1. Compute gradient of dx (in-place into gx, gy, gz) - periodic BCs matching MATLAB gradfp_mex
211    fgrad_periodic_inplace_f32(bufs.gx, bufs.gy, bufs.gz, dx, nx, ny, nz, vsx, vsy, vsz);
212
213    // 2. Apply per-direction weights: reg_i = m_i * P * m_i * g_i (SIMD accelerated)
214    // MATLAB: ux = mx .* P .* mx .* ux; uy = my .* P .* my .* uy; uz = mz .* P .* mz .* uz;
215    apply_gradient_weights_f32(
216        bufs.reg_x, bufs.reg_y, bufs.reg_z,
217        mx, my, mz, vr,
218        bufs.gx, bufs.gy, bufs.gz,
219    );
220
221    // 3. Compute divergence (in-place into div_buf) - periodic BCs matching MATLAB gradfp_adj_mex
222    bdiv_periodic_inplace_f32(bufs.div_buf, bufs.reg_x, bufs.reg_y, bufs.reg_z, nx, ny, nz, vsx, vsy, vsz);
223
224    // 4. Fidelity term: D^T(|w|^2 * D(dx))
225    apply_dipole_conv(fft_ws, dx, d_kernel, bufs.dipole_buf, bufs.complex_buf);
226
227    // Multiply by |w|^2 and convert to complex
228    for i in 0..n {
229        let w_mag_sq = w[i].norm_sqr();
230        bufs.complex_buf2[i] = Complex32::new(bufs.dipole_buf[i] * w_mag_sq, 0.0);
231    }
232
233    // Apply D^T (which is D for real symmetric kernel)
234    fft_ws.fft3d(bufs.complex_buf2);
235    for i in 0..n {
236        bufs.complex_buf2[i] *= d_kernel[i];
237    }
238    fft_ws.ifft3d(bufs.complex_buf2);
239
240    // 5. Combine: out = lambda*div_buf + real(complex_buf2) (matching MATLAB: y = D + R)
241    // Extract real parts for SIMD operation
242    for i in 0..n {
243        bufs.dipole_buf[i] = bufs.complex_buf2[i].re;
244    }
245    combine_terms_f32(out, bufs.div_buf, bufs.dipole_buf, lambda);
246}
247
248/// Conjugate gradient solver with buffer reuse
249/// Solves Ax = b where A is the MEDI operator
250///
251/// The optional progress callback receives (cg_iter, max_iter) for each CG iteration.
252/// Uses per-direction gradient masks (mx, my, mz) matching MATLAB MEDI.
253#[inline]
254fn cg_solve_medi<F>(
255    ws: &mut MediWorkspace,
256    w: &[Complex32],
257    d_kernel: &[f32],
258    mx: &[f32],  // Per-direction gradient mask for x
259    my: &[f32],  // Per-direction gradient mask for y
260    mz: &[f32],  // Per-direction gradient mask for z
261    vr: &[f32],
262    lambda: f32,
263    b: &[f32],
264    x: &mut [f32],
265    tol: f32,
266    max_iter: usize,
267    mut progress_callback: F,
268) where
269    F: FnMut(usize, usize),
270{
271    let n = ws.n_total;
272    let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
273    let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
274
275    // Initialize x to zero
276    x.fill(0.0);
277
278    // r = b - A*x = b (since x=0)
279    ws.cg_r.copy_from_slice(b);
280
281    // p = r
282    ws.cg_p.copy_from_slice(&ws.cg_r);
283
284    // rsold = r·r (SIMD accelerated)
285    let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
286
287    // b_norm for relative tolerance (SIMD accelerated)
288    let b_norm: f32 = norm_squared_f32(b).sqrt();
289    if b_norm < 1e-10 {
290        return; // b is zero, x=0 is the solution
291    }
292
293    // Buffer for p (to avoid borrow conflict)
294    let mut p_copy = vec![0.0f32; n];
295
296    for cg_iter in 0..max_iter {
297        // Report CG progress
298        progress_callback(cg_iter + 1, max_iter);
299
300        // Copy p to avoid borrow conflict
301        p_copy.copy_from_slice(&ws.cg_p);
302
303        // ap = A*p - use split borrowing
304        {
305            let mut bufs = MediOpBuffers {
306                gx: &mut ws.gx,
307                gy: &mut ws.gy,
308                gz: &mut ws.gz,
309                reg_x: &mut ws.reg_x,
310                reg_y: &mut ws.reg_y,
311                reg_z: &mut ws.reg_z,
312                div_buf: &mut ws.div_buf,
313                dipole_buf: &mut ws.dipole_buf,
314                complex_buf: &mut ws.complex_buf,
315                complex_buf2: &mut ws.complex_buf2,
316            };
317            apply_medi_operator_core(
318                &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
319                &p_copy, w, d_kernel, mx, my, mz, vr, lambda, &mut ws.cg_ap
320            );
321        }
322
323        // pap = p·ap (SIMD accelerated)
324        let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
325
326        if pap.abs() < 1e-15 {
327            break;
328        }
329
330        let alpha = rsold / pap;
331
332        // x = x + alpha*p (SIMD accelerated)
333        axpy_f32(x, alpha, &ws.cg_p);
334
335        // r = r - alpha*ap (SIMD accelerated)
336        axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
337
338        // rsnew = r·r (SIMD accelerated)
339        let rsnew: f32 = norm_squared_f32(&ws.cg_r);
340        let residual = rsnew.sqrt();
341
342        // Check convergence
343        if residual < tol * b_norm {
344            break;
345        }
346
347        let beta = rsnew / rsold;
348
349        // p = r + beta*p (SIMD accelerated)
350        xpby_f32(&mut ws.cg_p, &ws.cg_r, beta);
351
352        rsold = rsnew;
353    }
354}
355
356/// MEDI L1 dipole inversion (OPTIMIZED f32 VERSION)
357///
358/// # Arguments
359/// * `local_field` - Local field/phase (RDF) in radians (nx * ny * nz)
360/// * `n_std` - Noise standard deviation map (same size as local_field)
361/// * `magnitude` - Magnitude image for gradient weighting (nx * ny * nz)
362/// * `mask` - Binary mask (nx * ny * nz), 1 = brain
363/// * `grid` - Volume grid (dimensions and voxel sizes)
364/// * `bdir` - B0 field direction
365/// * `params` - MEDI parameters
366/// * `progress` - Progress callback `(current_step, total_steps)`
367///
368/// # Returns
369/// Susceptibility map (in same units as input field)
370pub fn medi(
371    local_field: &[f64],
372    n_std: &[f64],
373    magnitude: &[f64],
374    mask: &[u8],
375    grid: &Grid,
376    bdir: (f64, f64, f64),
377    params: &MediParams,
378    mut progress: impl FnMut(usize, usize),
379) -> Vec<f64> {
380    let (nx, ny, nz) = grid.dims;
381    let n_total = grid.n_total();
382
383    // Convert to f32 for internal computation (much faster on WASM)
384    let vsx_f32 = grid.vsx() as f32;
385    let vsy_f32 = grid.vsy() as f32;
386    let vsz_f32 = grid.vsz() as f32;
387    let lambda_f32 = params.lambda as f32;
388    let bdir_f32 = (bdir.0 as f32, bdir.1 as f32, bdir.2 as f32);
389    let smv_radius_f32 = params.smv_radius as f32;
390    let percentage_f32 = params.percentage as f32;
391    let cg_tol_f32 = params.cg_tol as f32;
392    let tol_f32 = params.tol as f32;
393    let max_iter = params.max_iter;
394    let cg_max_iter = params.cg_max_iter;
395    let data_weighting = params.data_weighting;
396
397    // Convert input arrays to f32
398    let local_field_f32: Vec<f32> = local_field.iter().map(|&v| v as f32).collect();
399    let n_std_f32: Vec<f32> = n_std.iter().map(|&v| v as f32).collect();
400    let magnitude_f32: Vec<f32> = magnitude.iter().map(|&v| v as f32).collect();
401
402    // Create workspace - this allocates all buffers ONCE
403    let mut ws = MediWorkspace::new(grid);
404
405    // Working copies that may be modified by SMV preprocessing
406    let mut rdf: Vec<f32> = local_field_f32.clone();
407    let mut work_mask: Vec<u8> = mask.to_vec();
408    let mut tempn: Vec<f32> = n_std_f32.clone();
409
410    // Apply mask to N_std
411    for i in 0..n_total {
412        if mask[i] == 0 {
413            tempn[i] = 0.0;
414        }
415    }
416
417    // Generate dipole kernel
418    let mut d_kernel = dipole_kernel_f32(grid, bdir_f32);
419
420    // SMV preprocessing (optional)
421    let sphere_k = if params.smv {
422        let sk = smv_kernel_f32(grid, smv_radius_f32);
423
424        // FFT of sphere kernel for convolution
425        let mut sk_fft: Vec<Complex32> = sk.iter()
426            .map(|&v| Complex32::new(v, 0.0))
427            .collect();
428        ws.fft_ws.fft3d(&mut sk_fft);
429
430        // Erode mask: SMV(mask) > 0.999
431        let mask_f32: Vec<f32> = work_mask.iter().map(|&m| m as f32).collect();
432        let smv_mask = apply_smv_kernel_ws(&mask_f32, &sk_fft, &mut ws);
433        for i in 0..n_total {
434            work_mask[i] = if smv_mask[i] > 0.999 { 1 } else { 0 };
435        }
436
437        // Modify dipole kernel: D = (1 - SphereK) * D
438        for i in 0..n_total {
439            d_kernel[i] *= 1.0 - sk[i];
440        }
441
442        // Modify RDF: RDF = RDF - SMV(RDF)
443        let smv_rdf = apply_smv_kernel_ws(&rdf, &sk_fft, &mut ws);
444        for i in 0..n_total {
445            rdf[i] -= smv_rdf[i];
446            if work_mask[i] == 0 {
447                rdf[i] = 0.0;
448            }
449        }
450
451        // Modify noise: tempn = sqrt(SMV(tempn^2) + tempn^2)
452        let tempn_sq: Vec<f32> = tempn.iter().map(|&t| t * t).collect();
453        let smv_tempn_sq = apply_smv_kernel_ws(&tempn_sq, &sk_fft, &mut ws);
454        for i in 0..n_total {
455            tempn[i] = (smv_tempn_sq[i] + tempn_sq[i]).sqrt();
456        }
457
458        Some(sk_fft)
459    } else {
460        None
461    };
462
463    // Compute data weighting
464    let mut m = dataterm_mask_f32(data_weighting, &tempn, &work_mask);
465
466    // b0 = m * exp(i * RDF)
467    let mut b0: Vec<Complex32> = rdf.iter()
468        .zip(m.iter())
469        .map(|(&f, &mi)| {
470            let phase = Complex32::new(0.0, f);
471            mi * phase.exp()
472        })
473        .collect();
474
475    // Compute per-direction gradient weighting masks from magnitude edges
476    // Returns (mx, my, mz) - separate masks for each gradient direction (matching MATLAB MEDI)
477    let (w_gx, w_gy, w_gz) = gradient_mask_f32(&magnitude_f32, &work_mask, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32, percentage_f32);
478
479    // Fallback: if any mask is all zeros, use magnitude image (matching MATLAB)
480    let w_gx = if w_gx.iter().any(|&v| v != 0.0) { w_gx } else { magnitude_f32.clone() };
481    let w_gy = if w_gy.iter().any(|&v| v != 0.0) { w_gy } else { magnitude_f32.clone() };
482    let w_gz = if w_gz.iter().any(|&v| v != 0.0) { w_gz } else { magnitude_f32.clone() };
483
484    // Initialize susceptibility
485    let mut chi = vec![0.0f32; n_total];
486    let mut dx = vec![0.0f32; n_total];  // Reusable buffer for CG solution
487    let mut rhs = vec![0.0f32; n_total]; // Reusable buffer for RHS
488    let mut vr = vec![0.0f32; n_total];  // Reusable buffer for Vr (P in MATLAB)
489    let mut w: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); n_total]; // Reusable buffer for w
490    let mut chi_prev = vec![0.0f32; n_total]; // Reusable buffer for convergence check
491    let mut badpoint = vec![0.0f32; n_total];
492    let mut n_std_work: Vec<f32> = n_std_f32.clone();
493
494    // MATLAB: beta = sqrt(eps(class(f))) where eps for f64 ≈ 2.22e-16, so sqrt(eps) ≈ 1.49e-8.
495    // This is a regularization parameter for the P weight denominator, not a precision limit.
496    // Using the same value as MATLAB (1.49e-8) is fine in f32 (representable, well above f32 eps).
497    let beta = 1.49e-8_f32;
498
499    // Total progress = GN iterations * CG iterations per GN
500    let total_steps = max_iter * cg_max_iter;
501
502    // Gauss-Newton iterations
503    for iter in 0..max_iter {
504        // Save chi_prev for convergence check
505        chi_prev.copy_from_slice(&chi);
506
507        // Compute P = 1 / sqrt(|m * grad(chi)|^2 + beta) using per-direction masks (SIMD accelerated)
508        // MATLAB: P = 1 ./ sqrt(ux.*ux + uy.*uy + uz.*uz + beta);
509        // where ux = mx .* grad_x(chi), uy = my .* grad_y(chi), uz = mz .* grad_z(chi)
510        // Uses periodic BCs matching MATLAB's grad_ (which calls gradfp_mex)
511        fgrad_periodic_inplace_f32(
512            &mut ws.gx, &mut ws.gy, &mut ws.gz,
513            &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
514        );
515
516        compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
517
518        // Compute w = m * exp(i * D*chi) using workspace
519        apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
520        for i in 0..n_total {
521            let phase = Complex32::new(0.0, ws.dipole_buf[i]);
522            w[i] = m[i] * phase.exp();
523        }
524
525        // Compute right-hand side using workspace
526        compute_rhs_inplace(&chi, &w, &b0, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda_f32, &mut rhs, &mut ws);
527
528        // Negate for CG (solving A*dx = -b) (SIMD accelerated)
529        negate_f32(&mut rhs);
530
531        // Solve A*dx = rhs using optimized CG with combined progress reporting
532        let gn_iter = iter;
533        cg_solve_medi(
534            &mut ws, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda_f32, &rhs, &mut dx, cg_tol_f32, cg_max_iter,
535            |cg_iter, cg_total| {
536                let current = gn_iter * cg_total + cg_iter;
537                progress(current, total_steps);
538            }
539        );
540
541        // Update: chi = chi + dx (SIMD accelerated)
542        axpy_f32(&mut chi, 1.0, &dx);
543
544        // Check convergence (SIMD accelerated)
545        let norm_dx_sq = norm_squared_f32(&dx);
546        let norm_chi_sq = norm_squared_f32(&chi_prev);
547        let rel_change = norm_dx_sq.sqrt() / (norm_chi_sq.sqrt() + 1e-6);
548
549        // Merit adjustment (optional)
550        if params.merit {
551            // Compute residual: wres = m * exp(i * D*chi) - b0
552            apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
553            let mut wres: Vec<Complex32> = ws.dipole_buf.iter()
554                .zip(m.iter())
555                .zip(b0.iter())
556                .map(|((&dc, &mi), &b0i)| {
557                    let phase = Complex32::new(0.0, dc);
558                    mi * phase.exp() - b0i
559                })
560                .collect();
561
562            // Subtract mean over mask
563            let mask_count = work_mask.iter().filter(|&&m| m != 0).count() as f32;
564            if mask_count > 0.0 {
565                let mean_wres: Complex32 = wres.iter()
566                    .zip(work_mask.iter())
567                    .filter(|(_, &m)| m != 0)
568                    .map(|(w, _)| w)
569                    .sum::<Complex32>() / mask_count;
570
571                for i in 0..n_total {
572                    if work_mask[i] != 0 {
573                        wres[i] -= mean_wres;
574                    }
575                }
576            }
577
578            // Compute factor = std(abs(wres[mask])) * 6
579            let abs_wres: Vec<f32> = wres.iter()
580                .zip(work_mask.iter())
581                .filter(|(_, &m)| m != 0)
582                .map(|(w, _)| w.norm())
583                .collect();
584
585            if !abs_wres.is_empty() {
586                let mean_abs: f32 = abs_wres.iter().sum::<f32>() / abs_wres.len() as f32;
587                let var: f32 = abs_wres.iter()
588                    .map(|&v| (v - mean_abs).powi(2))
589                    .sum::<f32>() / abs_wres.len() as f32;
590                let factor = var.sqrt() * 6.0;
591
592                if factor > 1e-10 {
593                    // Normalize wres by factor
594                    let mut wres_norm: Vec<f32> = wres.iter()
595                        .map(|w| w.norm() / factor)
596                        .collect();
597
598                    // Clamp values < 1 to 1
599                    for v in wres_norm.iter_mut() {
600                        if *v < 1.0 {
601                            *v = 1.0;
602                        }
603                    }
604
605                    // Mark bad points and update noise
606                    for i in 0..n_total {
607                        if wres_norm[i] > 1.0 {
608                            badpoint[i] = 1.0;
609                        }
610                        if work_mask[i] != 0 {
611                            n_std_work[i] *= wres_norm[i].powi(2);
612                        }
613                    }
614
615                    // Recompute tempn
616                    tempn = n_std_work.clone();
617                    if let Some(ref sk_fft) = sphere_k {
618                        let tempn_sq: Vec<f32> = tempn.iter().map(|&t| t * t).collect();
619                        let smv_tempn_sq = apply_smv_kernel_ws(&tempn_sq, sk_fft, &mut ws);
620                        for i in 0..n_total {
621                            tempn[i] = (smv_tempn_sq[i] + tempn_sq[i]).sqrt();
622                        }
623                    }
624
625                    // Recompute data weighting and b0
626                    m = dataterm_mask_f32(data_weighting, &tempn, &work_mask);
627                    b0 = rdf.iter()
628                        .zip(m.iter())
629                        .map(|(&f, &mi)| {
630                            let phase = Complex32::new(0.0, f);
631                            mi * phase.exp()
632                        })
633                        .collect();
634                }
635            }
636        }
637
638        if rel_change < tol_f32 {
639            // Report completion on early convergence
640            progress(total_steps, total_steps);
641            break;
642        }
643    }
644
645    // Suppress unused variable warning
646    let _ = badpoint;
647
648    // Apply mask and convert back to f64
649    chi.iter()
650        .zip(mask.iter())
651        .map(|(&c, &m)| if m == 0 { 0.0 } else { c as f64 })
652        .collect()
653}
654
655/// Apply SMV kernel using workspace buffers (f32)
656fn apply_smv_kernel_ws(
657    x: &[f32],
658    sk_fft: &[Complex32],
659    ws: &mut MediWorkspace,
660) -> Vec<f32> {
661    let n_total = ws.n_total;
662
663    // Copy to complex buffer
664    for (c, &r) in ws.complex_buf.iter_mut().zip(x.iter()) {
665        *c = Complex32::new(r, 0.0);
666    }
667
668    ws.fft_ws.fft3d(&mut ws.complex_buf);
669
670    for i in 0..n_total {
671        ws.complex_buf[i] *= sk_fft[i];
672    }
673
674    ws.fft_ws.ifft3d(&mut ws.complex_buf);
675
676    ws.complex_buf.iter().map(|c| c.re).collect()
677}
678
679/// Compute RHS in-place using workspace buffers (f32)
680/// Uses per-direction gradient masks (mx, my, mz) matching MATLAB MEDI
681/// SIMD-accelerated for element-wise operations
682pub(crate) fn compute_rhs_inplace(
683    chi: &[f32],
684    w: &[Complex32],
685    b0: &[Complex32],
686    d_kernel: &[f32],
687    mx: &[f32],  // Per-direction gradient mask for x
688    my: &[f32],  // Per-direction gradient mask for y
689    mz: &[f32],  // Per-direction gradient mask for z
690    vr: &[f32],
691    lambda: f32,
692    rhs: &mut [f32],
693    ws: &mut MediWorkspace,
694) {
695    let n = ws.n_total;
696
697    // Regularization term: div(m * P * m * grad(chi)) for each direction
698    // MATLAB: b = lam .* gradAdj_(ux, uy, uz, vsz);
699    // where ux = mx .* P .* mx .* grad_x(chi), etc.
700    // Uses periodic BCs matching MATLAB's gradfp_mex / gradfp_adj_mex
701    fgrad_periodic_inplace_f32(
702        &mut ws.gx, &mut ws.gy, &mut ws.gz,
703        chi, ws.nx, ws.ny, ws.nz,
704        ws.vsx, ws.vsy, ws.vsz,
705    );
706
707    // Apply per-direction weights: ux = mx * P * mx * gx (SIMD accelerated)
708    apply_gradient_weights_f32(
709        &mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
710        mx, my, mz, vr,
711        &ws.gx, &ws.gy, &ws.gz,
712    );
713
714    bdiv_periodic_inplace_f32(
715        &mut ws.div_buf,
716        &ws.reg_x, &ws.reg_y, &ws.reg_z,
717        ws.nx, ws.ny, ws.nz,
718        ws.vsx, ws.vsy, ws.vsz,
719    );
720
721    // Data term: D^T(conj(w) * (-i) * (w - b0))
722    // MATLAB: b = b + real(ifft3(conj(D) .* fft3(1i .* w2 .* (exp(1i.*(f - Dx)) - 1))));
723    for i in 0..n {
724        let diff = w[i] - b0[i];
725        let conj_w = w[i].conj();
726        let neg_i = Complex32::new(0.0, -1.0);
727        ws.complex_buf2[i] = conj_w * neg_i * diff;
728    }
729
730    // Apply D^T (which is D for real symmetric kernel)
731    ws.fft_ws.fft3d(&mut ws.complex_buf2);
732
733    for i in 0..n {
734        ws.complex_buf2[i] *= d_kernel[i];
735    }
736
737    ws.fft_ws.ifft3d(&mut ws.complex_buf2);
738
739    // Extract real parts for SIMD combine operation
740    for i in 0..n {
741        ws.dipole_buf[i] = ws.complex_buf2[i].re;
742    }
743
744    // Combine terms: rhs = lambda * reg_term + data_term (SIMD accelerated, matching MATLAB)
745    combine_terms_f32(rhs, &ws.div_buf, &ws.dipole_buf, lambda);
746}
747
748/// Generate data weighting mask (f32)
749///
750/// # Arguments
751/// * `mode` - 0 for uniform weighting, 1 for SNR weighting
752/// * `n_std` - Noise standard deviation
753/// * `mask` - Binary mask
754pub(crate) fn dataterm_mask_f32(mode: i32, n_std: &[f32], mask: &[u8]) -> Vec<f32> {
755    let n = n_std.len();
756
757    if mode == 0 {
758        // Uniform weighting
759        mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect()
760    } else {
761        // SNR weighting: w = mask / N_std, normalized so mean over ROI = 1
762        let mut w: Vec<f32> = n_std.iter()
763            .zip(mask.iter())
764            .map(|(&n, &m)| {
765                if m != 0 && n > 1e-10 {
766                    1.0 / n
767                } else {
768                    0.0
769                }
770            })
771            .collect();
772
773        // Compute mean over ROI
774        let mask_count = mask.iter().filter(|&&m| m != 0).count() as f32;
775        if mask_count > 0.0 {
776            let sum: f32 = w.iter()
777                .zip(mask.iter())
778                .filter(|(_, &m)| m != 0)
779                .map(|(&wi, _)| wi)
780                .sum();
781            let mean = sum / mask_count;
782
783            if mean > 1e-10 {
784                // Normalize so mean = 1
785                for i in 0..n {
786                    w[i] /= mean;
787                }
788            }
789        }
790
791        // Ensure zeros outside mask
792        for i in 0..n {
793            if mask[i] == 0 {
794                w[i] = 0.0;
795            }
796        }
797
798        w
799    }
800}
801
802/// Generate per-direction gradient weighting masks (f32)
803///
804/// Computes separate edge masks for each gradient direction from magnitude image,
805/// matching the MATLAB MEDI implementation (gradientMaskMedi.m).
806/// Returns (mx, my, mz) where each mask is 1 (regularize) for non-edges, 0 for edges.
807///
808/// # Arguments
809/// * `magnitude` - Magnitude image
810/// * `mask` - Binary mask
811/// * `nx`, `ny`, `nz` - Array dimensions
812/// * `vsx`, `vsy`, `vsz` - Voxel sizes
813/// * `percentage` - Percentage of voxels considered to be edges (0.0-1.0, e.g., 0.3 = 30% edges)
814///
815/// # Returns
816/// Tuple of (mx, my, mz) per-direction binary gradient masks
817pub(crate) fn gradient_mask_f32(
818    magnitude: &[f32],
819    mask: &[u8],
820    nx: usize, ny: usize, nz: usize,
821    vsx: f32, vsy: f32, vsz: f32,
822    percentage: f32,
823) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
824    let n_total = nx * ny * nz;
825
826    // Normalize magnitude by max value within mask (matching MATLAB)
827    let mag_max = magnitude.iter()
828        .zip(mask.iter())
829        .filter(|(_, &m)| m != 0)
830        .map(|(&v, _)| v.abs())
831        .fold(0.0_f32, f32::max);
832
833    let mag_normalized: Vec<f32> = magnitude.iter()
834        .zip(mask.iter())
835        .map(|(&m, &msk)| {
836            if msk != 0 && mag_max > 1e-10 {
837                m / mag_max
838            } else {
839                0.0
840            }
841        })
842        .collect();
843
844    // Compute gradient of normalized magnitude (using linear extrapolation BCs)
845    let (gx, gy, gz) = fgrad_linext_f32(&mag_normalized, nx, ny, nz, vsx, vsy, vsz);
846
847    // Take absolute values of each gradient direction
848    let abs_gx: Vec<f32> = gx.iter().map(|&v| v.abs()).collect();
849    let abs_gy: Vec<f32> = gy.iter().map(|&v| v.abs()).collect();
850    let abs_gz: Vec<f32> = gz.iter().map(|&v| v.abs()).collect();
851
852    // Collect all gradient values within mask for threshold computation
853    let mut all_grads: Vec<f32> = Vec::with_capacity(3 * n_total);
854    for i in 0..n_total {
855        if mask[i] != 0 {
856            all_grads.push(abs_gx[i]);
857            all_grads.push(abs_gy[i]);
858            all_grads.push(abs_gz[i]);
859        }
860    }
861
862    if all_grads.is_empty() {
863        return (vec![1.0; n_total], vec![1.0; n_total], vec![1.0; n_total]);
864    }
865
866    // Sort to find percentile threshold (100 - percentage)
867    // MATLAB: thr = prctile([mx(mask); my(mask); mz(mask)], 100 - p);
868    all_grads.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
869    let percentile_idx = ((1.0 - percentage) * (all_grads.len() - 1) as f32) as usize;
870    let threshold = all_grads[percentile_idx.min(all_grads.len() - 1)];
871
872    // Create per-direction masks: 1 where gradient < threshold (non-edges), 0 at edges
873    // MATLAB: mx = mx < thr; my = my < thr; mz = mz < thr;
874    let mx: Vec<f32> = abs_gx.iter()
875        .zip(mask.iter())
876        .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
877        .collect();
878
879    let my: Vec<f32> = abs_gy.iter()
880        .zip(mask.iter())
881        .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
882        .collect();
883
884    let mz: Vec<f32> = abs_gz.iter()
885        .zip(mask.iter())
886        .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
887        .collect();
888
889    (mx, my, mz)
890}
891
892/// Forward difference gradient with linear extrapolation boundary conditions (f32)
893/// Matches MATLAB's gradf behavior: dx(end) = dx(end-1)
894pub(crate) fn fgrad_linext_f32(
895    x: &[f32],
896    nx: usize, ny: usize, nz: usize,
897    vsx: f32, vsy: f32, vsz: f32,
898) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
899    let n_total = nx * ny * nz;
900    let mut gx = vec![0.0f32; n_total];
901    let mut gy = vec![0.0f32; n_total];
902    let mut gz = vec![0.0f32; n_total];
903    fgrad_linext_inplace_f32(&mut gx, &mut gy, &mut gz, x, nx, ny, nz, vsx, vsy, vsz);
904    (gx, gy, gz)
905}
906
907/// Forward difference gradient with linear extrapolation boundary conditions (f32, in-place)
908/// Matches MATLAB's gradf behavior: dx(end) = dx(end-1)
909#[inline]
910pub(crate) fn fgrad_linext_inplace_f32(
911    gx: &mut [f32], gy: &mut [f32], gz: &mut [f32],
912    x: &[f32],
913    nx: usize, ny: usize, nz: usize,
914    vsx: f32, vsy: f32, vsz: f32,
915) {
916    let hx = 1.0 / vsx;
917    let hy = 1.0 / vsy;
918    let hz = 1.0 / vsz;
919
920    for k in 0..nz {
921        let k_offset = k * nx * ny;
922
923        for j in 0..ny {
924            let j_offset = j * nx;
925
926            for i in 0..nx {
927                let idx = i + j_offset + k_offset;
928                let x_val = x[idx];
929
930                // Forward difference with linear extrapolation at boundary
931                // MATLAB: dx(end,:,:) = dx(end-1,:,:)
932                if i + 1 < nx {
933                    gx[idx] = (x[idx + 1] - x_val) * hx;
934                } else if i > 0 {
935                    // Copy from previous (linear extrapolation)
936                    gx[idx] = gx[idx - 1];
937                } else {
938                    gx[idx] = 0.0;
939                }
940
941                if j + 1 < ny {
942                    gy[idx] = (x[i + (j + 1) * nx + k_offset] - x_val) * hy;
943                } else if j > 0 {
944                    gy[idx] = gy[i + (j - 1) * nx + k_offset];
945                } else {
946                    gy[idx] = 0.0;
947                }
948
949                if k + 1 < nz {
950                    gz[idx] = (x[i + j_offset + (k + 1) * nx * ny] - x_val) * hz;
951                } else if k > 0 {
952                    gz[idx] = gz[i + j_offset + (k - 1) * nx * ny];
953                } else {
954                    gz[idx] = 0.0;
955                }
956            }
957        }
958    }
959}
960
961
962/// Forward difference gradient with periodic boundary conditions (f32, in-place)
963/// Matches MATLAB's gradfp_mex used inside MEDI iterations.
964/// At boundaries, wraps around: dx(end) = (x(1) - x(end)) / h
965///
966/// Parallelised one rayon task per z-slab (disjoint output chunks); `x` is only
967/// read, including across slab boundaries for the z-difference.
968pub(crate) fn fgrad_periodic_inplace_f32(
969    gx: &mut [f32], gy: &mut [f32], gz: &mut [f32],
970    x: &[f32],
971    nx: usize, ny: usize, nz: usize,
972    vsx: f32, vsy: f32, vsz: f32,
973) {
974    let hx = 1.0 / vsx;
975    let hy = 1.0 / vsy;
976    let hz = 1.0 / vsz;
977    let nxny = nx * ny;
978
979    maybe_par_chunks_mut!(gx, nxny)
980        .zip(maybe_par_chunks_mut!(gy, nxny))
981        .zip(maybe_par_chunks_mut!(gz, nxny))
982        .enumerate()
983        .for_each(|(k, ((gx_s, gy_s), gz_s))| {
984            let k_offset = k * nxny;
985
986            for j in 0..ny {
987                let j_offset = j * nx;
988
989                for i in 0..nx {
990                    let local = i + j_offset;
991                    let idx = local + k_offset;
992                    let x_val = x[idx];
993
994                    // x-direction: periodic wrap at i = nx-1
995                    let x_next = if i + 1 < nx { x[idx + 1] } else { x[j_offset + k_offset] };
996                    gx_s[local] = (x_next - x_val) * hx;
997
998                    // y-direction: periodic wrap at j = ny-1
999                    let y_next = if j + 1 < ny { x[i + (j + 1) * nx + k_offset] } else { x[i + k_offset] };
1000                    gy_s[local] = (y_next - x_val) * hy;
1001
1002                    // z-direction: periodic wrap at k = nz-1
1003                    let z_next = if k + 1 < nz { x[i + j_offset + (k + 1) * nxny] } else { x[i + j_offset] };
1004                    gz_s[local] = (z_next - x_val) * hz;
1005                }
1006            }
1007        });
1008}
1009
1010/// Backward divergence with periodic boundary conditions (f32, in-place)
1011/// Adjoint of fgrad_periodic_inplace_f32, matching MATLAB's gradfp_adj_mex.
1012/// At boundaries, wraps around: at i=0, uses gx(end) instead of zero.
1013///
1014/// Parallelised one rayon task per z-slab (disjoint output chunks); the
1015/// gradient inputs are only read, including across slab boundaries.
1016pub(crate) fn bdiv_periodic_inplace_f32(
1017    div: &mut [f32],
1018    gx: &[f32], gy: &[f32], gz: &[f32],
1019    nx: usize, ny: usize, nz: usize,
1020    vsx: f32, vsy: f32, vsz: f32,
1021) {
1022    let hx = -1.0 / vsx;  // Negative for adjoint
1023    let hy = -1.0 / vsy;
1024    let hz = -1.0 / vsz;
1025    let nxny = nx * ny;
1026
1027    maybe_par_chunks_mut!(div, nxny).enumerate().for_each(|(k, div_s)| {
1028        let k_offset = k * nxny;
1029
1030        for j in 0..ny {
1031            let j_offset = j * nx;
1032
1033            for i in 0..nx {
1034                let local = i + j_offset;
1035                let idx = local + k_offset;
1036
1037                // x-direction: at i=0, wrap to gx[nx-1,j,k]
1038                let gx_prev = if i > 0 { gx[idx - 1] } else { gx[(nx - 1) + j_offset + k_offset] };
1039                let gx_term = (gx[idx] - gx_prev) * hx;
1040
1041                // y-direction: at j=0, wrap to gy[i,ny-1,k]
1042                let gy_prev = if j > 0 { gy[i + (j - 1) * nx + k_offset] } else { gy[i + (ny - 1) * nx + k_offset] };
1043                let gy_term = (gy[idx] - gy_prev) * hy;
1044
1045                // z-direction: at k=0, wrap to gz[i,j,nz-1]
1046                let gz_prev = if k > 0 { gz[i + j_offset + (k - 1) * nxny] } else { gz[i + j_offset + (nz - 1) * nxny] };
1047                let gz_term = (gz[idx] - gz_prev) * hz;
1048
1049                div_s[local] = gx_term + gy_term + gz_term;
1050            }
1051        }
1052    });
1053}
1054
1055
1056#[cfg(test)]
1057mod tests {
1058    use super::*;
1059
1060    #[test]
1061    fn test_dataterm_mask_uniform() {
1062        let n_std = vec![1.0f32; 27];
1063        let mask = vec![1u8; 27];
1064
1065        let w = dataterm_mask_f32(0, &n_std, &mask);
1066
1067        for &wi in w.iter() {
1068            assert!((wi - 1.0).abs() < 1e-10);
1069        }
1070    }
1071
1072    #[test]
1073    fn test_dataterm_mask_snr() {
1074        let n_std = vec![2.0f32; 27];
1075        let mask = vec![1u8; 27];
1076
1077        let w = dataterm_mask_f32(1, &n_std, &mask);
1078
1079        // Mean should be 1
1080        let mean: f32 = w.iter().sum::<f32>() / 27.0;
1081        assert!((mean - 1.0).abs() < 1e-5);
1082    }
1083
1084    #[test]
1085    fn test_gradient_mask_constant() {
1086        // Constant magnitude should have no edges (all gradients are zero)
1087        let mag = vec![1.0f32; 8 * 8 * 8];
1088        let mask = vec![1u8; 8 * 8 * 8];
1089
1090        let (mx, my, mz) = gradient_mask_f32(&mag, &mask, 8, 8, 8, 1.0, 1.0, 1.0, 0.3);
1091
1092        // All should be binary masks (0 or 1)
1093        for i in 0..(8 * 8 * 8) {
1094            assert!(mx[i] == 0.0 || mx[i] == 1.0, "mx should be binary, got {}", mx[i]);
1095            assert!(my[i] == 0.0 || my[i] == 1.0, "my should be binary, got {}", my[i]);
1096            assert!(mz[i] == 0.0 || mz[i] == 1.0, "mz should be binary, got {}", mz[i]);
1097        }
1098    }
1099
1100    fn test_medi_params() -> MediParams {
1101        MediParams {
1102            lambda: 1000.0, percentage: 0.9, cg_tol: 0.1,
1103            cg_max_iter: 10, max_iter: 3, tol: 0.1,
1104            ..MediParams::default()
1105        }
1106    }
1107
1108    #[test]
1109    fn test_medi_zero_field() {
1110        let n = 8;
1111        let field = vec![0.0; n * n * n];
1112        let mask = vec![1u8; n * n * n];
1113        let mag = vec![1.0; n * n * n];
1114        let n_std = vec![1.0; n * n * n];
1115        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1116
1117        let chi = medi(
1118            &field, &n_std, &mag, &mask, &grid,
1119            (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1120        );
1121
1122        for &val in chi.iter() {
1123            assert!(val.abs() < 1e-4, "Zero field should give near-zero chi, got {}", val);
1124        }
1125    }
1126
1127    #[test]
1128    fn test_medi_finite() {
1129        let n = 8;
1130        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1131        let mask = vec![1u8; n * n * n];
1132        let mag = vec![1.0; n * n * n];
1133        let n_std = vec![1.0; n * n * n];
1134        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1135
1136        let chi = medi(
1137            &field, &n_std, &mag, &mask, &grid,
1138            (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1139        );
1140
1141        for (i, &val) in chi.iter().enumerate() {
1142            assert!(val.is_finite(), "Chi should be finite at index {}", i);
1143        }
1144    }
1145
1146    #[test]
1147    fn test_medi_with_smv() {
1148        let n = 8;
1149        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1150        let mask = vec![1u8; n * n * n];
1151        let mag = vec![1.0; n * n * n];
1152        let n_std = vec![1.0; n * n * n];
1153        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1154
1155        // Test with SMV enabled
1156        let params = MediParams { smv: true, smv_radius: 2.0, ..test_medi_params() };
1157        let chi = medi(
1158            &field, &n_std, &mag, &mask, &grid,
1159            (0.0, 0.0, 1.0), &params, |_, _| {},
1160        );
1161
1162        for (i, &val) in chi.iter().enumerate() {
1163            assert!(val.is_finite(), "Chi with SMV should be finite at index {}", i);
1164        }
1165    }
1166
1167    #[test]
1168    fn test_medi_mask() {
1169        let n = 8;
1170        let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1171        let mut mask = vec![1u8; n * n * n];
1172        let mag = vec![1.0; n * n * n];
1173        let n_std = vec![1.0; n * n * n];
1174        mask[0] = 0;
1175        mask[10] = 0;
1176        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1177
1178        let chi = medi(
1179            &field, &n_std, &mag, &mask, &grid,
1180            (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1181        );
1182
1183        assert_eq!(chi[0], 0.0, "Masked voxel should be zero");
1184        assert_eq!(chi[10], 0.0, "Masked voxel should be zero");
1185    }
1186
1187    /// Debug test: run MEDI step-by-step on real data, saving intermediates
1188    /// for comparison with Octave reference (other/medi_debug_octave.m).
1189    /// Run with: cargo test --release test_medi_debug -- --ignored --nocapture
1190    #[test]
1191    #[ignore]
1192    fn test_medi_debug() {
1193        let data_path = "/home/ashley/OUT/2bgRemoved.nii";
1194        if !std::path::Path::new(data_path).exists() {
1195            eprintln!("Skipping: {} not found", data_path);
1196            return;
1197        }
1198
1199        let outdir = "/home/ashley/OUT/debug";
1200        std::fs::create_dir_all(outdir).ok();
1201
1202        // Read NIfTI
1203        let bytes = std::fs::read(data_path).unwrap();
1204        let nifti_data = crate::io::load_nifti(&bytes).unwrap();
1205        let (nx, ny, nz) = nifti_data.dims;
1206        let (vsx, vsy, vsz) = nifti_data.voxel_size;
1207
1208        let n_total = nx * ny * nz;
1209        let vsx_f32 = vsx as f32;
1210        let vsy_f32 = vsy as f32;
1211        let vsz_f32 = vsz as f32;
1212
1213        eprintln!("Data: {}x{}x{}, voxel: {}x{}x{}", nx, ny, nz, vsx, vsy, vsz);
1214
1215        // Convert to f32
1216        let local_field: Vec<f32> = nifti_data.data.iter().map(|&v| v as f32).collect();
1217
1218        // Create mask from non-zero voxels
1219        let mask: Vec<u8> = local_field.iter()
1220            .map(|&v| if v.abs() > 1e-10 { 1 } else { 0 })
1221            .collect();
1222        let mask_count: usize = mask.iter().filter(|&&m| m != 0).count();
1223        eprintln!("Mask voxels: {} / {}", mask_count, n_total);
1224
1225        // Save inputs
1226        save_f32_raw(&local_field, &format!("{}/f_rust.raw", outdir));
1227        let mask_f32: Vec<f32> = mask.iter().map(|&m| m as f32).collect();
1228        save_f32_raw(&mask_f32, &format!("{}/mask_rust.raw", outdir));
1229
1230        // Parameters (matching MATLAB defaults)
1231        let lambda: f32 = 7.5e-5;
1232        let beta: f32 = 1.49e-8;
1233        let bdir = (0.0f32, 0.0f32, 1.0f32);
1234        let cg_tol: f32 = 0.1;  // Match MATLAB tolcg default
1235        let cg_max_iter: usize = 10;
1236
1237        // Workspace
1238        let debug_grid = Grid::new(nx, ny, nz, vsx, vsy, vsz);
1239
1240        // Dipole kernel
1241        let d_kernel = crate::kernels::dipole::dipole_kernel_f32(
1242            &debug_grid, bdir,
1243        );
1244        save_f32_raw(&d_kernel, &format!("{}/D_rust.raw", outdir));
1245        eprintln!("D: min={} max={} D[0]={}", fmin(&d_kernel), fmax(&d_kernel), d_kernel[0]);
1246        let mut ws = MediWorkspace::new(&debug_grid);
1247
1248        // Data weighting: uniform m = mask (matching w=ones in Octave)
1249        let m: Vec<f32> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
1250
1251        // b0 = m * exp(i * f)
1252        let b0: Vec<Complex32> = local_field.iter().zip(m.iter())
1253            .map(|(&f, &mi)| {
1254                let phase = Complex32::new(0.0, f);
1255                mi * phase.exp()
1256            })
1257            .collect();
1258
1259        // Gradient mask: uniform magnitude -> mx=my=mz=mask
1260        let w_gx: Vec<f32> = m.clone();
1261        let w_gy: Vec<f32> = m.clone();
1262        let w_gz: Vec<f32> = m.clone();
1263
1264        // ===== ITERATION 1 (chi = 0) =====
1265        eprintln!("\n=== Iteration 1 ===");
1266        let mut chi = vec![0.0f32; n_total];
1267        let mut dx = vec![0.0f32; n_total];
1268        let mut rhs = vec![0.0f32; n_total];
1269        let mut vr = vec![0.0f32; n_total];
1270        let mut w: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); n_total];
1271
1272        // P weights (chi=0 -> gradient=0 -> P = 1/sqrt(beta))
1273        fgrad_periodic_inplace_f32(
1274            &mut ws.gx, &mut ws.gy, &mut ws.gz,
1275            &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
1276        );
1277        compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
1278        save_f32_raw(&vr, &format!("{}/P1_rust.raw", outdir));
1279        eprintln!("P1: min={} max={} mean={}", fmin(&vr), fmax(&vr), fmean(&vr));
1280
1281        // w = m * exp(i * D*chi) = m (since chi=0)
1282        apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
1283        for i in 0..n_total {
1284            let phase = Complex32::new(0.0, ws.dipole_buf[i]);
1285            w[i] = m[i] * phase.exp();
1286        }
1287
1288        // RHS
1289        compute_rhs_inplace(
1290            &chi, &w, &b0, &d_kernel,
1291            &w_gx, &w_gy, &w_gz, &vr, lambda,
1292            &mut rhs, &mut ws,
1293        );
1294        save_f32_raw(&rhs, &format!("{}/rhs1_rust.raw", outdir));
1295        eprintln!("RHS1: min={} max={} norm={}", fmin(&rhs), fmax(&rhs), fnorm(&rhs));
1296
1297        // Negate for CG
1298        negate_f32(&mut rhs);
1299
1300        // CG solve (with iteration-level residual logging)
1301        let mut cg_residuals: Vec<f32> = Vec::new();
1302        {
1303            // Manual CG to capture residuals (matching cg_solve_medi but with logging)
1304            let n = ws.n_total;
1305            let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
1306            let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
1307
1308            dx.fill(0.0);
1309            ws.cg_r.copy_from_slice(&rhs);
1310            ws.cg_p.copy_from_slice(&ws.cg_r);
1311            let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
1312            let b_norm: f32 = norm_squared_f32(&rhs).sqrt();
1313
1314            let mut p_copy = vec![0.0f32; n];
1315            let mut prev_residual = rsold.sqrt();
1316
1317            for cg_iter in 0..cg_max_iter {
1318                let residual_before = rsold.sqrt();
1319                cg_residuals.push(residual_before);
1320
1321                p_copy.copy_from_slice(&ws.cg_p);
1322                {
1323                    let mut bufs = MediOpBuffers {
1324                        gx: &mut ws.gx, gy: &mut ws.gy, gz: &mut ws.gz,
1325                        reg_x: &mut ws.reg_x, reg_y: &mut ws.reg_y, reg_z: &mut ws.reg_z,
1326                        div_buf: &mut ws.div_buf, dipole_buf: &mut ws.dipole_buf,
1327                        complex_buf: &mut ws.complex_buf, complex_buf2: &mut ws.complex_buf2,
1328                    };
1329                    apply_medi_operator_core(
1330                        &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
1331                        &p_copy, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda, &mut ws.cg_ap,
1332                    );
1333                }
1334
1335                let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
1336                if pap.abs() < 1e-15 { break; }
1337                let alpha = rsold / pap;
1338
1339                axpy_f32(&mut dx, alpha, &ws.cg_p);
1340                axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
1341
1342                let rsnew: f32 = norm_squared_f32(&ws.cg_r);
1343                let residual = rsnew.sqrt();
1344
1345                eprintln!("  CG iter {}: res={:.6e}, alpha={:.6e}, pap={:.6e}",
1346                    cg_iter + 1, residual, alpha, pap);
1347
1348                if residual < cg_tol * b_norm { break; }
1349
1350                // No stall detection in this debug version (matching MATLAB)
1351
1352                let beta_cg = rsnew / rsold;
1353                xpby_f32(&mut ws.cg_p, &ws.cg_r, beta_cg);
1354                rsold = rsnew;
1355                prev_residual = residual;
1356            }
1357        }
1358        save_f32_raw(&dx, &format!("{}/dx1_rust.raw", outdir));
1359        eprintln!("dx1: min={} max={} norm={}", fmin(&dx), fmax(&dx), fnorm(&dx));
1360
1361        // Update chi
1362        axpy_f32(&mut chi, 1.0, &dx);
1363        save_f32_raw(&chi, &format!("{}/chi1_rust.raw", outdir));
1364        eprintln!("chi1: min={} max={} norm={}", fmin(&chi), fmax(&chi), fnorm(&chi));
1365
1366        // ===== ITERATION 2 =====
1367        eprintln!("\n=== Iteration 2 ===");
1368
1369        // P weights
1370        fgrad_periodic_inplace_f32(
1371            &mut ws.gx, &mut ws.gy, &mut ws.gz,
1372            &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
1373        );
1374        compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
1375        save_f32_raw(&vr, &format!("{}/P2_rust.raw", outdir));
1376        eprintln!("P2: min={} max={} mean={}", fmin(&vr), fmax(&vr), fmean(&vr));
1377
1378        // w = m * exp(i * D*chi)
1379        apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
1380        for i in 0..n_total {
1381            let phase = Complex32::new(0.0, ws.dipole_buf[i]);
1382            w[i] = m[i] * phase.exp();
1383        }
1384
1385        // RHS
1386        compute_rhs_inplace(
1387            &chi, &w, &b0, &d_kernel,
1388            &w_gx, &w_gy, &w_gz, &vr, lambda,
1389            &mut rhs, &mut ws,
1390        );
1391        save_f32_raw(&rhs, &format!("{}/rhs2_rust.raw", outdir));
1392        eprintln!("RHS2: min={} max={} norm={}", fmin(&rhs), fmax(&rhs), fnorm(&rhs));
1393
1394        negate_f32(&mut rhs);
1395
1396        // CG solve (no stall detection)
1397        {
1398            let n = ws.n_total;
1399            let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
1400            let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
1401
1402            dx.fill(0.0);
1403            ws.cg_r.copy_from_slice(&rhs);
1404            ws.cg_p.copy_from_slice(&ws.cg_r);
1405            let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
1406            let b_norm: f32 = norm_squared_f32(&rhs).sqrt();
1407            let mut p_copy = vec![0.0f32; n];
1408
1409            for cg_iter in 0..cg_max_iter {
1410                p_copy.copy_from_slice(&ws.cg_p);
1411                {
1412                    let mut bufs = MediOpBuffers {
1413                        gx: &mut ws.gx, gy: &mut ws.gy, gz: &mut ws.gz,
1414                        reg_x: &mut ws.reg_x, reg_y: &mut ws.reg_y, reg_z: &mut ws.reg_z,
1415                        div_buf: &mut ws.div_buf, dipole_buf: &mut ws.dipole_buf,
1416                        complex_buf: &mut ws.complex_buf, complex_buf2: &mut ws.complex_buf2,
1417                    };
1418                    apply_medi_operator_core(
1419                        &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
1420                        &p_copy, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda, &mut ws.cg_ap,
1421                    );
1422                }
1423                let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
1424                if pap.abs() < 1e-15 { break; }
1425                let alpha = rsold / pap;
1426                axpy_f32(&mut dx, alpha, &ws.cg_p);
1427                axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
1428                let rsnew: f32 = norm_squared_f32(&ws.cg_r);
1429                let residual = rsnew.sqrt();
1430                eprintln!("  CG iter {}: res={:.6e}", cg_iter + 1, residual);
1431                if residual < cg_tol * b_norm { break; }
1432                let beta_cg = rsnew / rsold;
1433                xpby_f32(&mut ws.cg_p, &ws.cg_r, beta_cg);
1434                rsold = rsnew;
1435            }
1436        }
1437        save_f32_raw(&dx, &format!("{}/dx2_rust.raw", outdir));
1438        eprintln!("dx2: min={} max={} norm={}", fmin(&dx), fmax(&dx), fnorm(&dx));
1439
1440        axpy_f32(&mut chi, 1.0, &dx);
1441        save_f32_raw(&chi, &format!("{}/chi2_rust.raw", outdir));
1442        eprintln!("chi2: min={} max={} norm={}", fmin(&chi), fmax(&chi), fnorm(&chi));
1443
1444        eprintln!("\nDone. Intermediates saved to {}", outdir);
1445    }
1446
1447    fn save_f32_raw(data: &[f32], path: &str) {
1448        use std::io::Write;
1449        let mut file = std::fs::File::create(path).unwrap();
1450        for &val in data {
1451            file.write_all(&val.to_le_bytes()).unwrap();
1452        }
1453    }
1454
1455    fn fmin(data: &[f32]) -> f32 { data.iter().cloned().fold(f32::MAX, f32::min) }
1456    fn fmax(data: &[f32]) -> f32 { data.iter().cloned().fold(f32::MIN, f32::max) }
1457    fn fmean(data: &[f32]) -> f32 { data.iter().sum::<f32>() / data.len() as f32 }
1458    fn fnorm(data: &[f32]) -> f32 { data.iter().map(|&v| v * v).sum::<f32>().sqrt() }
1459
1460    #[test]
1461    fn test_medi_small() {
1462        // 8x8x8 volume with synthetic local field
1463        let n = 8;
1464        let n_total = n * n * n;
1465
1466        // Create a synthetic local field from a dipole-like source
1467        let mut field = vec![0.0f64; n_total];
1468        let center = n / 2;
1469        for z in 0..n {
1470            for y in 0..n {
1471                for x in 0..n {
1472                    let idx = x + y * n + z * n * n;
1473                    let dx = (x as f64) - (center as f64);
1474                    let dy = (y as f64) - (center as f64);
1475                    let dz = (z as f64) - (center as f64);
1476                    let r2 = dx*dx + dy*dy + dz*dz;
1477                    if r2 > 1.0 {
1478                        // Dipole field pattern: (3*cos^2(theta) - 1) / r^3
1479                        let r = r2.sqrt();
1480                        let cos_theta = dz / r;
1481                        field[idx] = (3.0 * cos_theta * cos_theta - 1.0) / (r * r * r) * 0.01;
1482                    }
1483                }
1484            }
1485        }
1486
1487        let mask = vec![1u8; n_total];
1488        let mag = vec![1.0f64; n_total];
1489        let n_std = vec![1.0f64; n_total];
1490
1491        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1492
1493        // Run MEDI with few Gauss-Newton iterations to exercise the main loop
1494        let params = MediParams {
1495            lambda: 1e-4, percentage: 0.3,
1496            cg_tol: 0.01, cg_max_iter: 10, max_iter: 5, tol: 0.1,
1497            ..MediParams::default()
1498        };
1499        let chi = medi(
1500            &field, &n_std, &mag, &mask, &grid,
1501            (0.0, 0.0, 1.0), &params, |_, _| {},
1502        );
1503
1504        // Result should be finite and same size
1505        assert_eq!(chi.len(), n_total);
1506        for (i, &val) in chi.iter().enumerate() {
1507            assert!(val.is_finite(), "MEDI L1 chi should be finite at index {}", i);
1508        }
1509
1510        // Chi should be non-trivial (not all zero) for a dipole field input
1511        let chi_norm: f64 = chi.iter().map(|&v| v * v).sum::<f64>().sqrt();
1512        assert!(chi_norm > 1e-10, "MEDI L1 should produce non-zero susceptibility for dipole field, got norm={}", chi_norm);
1513    }
1514
1515    #[test]
1516    fn test_medi_weight_types() {
1517        let n = 8;
1518        let n_total = n * n * n;
1519
1520        let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1521        let mask = vec![1u8; n_total];
1522        let mag = vec![1.0f64; n_total];
1523        let n_std = vec![1.0f64; n_total];
1524
1525        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1526
1527        // Test uniform data weighting (mode 0)
1528        let params_uniform = MediParams { data_weighting: 0, ..test_medi_params() };
1529        let chi_uniform = medi(
1530            &field, &n_std, &mag, &mask, &grid,
1531            (0.0, 0.0, 1.0), &params_uniform, |_, _| {},
1532        );
1533
1534        // Test SNR data weighting (mode 1) with varying noise
1535        let n_std_varying: Vec<f64> = (0..n_total).map(|i| 0.5 + (i as f64) * 0.01).collect();
1536        let params_snr = MediParams { data_weighting: 1, ..test_medi_params() };
1537        let chi_snr = medi(
1538            &field, &n_std_varying, &mag, &mask, &grid,
1539            (0.0, 0.0, 1.0), &params_snr, |_, _| {},
1540        );
1541
1542        // Both should be finite
1543        for (i, &val) in chi_uniform.iter().enumerate() {
1544            assert!(val.is_finite(), "Uniform weighting chi should be finite at {}", i);
1545        }
1546        for (i, &val) in chi_snr.iter().enumerate() {
1547            assert!(val.is_finite(), "SNR weighting chi should be finite at {}", i);
1548        }
1549
1550        // Uniform and SNR weighting should give different results
1551        let diff_norm: f64 = chi_uniform.iter()
1552            .zip(chi_snr.iter())
1553            .map(|(&a, &b)| (a - b).powi(2))
1554            .sum::<f64>()
1555            .sqrt();
1556        // They may or may not differ depending on input, so just check finiteness
1557        assert!(diff_norm.is_finite(), "Difference between weight modes should be finite");
1558    }
1559
1560    #[test]
1561    fn test_medi_with_progress_small() {
1562        let n = 8;
1563        let n_total = n * n * n;
1564
1565        let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1566        let mask = vec![1u8; n_total];
1567        let mag = vec![1.0f64; n_total];
1568        let n_std = vec![1.0f64; n_total];
1569        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1570
1571        let mut progress_calls = 0usize;
1572        let chi = medi(
1573            &field, &n_std, &mag, &mask, &grid,
1574            (0.0, 0.0, 1.0), &test_medi_params(),
1575            |_iter, _max| { progress_calls += 1; },
1576        );
1577
1578        assert_eq!(chi.len(), n_total);
1579        for &val in &chi {
1580            assert!(val.is_finite(), "medi with progress output should be finite");
1581        }
1582        assert!(progress_calls > 0, "progress callback should be called at least once");
1583    }
1584
1585    #[test]
1586    fn test_medi_with_merit() {
1587        let n = 8;
1588        let n_total = n * n * n;
1589
1590        let mut field = vec![0.0f64; n_total];
1591        let center = n / 2;
1592        for z in 0..n {
1593            for y in 0..n {
1594                for x in 0..n {
1595                    let idx = x + y * n + z * n * n;
1596                    let dx = (x as f64) - (center as f64);
1597                    let dy = (y as f64) - (center as f64);
1598                    let dz = (z as f64) - (center as f64);
1599                    let r2 = dx * dx + dy * dy + dz * dz;
1600                    if r2 > 1.0 {
1601                        let r = r2.sqrt();
1602                        let cos_theta = dz / r;
1603                        field[idx] = (3.0 * cos_theta * cos_theta - 1.0) / (r * r * r) * 0.01;
1604                    }
1605                }
1606            }
1607        }
1608
1609        let mask = vec![1u8; n_total];
1610        let mag = vec![1.0f64; n_total];
1611        let n_std = vec![1.0f64; n_total];
1612
1613        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1614        // Test with merit=true to cover the merit adjustment code path
1615        let params = MediParams {
1616            lambda: 1e-4, merit: true, percentage: 0.3,
1617            cg_tol: 0.01, cg_max_iter: 10, max_iter: 5, tol: 0.1,
1618            ..MediParams::default()
1619        };
1620        let chi = medi(
1621            &field, &n_std, &mag, &mask, &grid,
1622            (0.0, 0.0, 1.0), &params, |_, _| {},
1623        );
1624
1625        assert_eq!(chi.len(), n_total);
1626        for &val in &chi {
1627            assert!(val.is_finite(), "MEDI with merit should produce finite results");
1628        }
1629    }
1630
1631    #[test]
1632    fn test_medi_with_smv_final() {
1633        let n = 8;
1634        let n_total = n * n * n;
1635
1636        let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1637        let mask = vec![1u8; n_total];
1638        let mag = vec![1.0f64; n_total];
1639        let n_std = vec![1.0f64; n_total];
1640        let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1641
1642        // Test with smv=true to cover the SMV-weighted code path
1643        let params = MediParams {
1644            smv: true, smv_radius: 3.0, percentage: 0.3,
1645            cg_tol: 0.01, cg_max_iter: 10, max_iter: 3, tol: 0.1,
1646            ..MediParams::default()
1647        };
1648        let chi = medi(
1649            &field, &n_std, &mag, &mask, &grid,
1650            (0.0, 0.0, 1.0), &params, |_, _| {},
1651        );
1652
1653        assert_eq!(chi.len(), n_total);
1654        for &val in &chi {
1655            assert!(val.is_finite(), "MEDI with SMV should produce finite results");
1656        }
1657    }
1658}