Skip to main content

qsm_core/separation/
hc_chisep.rs

1//! HC-ChiSep: hollow-cylinder χ-separation with signal-derived fiber orientation.
2//!
3//! HC-ChiSep separates paramagnetic (χ+, iron) and diamagnetic (χ−, myelin)
4//! susceptibility from a conventional QSM (χ_total), an R2' map, and multi-echo
5//! GRE magnitude — deriving white-matter fibre orientation from the magnitude's
6//! multi-compartment interference pattern (Wharton & Bowtell hollow-cylinder
7//! model) rather than requiring DTI.
8//!
9//! The pipeline (headline mode) is:
10//! 1. **Dr+ self-calibration** — the paramagnetic relaxivity is estimated from the
11//!    5th percentile of `R2'/χ_total` over confidently-paramagnetic voxels,
12//!    falling back to the field-scaled empirical `137·B0/3` Hz/ppm.
13//! 2. **Closed-form two-source solve** — `χ+ = R2'/Dr+`, `χ− = χ+ − χ_total`
14//!    (Ridani convention: no diamagnetic relaxivity outside myelinated WM),
15//!    clamped to keep χ+ ≥ 0, |χ−| ≥ 0. This is the non-WM branch.
16//! 3. **WM-likeness (beat) weighting** — a soft weight `w` from a hollow-cylinder
17//!    vs mono-exponential model-selection test on the magnitude decay, gated by
18//!    χ_total sign. Voxels whose magnitude shows the multi-compartment "beat" and
19//!    are diamagnetic-leaning get `w → 1`.
20//! 4. **θ + MWF grid fit** — for supported voxels, fibre angle θ and myelin-water
21//!    fraction (MWF) are found by an R2'-anchored grid search over the hollow-
22//!    cylinder magnitude library, then spatially regularised.
23//! 5. **Separation** — in WM, |χ−| = K_χ·MWF (myelin-content ↔ MWF anchor,
24//!    0.038 ppm ↔ MWF 0.12), χ+ = χ_total + |χ−| shrunk toward a self-calibrated
25//!    median; the two branches are blended by `w`.
26//!
27//! All the reference's env-gated optional regularisers (TV/guided joint inversion,
28//! R2' denoise, hard constraint) are **off** in headline mode and are omitted here.
29//!
30//! Inputs are in the project's standard units — χ_total in ppm, R2' in Hz, TE in
31//! seconds, B0 in Tesla — and outputs follow the χ-separation sign convention
32//! (χ+ ≥ 0, χ− ≤ 0, χ_total = χ+ + χ−).
33//!
34//! Reference:
35//! Wharton, S. & Bowtell, R. (2012). "Fiber orientation-dependent white matter
36//! contrast in gradient echo MRI." PNAS 109(45):18559-18564. HC-ChiSep is the
37//! QSM-CI submission building on this biophysical model.
38
39use crate::Grid;
40#[cfg(feature = "parallel")]
41use rayon::prelude::*;
42
43use std::f64::consts::PI;
44
45/// Reduced gyromagnetic ratio (Hz/T) used by the hollow-cylinder model.
46const GAMMA_BAR: f64 = 42.577e6;
47
48// --- Hollow-cylinder white-matter model constants (Wharton & Bowtell Table 3) ---
49const CHI_I: f64 = -0.06e-6; // isotropic susceptibility
50const CHI_A: f64 = -0.10e-6; // anisotropic susceptibility
51const E_EXCH: f64 = 0.02e-6; // isotropic exchange offset
52const G_RATIO: f64 = 0.7;
53const T2_M: f64 = 10e-3; // myelin-water T2 (s)
54const T2_A: f64 = 64e-3; // axonal-water T2 (s)
55const T2_E: f64 = 48e-3; // extra-axonal-water T2 (s)
56const F_AXON: f64 = 0.55; // axonal fraction of non-myelin water
57
58// --- Myelin-content anchor & MWF bounds ---
59const CHI_NEG_REF: f64 = 0.038; // ppm |χ−| at reference MWF
60const MWF_REF: f64 = 0.12;
61const MWF_MIN: f64 = 0.03;
62const MWF_MAX: f64 = 0.25;
63
64// --- Locked hyperparameters (headline mode) ---
65const SMOOTH_CLS: [f64; 2] = [0.75, 1.0]; // classification smoothing sigmas
66const SMOOTH_FIT: f64 = 1.0;
67const W_CENTER: f64 = 1.5;
68const W_SCALE: f64 = 0.3;
69const CHI_GATE_C: f64 = 0.01;
70const CHI_GATE_S: f64 = 0.01;
71const LAM: f64 = 0.7;
72const NCONV_SIGMA: f64 = 1.5;
73const W_SUPPORT: f64 = 0.02;
74const NO_BEAT_FRAC: f64 = 0.01;
75
76/// `K_χ`: ppm of |χ−| per unit MWF (= CHI_NEG_REF / MWF_REF).
77const K_CHI: f64 = CHI_NEG_REF / MWF_REF;
78
79// θ and MWF grids.
80const NT: usize = 61; // 0 .. 90 step 1.5
81const NM: usize = 45; // 0.03 .. 0.25 step 0.005
82
83/// Parameters for [`hc_chisep`].
84#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
85#[derive(Clone, Debug)]
86pub struct HcChisepParams {
87    /// Main field strength in Tesla.
88    pub b0: f64,
89    /// Spin-echo echo times in **seconds** (empty = no SE evidence used).
90    pub se_echo_times: Vec<f64>,
91    /// Paramagnetic relaxivity at 3 T in Hz/ppm (empirical Shin 2021: 137).
92    pub dr_pos_3t: f64,
93    /// R2' bin width (Hz) for the anchored grid search (reference default 0.25).
94    pub bin_hz: f64,
95}
96
97impl Default for HcChisepParams {
98    fn default() -> Self {
99        Self {
100            b0: 7.0,
101            se_echo_times: Vec::new(),
102            dr_pos_3t: 137.0,
103            bin_hz: 0.25,
104        }
105    }
106}
107
108/// HC-ChiSep source separation from a QSM, an R2' map and multi-echo magnitude.
109///
110/// # Arguments
111/// * `chi_total` — Conventional QSM χ_total in **ppm** (`nx·ny·nz`, column-major).
112/// * `r2prime` — R2' map in **Hz** (`nx·ny·nz`).
113/// * `magnitude` — Multi-echo GRE magnitude, flattened as `(n_voxels, n_echoes)`
114///   in row-major order (echo fastest per voxel). Normalised per voxel internally.
115/// * `echo_times` — GRE echo times in **seconds** (`n_echoes`).
116/// * `se_magnitude` — Optional multi-echo spin-echo magnitude, `(n_voxels,
117///   n_se_echoes)`; used as soft MWF/pool-T2 evidence when present (with
118///   `params.se_echo_times`).
119/// * `mask` — Binary brain mask (`nx·ny·nz`, 1 = inside).
120/// * `grid` — Volume dimensions and voxel sizes.
121/// * `params` — See [`HcChisepParams`].
122/// * `progress` — Progress callback `(stage_done, stage_total)`.
123///
124/// # Returns
125/// `(chi_pos, chi_neg, chi_total)` in ppm, restricted to `mask` — matching the
126/// χ-separation convention: `chi_pos` ≥ 0 (paramagnetic), `chi_neg` ≤ 0
127/// (diamagnetic, signed), and `chi_total = chi_pos + chi_neg`.
128#[allow(clippy::too_many_arguments)]
129pub fn hc_chisep(
130    chi_total: &[f64],
131    r2prime: &[f64],
132    magnitude: &[f64],
133    echo_times: &[f64],
134    se_magnitude: Option<&[f64]>,
135    mask: &[u8],
136    grid: &Grid,
137    params: &HcChisepParams,
138    mut progress: impl FnMut(usize, usize),
139) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
140    let dims = grid.dims;
141    let n = dims.0 * dims.1 * dims.2;
142    let ne = echo_times.len();
143    assert_eq!(chi_total.len(), n, "chi_total length must match grid");
144    assert_eq!(r2prime.len(), n, "r2prime length must match grid");
145    assert_eq!(
146        magnitude.len(),
147        n * ne,
148        "magnitude must be n_voxels * n_echoes"
149    );
150    assert_eq!(mask.len(), n, "mask length must match grid");
151
152    let b0 = params.b0;
153
154    // Convert the voxel-major public inputs (i*ne + e) to echo-major contiguous
155    // per-echo volumes (e*n + i) so the spatial filters and gathers are simple.
156    let mag_em = to_echo_major(magnitude, n, ne);
157    let se_e_len = params.se_echo_times.len();
158    let se_em: Option<Vec<f64>> = match se_magnitude {
159        Some(s) if se_e_len >= 2 && s.len() == n * se_e_len => Some(to_echo_major(s, n, se_e_len)),
160        _ => None,
161    };
162
163    let n_stages = 5;
164    let mut stage = 0usize;
165    let mut tick = |s: &mut usize, p: &mut dyn FnMut(usize, usize)| {
166        *s += 1;
167        p(*s, n_stages);
168    };
169
170    // --- Stage 1: Dr+ self-calibration ---------------------------------------
171    let dr_default = params.dr_pos_3t * b0 / 3.0;
172    let dr_pos = {
173        let ratios: Vec<f64> = (0..n)
174            .filter(|&i| mask[i] != 0 && chi_total[i] > 0.02)
175            .map(|i| r2prime[i] / chi_total[i])
176            .collect();
177        if ratios.len() > 5000 {
178            let p5 = percentile(&ratios, 5.0);
179            if (0.3 * dr_default..=3.0 * dr_default).contains(&p5) {
180                p5
181            } else {
182                dr_default
183            }
184        } else {
185            dr_default
186        }
187    };
188    tick(&mut stage, &mut progress);
189
190    // --- Stage 2: closed-form two-source solve (non-WM branch) ---------------
191    // cf_pos = clip(R2'/Dr+, 0), cf_neg = clip(cf_pos - χ_total, 0) with the
192    // χ_total-consistency fix where the pair is inconsistent. cf_neg is a
193    // positive magnitude here.
194    let mut cf_pos = vec![0.0_f64; n];
195    let mut cf_neg = vec![0.0_f64; n];
196    for i in 0..n {
197        let p = (r2prime[i] / dr_pos).max(0.0);
198        let neg = p - chi_total[i];
199        if neg < 0.0 {
200            cf_pos[i] = chi_total[i].max(0.0);
201        } else {
202            cf_pos[i] = p;
203        }
204        cf_neg[i] = (cf_pos[i] - chi_total[i]).max(0.0);
205    }
206    tick(&mut stage, &mut progress);
207
208    // Build the anchored fitter (hollow-cylinder magnitude library).
209    let fitter = AnchoredFitter::new(echo_times, b0, &params.se_echo_times);
210    let have_se = se_em.is_some() && fitter.seml.is_some();
211    let se_e = params.se_echo_times.len();
212
213    // Brain voxel indices.
214    let bii: Vec<usize> = (0..n).filter(|&i| mask[i] != 0).collect();
215    if bii.is_empty() {
216        return finish(&cf_pos, &cf_neg, mask, n);
217    }
218    let rho_b: Vec<f64> = bii.iter().map(|&i| r2prime[i]).collect();
219
220    // --- Stage 3: WM-likeness (beat) weighting -------------------------------
221    // For each classification smoothing scale, fit the HC library and compare its
222    // SSE to a mono-exponential fit; take the elementwise-best log-ratio.
223    let maskf: Vec<f64> = mask.iter().map(|&m| m as f64).collect();
224    let mut lr_best: Option<Vec<f64>> = None;
225    for &s in &SMOOTH_CLS {
226        let m_s = smooth_stack(&mag_em, dims, ne, s);
227        let sn = gather_normalised(&m_s, &bii, ne);
228        let (_t, _m, sse_hc) = fitter.fit(&sn, &rho_b, None, None, params.bin_hz);
229        let mut num = sse_hc;
230        let mut den: Vec<f64> = (0..bii.len())
231            .map(|v| mono_sse(&sn[v * ne..v * ne + ne], echo_times))
232            .collect();
233        if have_se {
234            let se_s = smooth_stack(se_em.as_ref().unwrap(), dims, se_e, s);
235            let sen = gather_normalised(&se_s, &bii, se_e);
236            let seml = fitter.seml.as_ref().unwrap();
237            for v in 0..bii.len() {
238                // min over MWF of SE pool SSE
239                let mut best = f64::INFINITY;
240                for mi in 0..NM {
241                    let mut acc = 0.0;
242                    for e in 0..se_e {
243                        let d = sen[v * se_e + e] - seml[mi * se_e + e];
244                        acc += d * d;
245                    }
246                    best = best.min(acc);
247                }
248                num[v] += best;
249                den[v] += mono_sse(&sen[v * se_e..v * se_e + se_e], &params.se_echo_times);
250            }
251        }
252        let lr: Vec<f64> = (0..bii.len())
253            .map(|v| (num[v] / den[v].max(1e-12)).max(1e-6).log10())
254            .collect();
255        lr_best = Some(match lr_best {
256            None => lr,
257            Some(prev) => (0..bii.len()).map(|v| prev[v].min(lr[v])).collect(),
258        });
259    }
260    let lr_best = lr_best.unwrap();
261
262    // lr_med = median_filter(volmap(lr_best, fill=10.0), 3)
263    let lr_vol = volmap(&lr_best, &bii, n, 10.0);
264    let lr_med = median_filter_3(&lr_vol, dims);
265
266    // chs = NC-smoothed χ_total (gaussian(χ·mask,1)/gaussian(mask,1)).
267    let chi_masked: Vec<f64> = (0..n).map(|i| chi_total[i] * maskf[i]).collect();
268    let chs_num = gaussian_filter_3d(&chi_masked, dims, 1.0);
269    let chs_den = gaussian_filter_3d(&maskf, dims, 1.0);
270    let chs: Vec<f64> = (0..n).map(|i| chs_num[i] / chs_den[i].max(1e-6)).collect();
271
272    // w = sigmoid(lr) * sigmoid(chi), masked.
273    let mut w = vec![0.0_f64; n];
274    for i in 0..n {
275        let s1 = 1.0 / (1.0 + ((lr_med[i] - W_CENTER) / W_SCALE).exp());
276        let s2 = 1.0 / (1.0 + ((chs[i] - CHI_GATE_C) / CHI_GATE_S).exp());
277        w[i] = s1 * s2 * maskf[i];
278    }
279    let beat_frac = {
280        let cnt = bii.iter().filter(|&&i| w[i] > 0.5).count();
281        cnt as f64 / bii.len() as f64
282    };
283    tick(&mut stage, &mut progress);
284    if beat_frac < NO_BEAT_FRAC {
285        // No detectable beat anywhere → closed-form everywhere.
286        return finish(&cf_pos, &cf_neg, mask, n);
287    }
288
289    // --- Stage 4: θ + MWF grid fit on supported voxels -----------------------
290    let sii: Vec<usize> = bii.iter().cloned().filter(|&i| w[i] > W_SUPPORT).collect();
291    let rho_s: Vec<f64> = sii.iter().map(|&i| r2prime[i]).collect();
292
293    let m_fit = smooth_stack(&mag_em, dims, ne, SMOOTH_FIT);
294    let sn_sm = gather_normalised(&m_fit, &sii, ne);
295    let sn_raw = gather_normalised(&mag_em, &sii, ne);
296    let (sen_sm, sen_raw) = if have_se {
297        let se = se_em.as_ref().unwrap();
298        let se_fit = smooth_stack(se, dims, se_e, SMOOTH_FIT);
299        (
300            Some(gather_normalised(&se_fit, &sii, se_e)),
301            Some(gather_normalised(se, &sii, se_e)),
302        )
303    } else {
304        (None, None)
305    };
306
307    // θ from the smoothed magnitude, spatially regularised (median, fill 45°).
308    let (th_sm, _m_sm, _s_sm) = fitter.fit(&sn_sm, &rho_s, sen_sm.as_deref(), None, params.bin_hz);
309    // Non-supported voxels filled with 45° before the median (matches the ref).
310    let th_fill = volmap(&th_sm, &sii, n, 45.0);
311    let th_reg = median_filter_3(&th_fill, dims);
312    let th_pin: Vec<f64> = sii.iter().map(|&i| th_reg[i]).collect();
313
314    // MWF from raw and smoothed magnitude, θ pinned; average with the physics bound.
315    let (_t1, mw_raw, _s1) = fitter.fit(
316        &sn_raw,
317        &rho_s,
318        sen_raw.as_deref(),
319        Some(&th_pin),
320        params.bin_hz,
321    );
322    let (_t2, mw_sm, _s2) = fitter.fit(
323        &sn_sm,
324        &rho_s,
325        sen_sm.as_deref(),
326        Some(&th_pin),
327        params.bin_hz,
328    );
329    let mw_est: Vec<f64> = (0..sii.len())
330        .map(|v| {
331            let bnd = (-chi_total[sii[v]] / K_CHI).clamp(MWF_MIN, MWF_MAX);
332            0.5 * (mw_raw[v].max(bnd) + mw_sm[v].max(bnd))
333        })
334        .collect();
335    let mw_vol = volmap(&mw_est, &sii, n, 0.0);
336
337    // Confidence-weighted normalised convolution of MWF.
338    let wv: Vec<f64> = w.iter().map(|&x| x.clamp(0.0, 1.0)).collect();
339    let mw_w: Vec<f64> = (0..n).map(|i| mw_vol[i] * wv[i]).collect();
340    let mw_num = gaussian_filter_3d(&mw_w, dims, NCONV_SIGMA);
341    let mw_den = gaussian_filter_3d(&wv, dims, NCONV_SIGMA);
342    let mw_reg: Vec<f64> = (0..n).map(|i| mw_num[i] / mw_den[i].max(1e-3)).collect();
343    tick(&mut stage, &mut progress);
344
345    // --- Stage 5: separation + blend -----------------------------------------
346    let route: Vec<f64> = (0..n).map(|i| chi_total[i] + K_CHI * mw_reg[i]).collect();
347    let conf: Vec<usize> = (0..n).filter(|&i| w[i] > 0.5).collect();
348    let c0 = if conf.len() > 100 {
349        let vals: Vec<f64> = conf.iter().map(|&i| route[i].max(0.0)).collect();
350        median(&vals)
351    } else {
352        0.005
353    };
354    let mut chi_pos = vec![0.0_f64; n];
355    let mut chi_neg = vec![0.0_f64; n];
356    let mut chi_out = vec![0.0_f64; n];
357    for i in 0..n {
358        if mask[i] == 0 {
359            continue;
360        }
361        let wm_pos = (LAM * route[i] + (1.0 - LAM) * c0).max(0.0);
362        let wm_neg = (wm_pos - chi_total[i]).max(0.0);
363        let pos = w[i] * wm_pos + (1.0 - w[i]) * cf_pos[i];
364        let neg = w[i] * wm_neg + (1.0 - w[i]) * cf_neg[i];
365        chi_pos[i] = pos;
366        chi_neg[i] = -neg; // signed χ− ≤ 0
367        chi_out[i] = pos - neg;
368    }
369    tick(&mut stage, &mut progress);
370    (chi_pos, chi_neg, chi_out)
371}
372
373/// Emit the closed-form branch as the final result (χ+ ≥ 0, χ− ≤ 0 signed).
374fn finish(cf_pos: &[f64], cf_neg: &[f64], mask: &[u8], n: usize) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
375    let mut chi_pos = vec![0.0_f64; n];
376    let mut chi_neg = vec![0.0_f64; n];
377    let mut chi_out = vec![0.0_f64; n];
378    for i in 0..n {
379        if mask[i] != 0 {
380            chi_pos[i] = cf_pos[i];
381            chi_neg[i] = -cf_neg[i];
382            chi_out[i] = cf_pos[i] - cf_neg[i];
383        }
384    }
385    (chi_pos, chi_neg, chi_out)
386}
387
388// ---------------------------------------------------------------------------
389// Hollow-cylinder model + anchored grid-search fitter.
390// ---------------------------------------------------------------------------
391
392/// Compartment frequencies (Hz) for fibre-to-B0 angle `theta` (radians).
393/// Returns `(Δf_myelin, Δf_axon, Δf_extra)`; extra-axonal is exactly 0.
394fn hc_compartment_freqs(theta: f64, b0: f64) -> (f64, f64) {
395    let s2 = theta.sin() * theta.sin();
396    let w0 = GAMMA_BAR * b0;
397    let ln_term = 0.75 * CHI_A * (1.0 / G_RATIO).ln() * s2;
398    let f_my = w0
399        * (CHI_I * (2.0 / 3.0 - s2) / 2.0
400            + CHI_A * (1.0 / 12.0 - 5.0 / 12.0 * s2)
401            + ln_term
402            + E_EXCH);
403    let f_ax = w0 * ln_term;
404    (f_my, f_ax)
405}
406
407/// |complex hollow-cylinder GRE magnitude| at echo time `te` (s), angle `theta`
408/// (rad), myelin-water fraction `mwf`, and mesoscopic reversible rate `r2p_meso`.
409fn hc_wm_signal_mag(te: f64, theta: f64, b0: f64, mwf: f64, r2p_meso: f64) -> f64 {
410    let f_m = mwf;
411    let rest = 1.0 - f_m;
412    let f_a = rest * F_AXON;
413    let f_e = rest * (1.0 - F_AXON);
414    let (dfm, dfa) = hc_compartment_freqs(theta, b0);
415    let (mut re, mut im) = (0.0, 0.0);
416    let terms = [(f_m, T2_M, dfm), (f_a, T2_A, dfa), (f_e, T2_E, 0.0)];
417    for (f, t2, df) in terms {
418        let amp = f * (-te / t2).exp();
419        let ph = 2.0 * PI * df * te;
420        re += amp * ph.cos();
421        im += amp * ph.sin();
422    }
423    let mag = (re * re + im * im).sqrt();
424    mag * (-r2p_meso * te).exp()
425}
426
427/// Spin-echo WM magnitude factor (offsets refocused → real pool-T2 mixture).
428fn hc_wm_se_signal(te: f64, mwf: f64) -> f64 {
429    let rest = 1.0 - mwf;
430    mwf * (-te / T2_M).exp()
431        + rest * F_AXON * (-te / T2_A).exp()
432        + rest * (1.0 - F_AXON) * (-te / T2_E).exp()
433}
434
435/// Mono-exp-equivalent reversible rate the beat itself contributes (Hz):
436/// weighted-LS slope of `log(|GRE|/SE)` vs centred TE, negated and clamped ≥ 0.
437fn hc_wm_r2prime(theta: f64, mwf: f64, tes: &[f64], b0: f64) -> f64 {
438    let tbar = tes.iter().sum::<f64>() / tes.len() as f64;
439    let denom: f64 = tes.iter().map(|&t| (t - tbar) * (t - tbar)).sum();
440    let mut acc = 0.0;
441    for &te in tes {
442        let gre = hc_wm_signal_mag(te, theta, b0, mwf, 0.0);
443        let se = hc_wm_se_signal(te, mwf).max(1e-12);
444        let ratio = (gre / se).max(1e-12);
445        acc += ratio.ln() * (te - tbar);
446    }
447    (-acc / denom).max(0.0)
448}
449
450/// Precomputed hollow-cylinder magnitude library over the (θ, MWF) grid.
451struct AnchoredFitter {
452    e: usize,
453    /// `Pn[ (ti*NM + mi)*E + e ]` — first-echo-normalised magnitude library.
454    pn: Vec<f64>,
455    /// `H[ ti*NM + mi ]` — R2'_hc per grid point (Hz).
456    h: Vec<f64>,
457    /// `dte[e] = TE[e] - TE[0]`.
458    dte: Vec<f64>,
459    /// Optional SE library `SEml[ mi*se_e + e ]` (first-SE-normalised).
460    seml: Option<Vec<f64>>,
461}
462
463impl AnchoredFitter {
464    fn new(tes: &[f64], b0: f64, se_tes: &[f64]) -> Self {
465        let e = tes.len();
466        let mut pn = vec![0.0_f64; NT * NM * e];
467        let mut h = vec![0.0_f64; NT * NM];
468        for ti in 0..NT {
469            let theta = (ti as f64 * 1.5).to_radians();
470            for mi in 0..NM {
471                let mwf = MWF_MIN + mi as f64 * 0.005;
472                let base = (ti * NM + mi) * e;
473                let first = hc_wm_signal_mag(tes[0], theta, b0, mwf, 0.0).max(1e-30);
474                for (k, &te) in tes.iter().enumerate() {
475                    pn[base + k] = hc_wm_signal_mag(te, theta, b0, mwf, 0.0) / first;
476                }
477                h[ti * NM + mi] = hc_wm_r2prime(theta, mwf, tes, b0);
478            }
479        }
480        let dte: Vec<f64> = tes.iter().map(|&t| t - tes[0]).collect();
481        let seml = if se_tes.len() >= 2 {
482            let se_e = se_tes.len();
483            let mut m = vec![0.0_f64; NM * se_e];
484            for mi in 0..NM {
485                let mwf = MWF_MIN + mi as f64 * 0.005;
486                let first = hc_wm_se_signal(se_tes[0], mwf).max(1e-30);
487                for (k, &te) in se_tes.iter().enumerate() {
488                    m[mi * se_e + k] = hc_wm_se_signal(te, mwf) / first;
489                }
490            }
491            Some(m)
492        } else {
493            None
494        };
495        Self {
496            e,
497            pn,
498            h,
499            dte,
500            seml,
501        }
502    }
503
504    /// Anchored grid search. `sig_n` is `(N, E)` first-echo-normalised; `rho` is
505    /// R2' per voxel (Hz). Returns `(theta_deg, mwf, sse)` per voxel.
506    fn fit(
507        &self,
508        sig_n: &[f64],
509        rho: &[f64],
510        se_n: Option<&[f64]>,
511        theta_pin: Option<&[f64]>,
512        bin_hz: f64,
513    ) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
514        let e = self.e;
515        let n = rho.len();
516        let l = NT * NM;
517
518        // Per-voxel SE SSE (N, NM), if SE evidence is present.
519        let se_e = se_n.map(|s| s.len() / n).unwrap_or(0);
520        let sse_se: Option<Vec<f64>> = match (se_n, &self.seml) {
521            (Some(sn), Some(seml)) => {
522                let mut out = vec![0.0_f64; n * NM];
523                for v in 0..n {
524                    for mi in 0..NM {
525                        let mut acc = 0.0;
526                        for k in 0..se_e {
527                            let d = sn[v * se_e + k] - seml[mi * se_e + k];
528                            acc += d * d;
529                        }
530                        out[v * NM + mi] = acc;
531                    }
532                }
533                Some(out)
534            }
535            _ => None,
536        };
537
538        // Group voxel indices by rho bin.
539        use std::collections::HashMap;
540        let mut groups: HashMap<i64, Vec<usize>> = HashMap::new();
541        for (v, &r) in rho.iter().enumerate() {
542            let b = (r.clamp(0.0, 200.0) / bin_hz).round() as i64;
543            groups.entry(b).or_default().push(v);
544        }
545
546        let group_vec: Vec<(i64, Vec<usize>)> = groups.into_iter().collect();
547        // Solve each bin (parallel), scatter results back.
548        let per_bin: Vec<Vec<(usize, f64, f64, f64)>> = maybe_par_iter!(group_vec)
549            .map(|(b, vs)| {
550                // Candidate curves for this bin: Pn · exp(-max(b·bin_hz - H, 0)·dte).
551                let mut flat = vec![0.0_f64; l * e];
552                let mut m2 = vec![0.0_f64; l];
553                let mut bad = vec![false; l];
554                let rho_bin = *b as f64 * bin_hz;
555                for li in 0..l {
556                    let meso = rho_bin - self.h[li];
557                    bad[li] = meso < -1.0;
558                    let decay = meso.max(0.0);
559                    let mut acc = 0.0;
560                    for k in 0..e {
561                        let val = self.pn[li * e + k] * (-decay * self.dte[k]).exp();
562                        flat[li * e + k] = val;
563                        acc += val * val;
564                    }
565                    m2[li] = 0.5 * acc;
566                }
567                let all_bad = bad.iter().all(|&x| x);
568
569                let mut out = Vec::with_capacity(vs.len());
570                for &v in vs {
571                    let sig = &sig_n[v * e..v * e + e];
572                    let se_row = sse_se.as_ref().map(|s| &s[v * NM..v * NM + NM]);
573                    let pin_row = theta_pin.map(|tp| {
574                        ((tp[v] / 1.5).round() as isize).clamp(0, NT as isize - 1) as usize
575                    });
576
577                    let (mut best_score, mut best_l) = (f64::NEG_INFINITY, 0usize);
578                    let (lo, hi) = match pin_row {
579                        Some(r) => (r * NM, r * NM + NM),
580                        None => (0, l),
581                    };
582                    for li in lo..hi {
583                        if bad[li] && !all_bad {
584                            continue;
585                        }
586                        let mut s = -m2[li];
587                        for k in 0..e {
588                            s += sig[k] * flat[li * e + k];
589                        }
590                        if let Some(se_row) = se_row {
591                            s -= 0.5 * se_row[li % NM];
592                        }
593                        if s > best_score {
594                            best_score = s;
595                            best_l = li;
596                        }
597                    }
598                    // True SSE at the winner.
599                    let mut sse = 0.0;
600                    for k in 0..e {
601                        let d = sig[k] - flat[best_l * e + k];
602                        sse += d * d;
603                    }
604                    let theta_deg = (best_l / NM) as f64 * 1.5;
605                    let mwf = MWF_MIN + (best_l % NM) as f64 * 0.005;
606                    out.push((v, theta_deg, mwf, sse));
607                }
608                out
609            })
610            .collect();
611
612        let mut out_t = vec![0.0_f64; n];
613        let mut out_m = vec![MWF_REF; n];
614        let mut out_s = vec![f64::INFINITY; n];
615        for bin in per_bin {
616            for (v, t, m, s) in bin {
617                out_t[v] = t;
618                out_m[v] = m;
619                out_s[v] = s;
620            }
621        }
622        (out_t, out_m, out_s)
623    }
624}
625
626/// Mono-exponential fit SSE for one voxel's first-echo-normalised signal.
627fn mono_sse(sig: &[f64], tes: &[f64]) -> f64 {
628    let e = tes.len();
629    let tbar = tes.iter().sum::<f64>() / e as f64;
630    let logs: Vec<f64> = sig.iter().map(|&s| s.max(1e-9).ln()).collect();
631    let den: f64 = tes.iter().map(|&t| (t - tbar) * (t - tbar)).sum();
632    let b: f64 = (0..e).map(|k| logs[k] * (tes[k] - tbar)).sum::<f64>() / den;
633    let a: f64 = logs.iter().sum::<f64>() / e as f64;
634    (0..e)
635        .map(|k| {
636            let pred = (a + b * (tes[k] - tbar)).exp();
637            let d = sig[k] - pred;
638            d * d
639        })
640        .sum()
641}
642
643/// Convert a voxel-major multi-echo stack `vm[i*ne + e]` to echo-major
644/// (contiguous per-echo volumes) `em[e*n + i]`.
645fn to_echo_major(vm: &[f64], n: usize, ne: usize) -> Vec<f64> {
646    let mut out = vec![0.0_f64; n * ne];
647    for i in 0..n {
648        for e in 0..ne {
649            out[e * n + i] = vm[i * ne + e];
650        }
651    }
652    out
653}
654
655/// Gather `(len, E)` first-echo-normalised signals for the given voxel indices
656/// from an echo-major volume stack `stack[e*n + i]`.
657fn gather_normalised(stack: &[f64], idx: &[usize], ne: usize) -> Vec<f64> {
658    let n = stack.len() / ne;
659    let mut out = vec![0.0_f64; idx.len() * ne];
660    for (k, &i) in idx.iter().enumerate() {
661        let first = stack[i].max(1e-9); // e = 0
662        for e in 0..ne {
663            out[k * ne + e] = stack[e * n + i] / first;
664        }
665    }
666    out
667}
668
669/// Scatter per-voxel values back into a full volume, filling the rest with `fill`.
670fn volmap(vals: &[f64], idx: &[usize], n: usize, fill: f64) -> Vec<f64> {
671    let mut out = vec![fill; n];
672    for (k, &i) in idx.iter().enumerate() {
673        out[i] = vals[k];
674    }
675    out
676}
677
678// ---------------------------------------------------------------------------
679// Spatial filters (scipy-matching: reflect boundary).
680// ---------------------------------------------------------------------------
681
682/// Reflect an index into `[0, len)` using scipy's `reflect` mode (edge repeated).
683#[inline]
684fn reflect_index(mut p: isize, len: usize) -> usize {
685    let l = len as isize;
686    loop {
687        if p < 0 {
688            p = -p - 1;
689        } else if p >= l {
690            p = 2 * l - p - 1;
691        } else {
692            return p as usize;
693        }
694    }
695}
696
697/// Smooth each echo of an echo-major stack `stack[e*n + i]` by a 3D Gaussian.
698fn smooth_stack(stack: &[f64], dims: (usize, usize, usize), ne: usize, sigma: f64) -> Vec<f64> {
699    let n = dims.0 * dims.1 * dims.2;
700    let mut out = vec![0.0_f64; n * ne];
701    for e in 0..ne {
702        let sm = gaussian_filter_3d(&stack[e * n..e * n + n], dims, sigma);
703        out[e * n..e * n + n].copy_from_slice(&sm);
704    }
705    out
706}
707
708/// 3D separable Gaussian filter, scipy `gaussian_filter` semantics: `reflect`
709/// boundary, `truncate = 4.0`.
710fn gaussian_filter_3d(data: &[f64], dims: (usize, usize, usize), sigma: f64) -> Vec<f64> {
711    if sigma <= 0.0 {
712        return data.to_vec();
713    }
714    let (nx, ny, nz) = dims;
715    let radius = (4.0 * sigma + 0.5) as usize;
716    let mut kernel = vec![0.0_f64; 2 * radius + 1];
717    let mut ksum = 0.0;
718    for (k, w) in kernel.iter_mut().enumerate() {
719        let x = k as f64 - radius as f64;
720        *w = (-0.5 * (x / sigma) * (x / sigma)).exp();
721        ksum += *w;
722    }
723    for w in kernel.iter_mut() {
724        *w /= ksum;
725    }
726
727    let idx = |x: usize, y: usize, z: usize| (z * ny + y) * nx + x;
728    let mut a = data.to_vec();
729    let mut b = vec![0.0_f64; a.len()];
730
731    // X axis.
732    for z in 0..nz {
733        for y in 0..ny {
734            for x in 0..nx {
735                let mut s = 0.0;
736                for (kk, &w) in kernel.iter().enumerate() {
737                    let xi = reflect_index(x as isize + kk as isize - radius as isize, nx);
738                    s += w * a[idx(xi, y, z)];
739                }
740                b[idx(x, y, z)] = s;
741            }
742        }
743    }
744    std::mem::swap(&mut a, &mut b);
745    // Y axis.
746    for z in 0..nz {
747        for y in 0..ny {
748            for x in 0..nx {
749                let mut s = 0.0;
750                for (kk, &w) in kernel.iter().enumerate() {
751                    let yi = reflect_index(y as isize + kk as isize - radius as isize, ny);
752                    s += w * a[idx(x, yi, z)];
753                }
754                b[idx(x, y, z)] = s;
755            }
756        }
757    }
758    std::mem::swap(&mut a, &mut b);
759    // Z axis.
760    for z in 0..nz {
761        for y in 0..ny {
762            for x in 0..nx {
763                let mut s = 0.0;
764                for (kk, &w) in kernel.iter().enumerate() {
765                    let zi = reflect_index(z as isize + kk as isize - radius as isize, nz);
766                    s += w * a[idx(x, y, zi)];
767                }
768                b[idx(x, y, z)] = s;
769            }
770        }
771    }
772    b
773}
774
775/// 3D median filter with a 3×3×3 window (scipy `median_filter` size 3, `reflect`).
776fn median_filter_3(data: &[f64], dims: (usize, usize, usize)) -> Vec<f64> {
777    let (nx, ny, nz) = dims;
778    let idx = |x: usize, y: usize, z: usize| (z * ny + y) * nx + x;
779    let mut out = vec![0.0_f64; data.len()];
780    for z in 0..nz {
781        for y in 0..ny {
782            for x in 0..nx {
783                let mut win = [0.0_f64; 27];
784                let mut c = 0;
785                for dz in -1..=1_isize {
786                    let zi = reflect_index(z as isize + dz, nz);
787                    for dy in -1..=1_isize {
788                        let yi = reflect_index(y as isize + dy, ny);
789                        for dx in -1..=1_isize {
790                            let xi = reflect_index(x as isize + dx, nx);
791                            win[c] = data[idx(xi, yi, zi)];
792                            c += 1;
793                        }
794                    }
795                }
796                win.sort_by(|a, b| a.partial_cmp(b).unwrap());
797                out[idx(x, y, z)] = win[13];
798            }
799        }
800    }
801    out
802}
803
804// ---------------------------------------------------------------------------
805// Small statistics helpers.
806// ---------------------------------------------------------------------------
807
808/// numpy-style linear-interpolation percentile (`q` in [0, 100]).
809fn percentile(vals: &[f64], q: f64) -> f64 {
810    if vals.is_empty() {
811        return 0.0;
812    }
813    let mut v = vals.to_vec();
814    v.sort_by(|a, b| a.partial_cmp(b).unwrap());
815    let rank = q / 100.0 * (v.len() as f64 - 1.0);
816    let lo = rank.floor() as usize;
817    let hi = rank.ceil() as usize;
818    if lo == hi {
819        v[lo]
820    } else {
821        let frac = rank - lo as f64;
822        v[lo] * (1.0 - frac) + v[hi] * frac
823    }
824}
825
826fn median(vals: &[f64]) -> f64 {
827    if vals.is_empty() {
828        return 0.0;
829    }
830    let mut v = vals.to_vec();
831    v.sort_by(|a, b| a.partial_cmp(b).unwrap());
832    let m = v.len() / 2;
833    if v.len() % 2 == 1 {
834        v[m]
835    } else {
836        0.5 * (v[m - 1] + v[m])
837    }
838}
839
840#[cfg(test)]
841mod tests {
842    use super::*;
843
844    #[test]
845    fn reflect_index_matches_scipy() {
846        // scipy 'reflect' (edge repeated): [-2,-1,0,1,2,3,4,5] over len=4
847        // -> [1,0,0,1,2,3,3,2]
848        let len = 4;
849        let got: Vec<usize> = (-2..6).map(|p| reflect_index(p, len)).collect();
850        assert_eq!(got, vec![1, 0, 0, 1, 2, 3, 3, 2]);
851    }
852
853    #[test]
854    fn gaussian_preserves_constant() {
855        let dims = (8, 8, 8);
856        let n = 8 * 8 * 8;
857        let data = vec![3.0_f64; n];
858        let out = gaussian_filter_3d(&data, dims, 1.0);
859        for &v in &out {
860            assert!((v - 3.0).abs() < 1e-9, "constant not preserved: {v}");
861        }
862    }
863
864    #[test]
865    fn median_removes_impulse() {
866        let dims = (5, 5, 5);
867        let n = 125;
868        let mut data = vec![1.0_f64; n];
869        let idx = |x: usize, y: usize, z: usize| (z * 5 + y) * 5 + x;
870        data[idx(2, 2, 2)] = 100.0; // single impulse
871        let out = median_filter_3(&data, dims);
872        assert!(
873            (out[idx(2, 2, 2)] - 1.0).abs() < 1e-9,
874            "impulse survived median"
875        );
876    }
877
878    #[test]
879    fn percentile_and_median_basic() {
880        let v = [1.0, 2.0, 3.0, 4.0];
881        assert!((percentile(&v, 0.0) - 1.0).abs() < 1e-12);
882        assert!((percentile(&v, 100.0) - 4.0).abs() < 1e-12);
883        assert!((median(&v) - 2.5).abs() < 1e-12);
884        let v2 = [1.0, 2.0, 3.0];
885        assert!((median(&v2) - 2.0).abs() < 1e-12);
886    }
887
888    /// The R2'-anchored fitter should recover a planted (θ, MWF) from a synthetic
889    /// hollow-cylinder magnitude, given the matching R2'.
890    #[test]
891    fn fitter_recovers_planted_theta_mwf() {
892        let tes = [0.004, 0.012, 0.020, 0.028];
893        let b0 = 7.0;
894        let fitter = AnchoredFitter::new(&tes, b0, &[]);
895
896        let theta_deg = 30.0_f64;
897        let mwf = 0.10_f64;
898        let theta = theta_deg.to_radians();
899        // R2' = H(θ,mwf) + a chosen mesoscopic rate.
900        let meso = 8.0;
901        let h = hc_wm_r2prime(theta, mwf, &tes, b0);
902        let rho = h + meso;
903        // Build the first-echo-normalised magnitude with that meso decay.
904        let first = hc_wm_signal_mag(tes[0], theta, b0, mwf, meso);
905        let sig: Vec<f64> = tes
906            .iter()
907            .map(|&t| hc_wm_signal_mag(t, theta, b0, mwf, meso) / first)
908            .collect();
909
910        let (t, m, _s) = fitter.fit(&sig, &[rho], None, None, 0.25);
911        // MWF is well-determined; θ is only weakly identifiable from magnitude at a
912        // few echoes (near-ties across neighbouring angles), which is why the full
913        // algorithm heavily regularises and pins θ — so tolerate a loose θ here.
914        assert!((m[0] - mwf).abs() <= 0.02, "mwf {} vs {}", m[0], mwf);
915        assert!(
916            (t[0] - theta_deg).abs() <= 12.0,
917            "theta {} vs {}",
918            t[0],
919            theta_deg
920        );
921    }
922
923    /// Build a synthetic phantom: voxels with `x < wm_x_max` are WM-like (hollow-
924    /// cylinder beat magnitude, diamagnetic χ_total, R2' = H+meso so the anchored
925    /// fit matches), the rest GM-like (mono-exponential magnitude, paramagnetic
926    /// χ_total, R2' = Dr+·χ). Returns `(chi_total, r2prime, mag_vm, se_mag_vm, mask)`
927    /// with magnitudes in voxel-major `(n, ne)` layout.
928    fn synth_phantom(
929        dims: (usize, usize, usize),
930        wm_x_max: usize,
931        b0: f64,
932        tes: &[f64],
933        se_tes: &[f64],
934    ) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<u8>) {
935        let (nx, ny, nz) = dims;
936        let n = nx * ny * nz;
937        let ne = tes.len();
938        let se_e = se_tes.len();
939        let dr = 137.0 * b0 / 3.0;
940        let theta = 40.0_f64.to_radians();
941        let mwf = 0.12;
942        let meso = 6.0;
943        let h = hc_wm_r2prime(theta, mwf, tes, b0);
944        let r2star = 25.0;
945
946        let mut chi = vec![0.0_f64; n];
947        let mut r2p = vec![0.0_f64; n];
948        let mut mag = vec![0.0_f64; n * ne];
949        let mut se = vec![0.0_f64; n * se_e];
950        let mask = vec![1u8; n];
951        for k in 0..nz {
952            for j in 0..ny {
953                for i in 0..nx {
954                    let idx = (k * ny + j) * nx + i;
955                    if i < wm_x_max {
956                        chi[idx] = -0.03;
957                        r2p[idx] = h + meso;
958                        for (e, &t) in tes.iter().enumerate() {
959                            mag[idx * ne + e] = hc_wm_signal_mag(t, theta, b0, mwf, meso);
960                        }
961                        for (e, &t) in se_tes.iter().enumerate() {
962                            se[idx * se_e + e] = hc_wm_se_signal(t, mwf);
963                        }
964                    } else {
965                        chi[idx] = 0.05;
966                        r2p[idx] = dr * 0.05;
967                        for (e, &t) in tes.iter().enumerate() {
968                            mag[idx * ne + e] = (-r2star * t).exp();
969                        }
970                        for (e, &t) in se_tes.iter().enumerate() {
971                            se[idx * se_e + e] = (-r2star * t).exp();
972                        }
973                    }
974                }
975            }
976        }
977        (chi, r2p, mag, se, mask)
978    }
979
980    /// Full pipeline (with SE evidence): WM slab triggers the beat branch (stages
981    /// 3-5), GM uses the closed form; exercises Dr+ self-calibration on a >5000
982    /// paramagnetic-voxel pool.
983    #[test]
984    fn hc_chisep_end_to_end_wm_and_gm() {
985        let tes = [0.004, 0.012, 0.020, 0.028];
986        let se_tes = [0.01, 0.03, 0.05, 0.07];
987        let dims = (20, 20, 20);
988        let n = dims.0 * dims.1 * dims.2;
989        let (chi, r2p, mag, se, mask) = synth_phantom(dims, 6, 7.0, &tes, &se_tes);
990        let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
991        let params = HcChisepParams { b0: 7.0, se_echo_times: se_tes.to_vec(), ..Default::default() };
992
993        let mut stages = 0;
994        let (pos, neg, tot) = hc_chisep(
995            &chi, &r2p, &mag, &tes, Some(&se), &mask, &grid, &params, |i, _| stages = i,
996        );
997        assert_eq!(stages, 5, "all pipeline stages ran (beat detected)");
998
999        // Sign convention + invariant everywhere.
1000        for i in 0..n {
1001            assert!(pos[i] >= 0.0 && neg[i] <= 0.0, "signs at {i}");
1002            assert!((tot[i] - (pos[i] + neg[i])).abs() < 1e-9, "invariant at {i}");
1003        }
1004        // Separation happened: paramagnetic in GM, diamagnetic in WM.
1005        assert!(pos.iter().any(|&v| v > 1e-3), "some χ+ present");
1006        assert!(neg.iter().any(|&v| v < -1e-3), "some χ− present (WM)");
1007    }
1008
1009    /// All-GM phantom: no magnitude beat → `beat_frac` below threshold → the
1010    /// closed-form branch is returned for every voxel.
1011    #[test]
1012    fn hc_chisep_no_beat_falls_back_to_closed_form() {
1013        let tes = [0.004, 0.012, 0.020, 0.028];
1014        let dims = (10, 10, 10);
1015        let n = dims.0 * dims.1 * dims.2;
1016        let (chi, r2p, mag, _se, mask) = synth_phantom(dims, 0, 7.0, &tes, &[]);
1017        let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
1018        let params = HcChisepParams { b0: 7.0, ..Default::default() };
1019
1020        let (pos, neg, _tot) =
1021            hc_chisep(&chi, &r2p, &mag, &tes, None, &mask, &grid, &params, |_, _| {});
1022        // Closed form on this phantom: χ+ = R2'/Dr+ = 0.05, χ− = 0.
1023        for i in 0..n {
1024            assert!((pos[i] - 0.05).abs() < 1e-6, "closed-form χ+ at {i}: {}", pos[i]);
1025            assert!(neg[i].abs() < 1e-9, "closed-form χ− at {i}: {}", neg[i]);
1026        }
1027    }
1028
1029    /// Empty mask returns all zeros via the early exit.
1030    #[test]
1031    fn hc_chisep_empty_mask_returns_zero() {
1032        let tes = [0.004, 0.012, 0.020, 0.028];
1033        let dims = (6, 6, 6);
1034        let n = dims.0 * dims.1 * dims.2;
1035        let (chi, r2p, mag, _se, _m) = synth_phantom(dims, 0, 7.0, &tes, &[]);
1036        let mask = vec![0u8; n];
1037        let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
1038        let params = HcChisepParams { b0: 7.0, ..Default::default() };
1039
1040        let (pos, neg, tot) =
1041            hc_chisep(&chi, &r2p, &mag, &tes, None, &mask, &grid, &params, |_, _| {});
1042        assert!(pos.iter().all(|&v| v == 0.0), "χ+ all zero");
1043        assert!(neg.iter().all(|&v| v == 0.0), "χ− all zero");
1044        assert!(tot.iter().all(|&v| v == 0.0), "χ_total all zero");
1045    }
1046}