Skip to main content

qsm_core/separation/
chi_sep_medi.rs

1//! Chi-separation using MEDI-based coupled optimization
2//!
3//! Separates total susceptibility into paramagnetic (chi+) and diamagnetic (chi-)
4//! components using chi_pos + chi_neg formulation:
5//!   chi_pos >= 0 (paramagnetic, iron), in Hz internally
6//!   chi_neg <= 0 (diamagnetic, myelin), in Hz internally
7//!   chi_total = chi_pos + chi_neg
8//!
9//! Forward model (all in Hz):
10//!   field = D * (chi_pos + chi_neg)
11//!   R2'(Hz) = dr_p_eff * chi_pos + dr_q_eff * (-chi_neg)
12//!           = dr_p_eff * |chi_pos| + dr_q_eff * |chi_neg|
13//!   where dr_eff = ppm_factor * Dr (dimensionless effective relaxivity)
14//!
15//! The constraints (chi_pos >= 0, chi_neg <= 0) naturally break the gauge
16//! freedom of the chi_pos + chi_neg formulation. In most voxels, either
17//! chi_pos = 0 or chi_neg = 0, which pins one variable to its constraint
18//! boundary and prevents correlated drift.
19//!
20//! Reference:
21//! Shin, H., et al. (2021). "chi-separation: Magnetic susceptibility source
22//! separation toward iron and myelin mapping in the brain." NeuroImage, 240:118371.
23
24use num_complex::Complex32;
25use crate::Grid;
26use crate::fft::Fft3dWorkspaceF32;
27use crate::kernels::dipole::dipole_kernel_f32;
28use crate::inversion::medi::{
29    gradient_mask_f32,
30    fgrad_periodic_inplace_f32,
31    bdiv_periodic_inplace_f32,
32};
33use crate::utils::padding::{next_fast_fft_size, pad3d, unpad3d};
34use crate::utils::simd_ops::{
35    dot_product_f32, norm_squared_f32, axpy_f32, xpby_f32,
36    apply_gradient_weights_f32, compute_p_weights_f32,
37};
38
39/// Workspace for chi-separation — holds all reusable buffers (f32).
40struct ChiSepWorkspace {
41    n: usize,
42    nx: usize, ny: usize, nz: usize,
43    vsx: f32, vsy: f32, vsz: f32,
44
45    fft_ws: Fft3dWorkspaceF32,
46
47    gx: Vec<f32>,
48    gy: Vec<f32>,
49    gz: Vec<f32>,
50
51    reg_x: Vec<f32>,
52    reg_y: Vec<f32>,
53    reg_z: Vec<f32>,
54
55    div_buf: Vec<f32>,
56
57    complex_buf: Vec<Complex32>,
58    dipole_buf: Vec<f32>,
59}
60
61impl ChiSepWorkspace {
62    fn new(nx: usize, ny: usize, nz: usize, vsx: f32, vsy: f32, vsz: f32) -> Self {
63        let n = nx * ny * nz;
64        Self {
65            n, nx, ny, nz, vsx, vsy, vsz,
66            fft_ws: Fft3dWorkspaceF32::new(nx, ny, nz),
67            gx: vec![0.0; n],
68            gy: vec![0.0; n],
69            gz: vec![0.0; n],
70            reg_x: vec![0.0; n],
71            reg_y: vec![0.0; n],
72            reg_z: vec![0.0; n],
73            div_buf: vec![0.0; n],
74            complex_buf: vec![Complex32::new(0.0, 0.0); n],
75            dipole_buf: vec![0.0; n],
76        }
77    }
78}
79
80/// Chi-separation algorithm parameters.
81///
82/// `cf`, `dr_pos`, and `dr_neg` are acquisition/tissue dependent and should be
83/// set for your scan; the remaining fields are optimization knobs with sensible
84/// defaults.
85#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
86#[derive(Clone, Debug)]
87pub struct ChiSepParams {
88    /// Central frequency in Hz (e.g. 123.2e6 for 3T)
89    pub cf: f64,
90    /// Paramagnetic L1 regularization weight
91    pub lambda_para: f64,
92    /// Diamagnetic L1 regularization weight
93    pub lambda_dia: f64,
94    /// Field/R2' coupling weight
95    pub lambda_cpl: f64,
96    /// Paramagnetic relaxivity in Hz/ppm
97    pub dr_pos: f64,
98    /// Diamagnetic relaxivity in Hz/ppm
99    pub dr_neg: f64,
100    /// Edge-mask percentage for morphology weighting
101    pub percentage: f64,
102    /// Inner conjugate-gradient tolerance
103    pub cg_tol: f64,
104    /// Inner conjugate-gradient max iterations
105    pub cg_max_iter: usize,
106    /// Outer Gauss-Newton max iterations
107    pub max_iter: usize,
108    /// Outer convergence tolerance
109    pub tol: f64,
110}
111
112impl Default for ChiSepParams {
113    fn default() -> Self {
114        Self {
115            cf: 123.2e6,
116            lambda_para: 1000.0,
117            lambda_dia: 1000.0,
118            lambda_cpl: 100.0,
119            dr_pos: 114.0,
120            dr_neg: 30.0,
121            percentage: 0.3,
122            cg_tol: 0.01,
123            cg_max_iter: 100,
124            max_iter: 10,
125            tol: 0.1,
126        }
127    }
128}
129
130/// Chi-separation using MEDI-based coupled optimization.
131///
132/// # Arguments
133/// * `local_field` - Local field map in ppm (same units convention as the
134///   dipole-inversion algorithms; converted to Hz internally via `params.cf`)
135/// * `r2prime` - R2' map in Hz
136/// * `magnitude` - Magnitude image for edge weighting
137/// * `mask` - Binary brain mask, 1 = brain
138/// * `grid` - Volume grid (dimensions and voxel sizes)
139/// * `bdir` - B0 field direction
140/// * `params` - Chi-separation parameters (see [`ChiSepParams`])
141/// * `progress` - Progress callback: `(iteration, total_iterations)`
142///
143/// # Returns
144/// `(chi_pos, chi_neg, chi_total)` — susceptibility maps in ppm
145///
146/// Volumes whose dimensions are not FFT-friendly (2ᵃ·3ᵇ·5ᶜ) are transparently
147/// zero-padded to the next fast size for the internal FFTs and cropped back
148/// (see [`chi_sep_ilsqr`](crate::separation::chi_sep_ilsqr) for details).
149pub fn chi_sep_medi<F>(
150    local_field: &[f64],
151    r2prime: &[f64],
152    magnitude: &[f64],
153    mask: &[u8],
154    grid: &Grid,
155    bdir: (f64, f64, f64),
156    params: &ChiSepParams,
157    progress: F,
158) -> (Vec<f64>, Vec<f64>, Vec<f64>)
159where
160    F: FnMut(usize, usize),
161{
162    let dims = grid.dims;
163    let fast = (
164        next_fast_fft_size(dims.0),
165        next_fast_fft_size(dims.1),
166        next_fast_fft_size(dims.2),
167    );
168    if fast == dims {
169        return chi_sep_medi_core(local_field, r2prime, magnitude, mask, grid, bdir, params, progress);
170    }
171    let (vsx, vsy, vsz) = grid.voxel_size;
172    let pgrid = Grid::new(fast.0, fast.1, fast.2, vsx, vsy, vsz);
173    let (chi_pos, chi_neg, chi_total) = chi_sep_medi_core(
174        &pad3d(local_field, dims, fast),
175        &pad3d(r2prime, dims, fast),
176        &pad3d(magnitude, dims, fast),
177        &pad3d(mask, dims, fast),
178        &pgrid,
179        bdir,
180        params,
181        progress,
182    );
183    (
184        unpad3d(&chi_pos, fast, dims),
185        unpad3d(&chi_neg, fast, dims),
186        unpad3d(&chi_total, fast, dims),
187    )
188}
189
190#[allow(clippy::too_many_arguments)]
191fn chi_sep_medi_core<F>(
192    local_field: &[f64],
193    r2prime: &[f64],
194    magnitude: &[f64],
195    mask: &[u8],
196    grid: &Grid,
197    bdir: (f64, f64, f64),
198    params: &ChiSepParams,
199    mut progress: F,
200) -> (Vec<f64>, Vec<f64>, Vec<f64>)
201where
202    F: FnMut(usize, usize),
203{
204    let cf = params.cf;
205    let lambda_para = params.lambda_para;
206    let lambda_dia = params.lambda_dia;
207    let lambda_cpl = params.lambda_cpl;
208    let dr_pos = params.dr_pos;
209    let dr_neg = params.dr_neg;
210    let percentage = params.percentage;
211    let cg_tol = params.cg_tol;
212    let cg_max_iter = params.cg_max_iter;
213    let max_iter = params.max_iter;
214    let tol = params.tol;
215    let (nx, ny, nz) = grid.dims;
216    let (vsx, vsy, vsz) = grid.voxel_size;
217    let n = nx * ny * nz;
218    let ppm_factor = (1.0e6 / cf) as f32;
219
220    let vsx_f32 = vsx as f32;
221    let vsy_f32 = vsy as f32;
222    let vsz_f32 = vsz as f32;
223    let bdir_f32 = (bdir.0 as f32, bdir.1 as f32, bdir.2 as f32);
224    let lambda_para_f32 = lambda_para as f32;
225    let lambda_dia_f32 = lambda_dia as f32;
226    let lambda_cpl_f32 = lambda_cpl as f32;
227    let cg_tol_f32 = cg_tol as f32;
228    let tol_f32 = tol as f32;
229
230    // Effective relaxivities: Dr(Hz/ppm) * ppm_factor(ppm/Hz) = dimensionless
231    // R2'(Hz) = Dr_pos * chi_ppm = Dr_pos * (chi_Hz * ppm_factor) = dr_p_eff * chi_Hz
232    let dr_p_eff = ppm_factor * dr_pos as f32;
233    let dr_q_eff = ppm_factor * dr_neg as f32;
234
235    // Auto-tune R2' normalization for field-strength-independent gauge breaking.
236    //
237    // The chi-sep gauge mode (χ+ grows, χ- shrinks equally) lives in the null space
238    // of the field Hessian and is ONLY constrained by R2'. The gauge eigenvalue is
239    // λ_cpl * (dr_p + dr_q)² which at 7T is only ~23 vs ~1000 for TV modes.
240    //
241    // We scale the R2' equation (both data and relaxivities) by r2_scale so that the
242    // gauge eigenvalue = max(λ_para, λ_dia), matching the TV regularization strength.
243    // This doesn't change the solution (R2' residual is still 0 at truth).
244    //
245    // r2_scale = sqrt(target / (λ_cpl * dr_sum²))
246    // → eigenvalue = λ_cpl * (r2_scale * dr_sum)² = target
247    let dr_sum = dr_p_eff + dr_q_eff;
248    let target_eig = 10.0 * lambda_para_f32.max(lambda_dia_f32);
249    let r2_scale = (target_eig / (lambda_cpl_f32 * dr_sum * dr_sum)).sqrt();
250    let dr_p_use = dr_p_eff * r2_scale;
251    let dr_q_use = dr_q_eff * r2_scale;
252
253    // Local field arrives in ppm (library-wide convention); the solver works in
254    // Hz internally, so convert on entry.
255    let field_f32: Vec<f32> = local_field.iter()
256        .zip(mask.iter())
257        .map(|(&v, &m)| if m != 0 { (v * cf * 1.0e-6) as f32 } else { 0.0 })
258        .collect();
259    let r2p_f32: Vec<f32> = r2prime.iter()
260        .zip(mask.iter())
261        .map(|(&v, &m)| if m != 0 { (v as f32) * r2_scale } else { 0.0 })
262        .collect();
263    let mag_f32: Vec<f32> = magnitude.iter().map(|&v| v as f32).collect();
264
265    let mut ws = ChiSepWorkspace::new(nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
266    let d_kernel = dipole_kernel_f32(grid, bdir_f32);
267    let (mx, my, mz) = gradient_mask_f32(
268        &mag_f32, mask, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32, percentage as f32,
269    );
270
271    // chi_pos >= 0 (paramagnetic), chi_neg <= 0 (diamagnetic), both in Hz
272    let mut chi_pos = vec![0.0f32; n];
273    let mut chi_neg = vec![0.0f32; n];
274
275    let mut vr_pos = vec![0.0f32; n];
276    let mut vr_neg = vec![0.0f32; n];
277    let mut vr_sum = vec![0.0f32; n];
278    let n2 = 2 * n;
279    let mut dx = vec![0.0f32; n2];
280    let mut rhs = vec![0.0f32; n2];
281    let mut chi_sum_buf = vec![0.0f32; n];
282    let mut field_residual = vec![0.0f32; n];
283    let mut r2_residual = vec![0.0f32; n];
284    // Scratch for the CG operator (avoids per-iteration buffer clones).
285    let mut stage = vec![0.0f32; n];
286
287    // TV weight on (chi+ + chi-) sum — paper uses 2*lambda for sum term
288    // Disabled for now (0.0) pending parameter tuning; the sum TV can over-couple components
289    let lambda_sum_f32 = 0.0_f32;
290
291    let eps = 1.0e-6_f32;
292
293    for iter in 0..max_iter {
294        progress(iter + 1, max_iter);
295
296        // --- TV reweighting ---
297        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
298            &chi_pos, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
299        compute_p_weights_f32(&mut vr_pos, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
300
301        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
302            &chi_neg, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
303        compute_p_weights_f32(&mut vr_neg, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
304
305        // TV weights for the sum (chi+ + chi-)
306        for i in 0..n {
307            chi_sum_buf[i] = chi_pos[i] + chi_neg[i];
308        }
309        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
310            &chi_sum_buf, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
311        compute_p_weights_f32(&mut vr_sum, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
312
313        // --- Residuals (chi_sum_buf still holds chi_pos + chi_neg) ---
314        // field_residual = field - D*(chi_pos + chi_neg)
315        ws.fft_ws.apply_dipole_inplace(&chi_sum_buf, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
316        for i in 0..n {
317            field_residual[i] = field_f32[i] - ws.dipole_buf[i];
318        }
319
320        // R2' residual (normalized): r2_norm - dr_p_use * chi_pos + dr_q_use * chi_neg
321        // where r2_norm = R2' / (dr_p_eff + dr_q_eff), dr_p_use = dr_p_eff / dr_sum
322        for i in 0..n {
323            r2_residual[i] = r2p_f32[i] - dr_p_use * chi_pos[i] + dr_q_use * chi_neg[i];
324        }
325
326        // --- Build gradient (MEDI convention: b_orig = gradient, then negate) ---
327        //
328        // Gradient derived from first principles:
329        //   J = λ_para*TV(χ+) + λ_dia*TV(χ-) + λ_cpl/2*||field_res||² + λ_cpl/2*||r2_res||²
330        //
331        // ∂J/∂χ+ = λ_para*TV_grad(χ+) - λ_cpl*D(field_res) - λ_cpl*r2_res*dr_p_eff
332        // ∂J/∂χ- = λ_dia*TV_grad(χ-)  - λ_cpl*D(field_res) + λ_cpl*r2_res*dr_q_eff
333        //
334        // TV_grad(χ) = bdiv(wG*Vr*wG*fgrad(χ)) [MEDI convention, IS the gradient]
335        //   because bdiv = -div, so bdiv(wG*Vr*wG*∇χ) = -div(wG*Vr*wG*∇χ) = ∂TV/∂χ
336        //
337        // Field: ∂||field_res||²/∂χ+ = -2*D(field_res), with 1/2 → -D(field_res)
338        //   Same for χ- since chi_total = χ+ + χ-, ∂/∂χ- has same sign
339        //
340        // R2': r2_res = R2' - dr_p*χ+ + dr_q*χ- (since |χ-| = -χ-)
341        //   ∂r2_res/∂χ+ = -dr_p → ∂||r2_res||²/∂χ+ = -2*r2_res*dr_p, with 1/2 → -r2_res*dr_p
342        //   ∂r2_res/∂χ- = +dr_q → ∂||r2_res||²/∂χ- = +2*r2_res*dr_q, with 1/2 → +r2_res*dr_q
343
344        // TV gradient for chi_pos
345        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
346            &chi_pos, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
347        apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
348            &mx, &my, &mz, &vr_pos, &ws.gx, &ws.gy, &ws.gz);
349        bdiv_periodic_inplace_f32(&mut ws.div_buf,
350            &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
351        for i in 0..n {
352            rhs[i] = lambda_para_f32 * ws.div_buf[i];
353        }
354
355        // TV gradient for chi_neg
356        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
357            &chi_neg, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
358        apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
359            &mx, &my, &mz, &vr_neg, &ws.gx, &ws.gy, &ws.gz);
360        bdiv_periodic_inplace_f32(&mut ws.div_buf,
361            &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
362        for i in 0..n {
363            rhs[n + i] = lambda_dia_f32 * ws.div_buf[i];
364        }
365
366        // TV gradient for (chi+ + chi-) sum — applied equally to both components
367        for i in 0..n {
368            chi_sum_buf[i] = chi_pos[i] + chi_neg[i];
369        }
370        fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
371            &chi_sum_buf, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
372        apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
373            &mx, &my, &mz, &vr_sum, &ws.gx, &ws.gy, &ws.gz);
374        bdiv_periodic_inplace_f32(&mut ws.div_buf,
375            &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
376        for i in 0..n {
377            let tv_sum = lambda_sum_f32 * ws.div_buf[i];
378            rhs[i] += tv_sum;
379            rhs[n + i] += tv_sum;
380        }
381
382        // Field fidelity gradient: -λ_cpl * D(field_res), SAME sign for both
383        ws.fft_ws.apply_dipole_inplace(&field_residual, &d_kernel,
384            &mut ws.dipole_buf, &mut ws.complex_buf);
385        for i in 0..n {
386            let fg = lambda_cpl_f32 * ws.dipole_buf[i];
387            rhs[i] -= fg;
388            rhs[n + i] -= fg;
389        }
390
391        // R2' fidelity gradient (normalized):
392        //   ∂J/∂χ+ contribution: -λ_cpl * r2_res * dr_p_use
393        //   ∂J/∂χ- contribution: +λ_cpl * r2_res * dr_q_use
394        for i in 0..n {
395            if mask[i] == 0 { continue; }
396            rhs[i] -= lambda_cpl_f32 * r2_residual[i] * dr_p_use;
397            rhs[n + i] += lambda_cpl_f32 * r2_residual[i] * dr_q_use;
398        }
399
400        // Negate for CG: b = -gradient
401        for v in rhs.iter_mut() {
402            *v = -*v;
403        }
404
405        // --- CG solve ---
406        cg_solve_chisep(
407            &mut ws, &d_kernel,
408            &mx, &my, &mz,
409            &vr_pos, &vr_neg, &vr_sum,
410            lambda_para_f32, lambda_dia_f32, lambda_sum_f32, lambda_cpl_f32,
411            dr_p_use, dr_q_use,
412            mask,
413            &rhs, &mut dx, &mut stage,
414            cg_tol_f32, cg_max_iter,
415        );
416
417        // --- Update (half Newton step for stability with sign constraints) ---
418        for i in 0..n {
419            chi_pos[i] += 0.5 * dx[i];
420            chi_neg[i] += 0.5 * dx[n + i];
421        }
422
423        // --- Enforce constraints: chi_pos >= 0, chi_neg <= 0 ---
424        for i in 0..n {
425            if mask[i] == 0 {
426                chi_pos[i] = 0.0;
427                chi_neg[i] = 0.0;
428            } else {
429                chi_pos[i] = chi_pos[i].max(0.0);
430                chi_neg[i] = chi_neg[i].min(0.0);
431            }
432        }
433
434        // --- Convergence check ---
435        let update_norm = norm_squared_f32(&dx).sqrt();
436        let sol_norm = (norm_squared_f32(&chi_pos) + norm_squared_f32(&chi_neg)).sqrt();
437        let ratio = update_norm / (sol_norm + 1e-6);
438
439        if ratio < tol_f32 {
440            break;
441        }
442    }
443
444    // Convert Hz -> ppm
445    let chi_pos_out: Vec<f64> = chi_pos.iter()
446        .zip(mask.iter())
447        .map(|(&v, &m)| if m == 0 { 0.0 } else { (v * ppm_factor) as f64 })
448        .collect();
449    let chi_neg_out: Vec<f64> = chi_neg.iter()
450        .zip(mask.iter())
451        .map(|(&v, &m)| if m == 0 { 0.0 } else { (v * ppm_factor) as f64 })
452        .collect();
453    let chi_total: Vec<f64> = chi_pos_out.iter()
454        .zip(chi_neg_out.iter())
455        .map(|(&p, &n)| p + n)
456        .collect();
457
458    (chi_pos_out, chi_neg_out, chi_total)
459}
460
461/// Apply chi-sep Hessian operator A to doubled vector dx = [d_pos; d_neg].
462///
463/// A_pos = λ_para * TV_hess(d_pos) + λ_sum * TV_hess_sum(d_pos+d_neg)
464///       + λ_cpl * D²(d_pos + d_neg) + λ_cpl * dr_p * (dr_p * d_pos - dr_q * d_neg)
465/// A_neg = λ_dia * TV_hess(d_neg) + λ_sum * TV_hess_sum(d_pos+d_neg)
466///       + λ_cpl * D²(d_pos + d_neg) - λ_cpl * dr_q * (dr_p * d_pos - dr_q * d_neg)
467///
468/// TV_hessian uses IRLS: bdiv(wG*Vr*wG*fgrad(.)), positive semi-definite.
469/// Field fidelity: D², same for both, positive semi-definite.
470/// R2': rank-1 structure [dr_p, -dr_q]^T * [dr_p, -dr_q], positive semi-definite.
471#[allow(clippy::too_many_arguments)]
472fn apply_chisep_operator(
473    ws: &mut ChiSepWorkspace,
474    d_kernel: &[f32],
475    mx: &[f32], my: &[f32], mz: &[f32],
476    vr_pos: &[f32], vr_neg: &[f32], vr_sum: &[f32],
477    lambda_para: f32, lambda_dia: f32, lambda_sum: f32, lambda_cpl: f32,
478    dr_p: f32, dr_q: f32,
479    mask: &[u8],
480    dx: &[f32],
481    out: &mut [f32],
482    stage: &mut [f32],
483) {
484    let n = ws.n;
485    let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
486    let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
487
488    let d_pos = &dx[..n];
489    let d_neg = &dx[n..];
490
491    // TV for chi_pos: λ_para * bdiv(wG*Vr_pos*wG*fgrad(d_pos))
492    fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
493        d_pos, nx, ny, nz, vsx, vsy, vsz);
494    apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
495        mx, my, mz, vr_pos, &ws.gx, &ws.gy, &ws.gz);
496    bdiv_periodic_inplace_f32(&mut ws.div_buf,
497        &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
498    for i in 0..n {
499        out[i] = lambda_para * ws.div_buf[i];
500    }
501
502    // TV for chi_neg: λ_dia * bdiv(wG*Vr_neg*wG*fgrad(d_neg))
503    fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
504        d_neg, nx, ny, nz, vsx, vsy, vsz);
505    apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
506        mx, my, mz, vr_neg, &ws.gx, &ws.gy, &ws.gz);
507    bdiv_periodic_inplace_f32(&mut ws.div_buf,
508        &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
509    for i in 0..n {
510        out[n + i] = lambda_dia * ws.div_buf[i];
511    }
512
513    // TV for sum (d_pos + d_neg): λ_sum * bdiv(wG*Vr_sum*wG*fgrad(d_pos+d_neg))
514    // Applied equally to both components. `stage` holds the sum and stays valid
515    // through the field-fidelity term below (no clones in this CG hot path).
516    for i in 0..n {
517        stage[i] = d_pos[i] + d_neg[i];
518    }
519    fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
520        stage, nx, ny, nz, vsx, vsy, vsz);
521    apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
522        mx, my, mz, vr_sum, &ws.gx, &ws.gy, &ws.gz);
523    bdiv_periodic_inplace_f32(&mut ws.div_buf,
524        &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
525    for i in 0..n {
526        let tv_s = lambda_sum * ws.div_buf[i];
527        out[i] += tv_s;
528        out[n + i] += tv_s;
529    }
530
531    // Field fidelity: λ_cpl * D²(d_pos + d_neg), SAME for both
532    ws.fft_ws.apply_dipole_inplace(stage, d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
533    stage.copy_from_slice(&ws.dipole_buf);
534    ws.fft_ws.apply_dipole_inplace(stage, d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
535
536    for i in 0..n {
537        let ff = lambda_cpl * ws.dipole_buf[i];
538        out[i] += ff;
539        out[n + i] += ff;
540    }
541
542    // R2' fidelity: rank-1 Hessian [dr_p, -dr_q]^T * [dr_p, -dr_q]
543    // r2_lin = dr_p * d_pos - dr_q * d_neg
544    // out_pos += λ_cpl * dr_p * r2_lin
545    // out_neg -= λ_cpl * dr_q * r2_lin
546    for i in 0..n {
547        if mask[i] == 0 { continue; }
548        let r2_lin = dr_p * d_pos[i] - dr_q * d_neg[i];
549        out[i] += lambda_cpl * dr_p * r2_lin;
550        out[n + i] -= lambda_cpl * dr_q * r2_lin;
551    }
552}
553
554/// CG solver for the doubled chi-sep system.
555#[allow(clippy::too_many_arguments)]
556fn cg_solve_chisep(
557    ws: &mut ChiSepWorkspace,
558    d_kernel: &[f32],
559    mx: &[f32], my: &[f32], mz: &[f32],
560    vr_pos: &[f32], vr_neg: &[f32], vr_sum: &[f32],
561    lambda_para: f32, lambda_dia: f32, lambda_sum: f32, lambda_cpl: f32,
562    dr_p: f32, dr_q: f32,
563    mask: &[u8],
564    b: &[f32],
565    x: &mut [f32],
566    stage: &mut [f32],
567    tol: f32,
568    max_iter: usize,
569) {
570    let n2 = 2 * ws.n;
571    x.fill(0.0);
572
573    let mut cg_r = vec![0.0f32; n2];
574    let mut cg_p = vec![0.0f32; n2];
575    let mut cg_ap = vec![0.0f32; n2];
576
577    cg_r.copy_from_slice(&b[..n2]);
578    cg_p.copy_from_slice(&cg_r);
579
580    let mut rsold = dot_product_f32(&cg_r, &cg_r);
581    let b_norm = dot_product_f32(b, b).sqrt();
582
583    if b_norm < 1e-10 {
584        return;
585    }
586
587    for _cg_iter in 0..max_iter {
588        apply_chisep_operator(
589            ws, d_kernel, mx, my, mz,
590            vr_pos, vr_neg, vr_sum,
591            lambda_para, lambda_dia, lambda_sum, lambda_cpl,
592            dr_p, dr_q,
593            mask,
594            &cg_p, &mut cg_ap, stage,
595        );
596
597        let pap = dot_product_f32(&cg_p, &cg_ap);
598        if pap.abs() < 1e-15 {
599            break;
600        }
601
602        let alpha = rsold / pap;
603        axpy_f32(x, alpha, &cg_p);
604        axpy_f32(&mut cg_r, -alpha, &cg_ap);
605
606        let rsnew = dot_product_f32(&cg_r, &cg_r);
607        if rsnew.sqrt() < tol * b_norm {
608            break;
609        }
610
611        let beta_cg = rsnew / rsold;
612        xpby_f32(&mut cg_p, &cg_r, beta_cg);
613        rsold = rsnew;
614    }
615}
616
617#[cfg(test)]
618mod tests {
619    use super::*;
620    use crate::Grid;
621    use crate::kernels::dipole::dipole_kernel;
622    use crate::fft::{fft3d_real, ifft3d_real};
623
624    fn make_sphere(nx: usize, ny: usize, nz: usize, cx: f64, cy: f64, cz: f64, r: f64) -> Vec<f64> {
625        let mut vol = vec![0.0; nx * ny * nz];
626        for k in 0..nz {
627            for j in 0..ny {
628                for i in 0..nx {
629                    let dx = i as f64 - cx;
630                    let dy = j as f64 - cy;
631                    let dz = k as f64 - cz;
632                    if dx * dx + dy * dy + dz * dz <= r * r {
633                        vol[i + j * nx + k * nx * ny] = 1.0;
634                    }
635                }
636            }
637        }
638        vol
639    }
640
641    #[test]
642    fn test_chi_sep_medi_basic() {
643        let (nx, ny, nz) = (32, 32, 32);
644        let n = nx * ny * nz;
645        let grid = Grid::new(nx, ny, nz, 1.0, 1.0, 1.0);
646        let bdir = (0.0, 0.0, 1.0);
647        let cf: f64 = 123.2e6; // 3T
648
649        let chi_pos_true_ppm = 0.05;
650        let chi_neg_true_ppm = -0.03;
651
652        let sphere_inner = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 4.0);
653        let sphere_outer = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 8.0);
654        let brain_mask = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 12.0);
655
656        let mut chi_pos_ppm = vec![0.0f64; n];
657        let mut chi_neg_ppm = vec![0.0f64; n];
658        for i in 0..n {
659            if sphere_inner[i] > 0.5 {
660                chi_pos_ppm[i] = chi_pos_true_ppm;
661            }
662            if sphere_outer[i] > 0.5 && sphere_inner[i] < 0.5 {
663                chi_neg_ppm[i] = chi_neg_true_ppm;
664            }
665        }
666
667        // Forward model: field_ppm = D * chi_total_ppm (library units convention)
668        let chi_total_ppm: Vec<f64> = chi_pos_ppm.iter()
669            .zip(chi_neg_ppm.iter())
670            .map(|(&p, &n)| p + n)
671            .collect();
672        let d = dipole_kernel(&grid, bdir);
673        let chi_fft = fft3d_real(&chi_total_ppm, nx, ny, nz);
674        let field_fft: Vec<_> = chi_fft.iter()
675            .zip(d.iter())
676            .map(|(&c, &dk)| c * dk)
677            .collect();
678        let local_field = ifft3d_real(&field_fft, nx, ny, nz);
679
680        // R2'(Hz) = Dr_pos * |chi+_ppm| + Dr_neg * |chi-_ppm|
681        let dr_pos: f64 = 114.0;
682        let dr_neg: f64 = 30.0;
683        let r2prime: Vec<f64> = (0..n).map(|i| {
684            dr_pos * chi_pos_ppm[i].abs() + dr_neg * chi_neg_ppm[i].abs()
685        }).collect();
686
687        let mask: Vec<u8> = brain_mask.iter()
688            .map(|&v| if v > 0.5 { 1 } else { 0 })
689            .collect();
690
691        let magnitude: Vec<f64> = (0..n).map(|i| {
692            if mask[i] == 0 { return 0.0; }
693            let base = 100.0;
694            if sphere_inner[i] > 0.5 {
695                base * 1.5
696            } else if sphere_outer[i] > 0.5 {
697                base * 0.7
698            } else {
699                base
700            }
701        }).collect();
702
703        let params = ChiSepParams {
704            cf,
705            lambda_para: 1000.0, lambda_dia: 1000.0, lambda_cpl: 100.0,
706            dr_pos, dr_neg,
707            percentage: 0.3, cg_tol: 0.01, cg_max_iter: 100, max_iter: 10, tol: 0.1,
708        };
709        let (chi_pos_out, chi_neg_out, chi_total_out) = chi_sep_medi(
710            &local_field, &r2prime, &magnitude, &mask,
711            &grid, bdir, &params,
712            |_, _| {},
713        );
714
715        // chi+ should be non-negative, chi- non-positive
716        for i in 0..n {
717            if mask[i] != 0 {
718                assert!(chi_pos_out[i] >= -1e-10,
719                    "chi+ should be non-negative at voxel {}, got {}", i, chi_pos_out[i]);
720                assert!(chi_neg_out[i] <= 1e-10,
721                    "chi- should be non-positive at voxel {}, got {}", i, chi_neg_out[i]);
722            }
723        }
724
725        for i in 0..n {
726            let diff = (chi_total_out[i] - chi_pos_out[i] - chi_neg_out[i]).abs();
727            assert!(diff < 1e-10, "chi_total != chi+ + chi- at voxel {}", i);
728        }
729
730        let pos_max = chi_pos_out.iter().cloned().fold(0.0_f64, f64::max);
731        let neg_min = chi_neg_out.iter().cloned().fold(0.0_f64, f64::min);
732        assert!(pos_max > 0.0, "chi+ should have positive values, max={}", pos_max);
733        assert!(neg_min < 0.0, "chi- should have negative values, min={}", neg_min);
734    }
735}