Skip to main content

qsm_core/pipeline/
separation.rs

1//! Susceptibility source-separation stage.
2//!
3//! Config-driven dispatcher over the χ-separation methods (χ+ / χ−). Unlike the
4//! main phase→field→BFR→inversion pipeline, separation is a downstream stage that
5//! needs extra inputs (a conventional QSM plus a relaxometry map / multi-echo
6//! magnitude), so it is invoked explicitly with a [`SeparationInputs`] bundle
7//! rather than folded into the top-level runner.
8//!
9//! Units follow the library convention: local field & QSM in ppm, R2'/R2* in Hz,
10//! echo times in seconds. Output is `(chi_pos ≥ 0, chi_neg ≤ 0, chi_total)` in ppm.
11
12use super::config::*;
13use crate::separation::{
14    ChiSepIlsqrParams, ChiSepParams, DecomposeParams, HcChisepParams, R2starQsmParams,
15};
16
17/// Proton gyromagnetic ratio in Hz/T (for the central frequency `cf = γ·B0`).
18const GAMMA_HZ_PER_T: f64 = 42.576e6;
19
20/// Inputs for [`run_separation`]. Provide whatever the chosen algorithm needs;
21/// the dispatcher errors if a required input is missing.
22#[derive(Clone, Copy)]
23pub struct SeparationInputs<'a> {
24    /// Local (tissue) field in ppm — for `ChiSepIlsqr` / `ChiSepMedi`.
25    pub local_field_ppm: &'a [f64],
26    /// Conventional QSM χ_total in ppm — the QSM-init for `ChiSepIlsqr` and the
27    /// input for `R2starQsm` / `WaveSep` / `Decompose` / `HcChisep`.
28    pub qsm: &'a [f64],
29    /// Binary brain mask (`n_voxels`, 1 = inside).
30    pub mask: &'a [u8],
31    /// R2' map in Hz (`ChiSepIlsqr` / `ChiSepMedi` / `WaveSep` / `HcChisep`).
32    pub r2prime: Option<&'a [f64]>,
33    /// R2* map in Hz (`R2starQsm`); if absent, fit from `magnitude_multi`.
34    pub r2star: Option<&'a [f64]>,
35    /// Root-sum-of-squares magnitude (`ChiSepIlsqr` / `ChiSepMedi`).
36    pub magnitude_rss: Option<&'a [f64]>,
37    /// Multi-echo magnitude, voxel-major `(n_voxels, n_echoes)`
38    /// (`R2starQsm` / `Decompose` / `HcChisep`).
39    pub magnitude_multi: Option<&'a [f64]>,
40    /// Multi-echo spin-echo magnitude, voxel-major `(n_voxels, n_se)` (`HcChisep`, optional).
41    pub se_magnitude_multi: Option<&'a [f64]>,
42}
43
44/// Result of [`run_separation`]: paramagnetic / diamagnetic / total maps in ppm.
45#[derive(Clone, Debug)]
46pub struct SeparationResult {
47    /// χ+ ≥ 0 (paramagnetic).
48    pub chi_pos: Vec<f64>,
49    /// χ− ≤ 0 (diamagnetic, signed).
50    pub chi_neg: Vec<f64>,
51    /// χ_total = χ+ + χ−.
52    pub chi_total: Vec<f64>,
53}
54
55/// Run the configured χ-separation algorithm.
56///
57/// # Arguments
58/// * `inputs` — See [`SeparationInputs`]; provide what the algorithm requires.
59/// * `metadata` — Scan metadata (grid, B0 direction, field strength, echo times).
60/// * `config` — See [`SeparationConfig`].
61/// * `progress` — Progress callback `(current_iter, max_iter)`.
62pub fn run_separation(
63    inputs: SeparationInputs,
64    metadata: &ScanMetadata,
65    config: &SeparationConfig,
66    progress: &mut dyn FnMut(usize, usize),
67) -> Result<SeparationResult, PipelineError> {
68    let grid = metadata.grid();
69    let n = grid.n_total();
70    let bdir = metadata.b0_direction;
71    let cf = GAMMA_HZ_PER_T * metadata.field_strength;
72    let b0 = metadata.field_strength;
73
74    if inputs.qsm.len() != n || inputs.mask.len() != n {
75        return Err(PipelineError::DimensionMismatch { expected: n, got: inputs.qsm.len() });
76    }
77    let mask = inputs.mask;
78
79    let (chi_pos, chi_neg, chi_total) = match config.algorithm {
80        SeparationAlgorithm::ChiSepIlsqr => {
81            let r2prime = need(inputs.r2prime, "r2prime")?;
82            let magnitude = need(inputs.magnitude_rss, "magnitude_rss")?;
83            let params = ChiSepIlsqrParams { cf, ..config.chi_sep_ilsqr.clone() };
84            crate::separation::chi_sep_ilsqr(
85                inputs.local_field_ppm, r2prime, magnitude, inputs.qsm, mask,
86                &grid, bdir, &params, |i, k| progress(i, k),
87            )
88        }
89        SeparationAlgorithm::ChiSepMedi => {
90            let r2prime = need(inputs.r2prime, "r2prime")?;
91            let magnitude = need(inputs.magnitude_rss, "magnitude_rss")?;
92            let params = ChiSepParams { cf, ..config.chi_sep_medi.clone() };
93            crate::separation::chi_sep_medi(
94                inputs.local_field_ppm, r2prime, magnitude, mask,
95                &grid, bdir, &params, |i, k| progress(i, k),
96            )
97        }
98        SeparationAlgorithm::R2starQsm => {
99            let params = R2starQsmParams { b0, ..config.r2star_qsm.clone() };
100            if let Some(r2star) = inputs.r2star {
101                crate::separation::r2star_qsm(inputs.qsm, r2star, mask, &params)
102            } else {
103                let magnitude = need(inputs.magnitude_multi, "r2star or magnitude_multi")?;
104                crate::separation::r2star_qsm_from_magnitude(
105                    magnitude, &metadata.echo_times, inputs.qsm, mask, &params,
106                )
107            }
108        }
109        SeparationAlgorithm::WaveSep => {
110            let r2prime = need(inputs.r2prime, "r2prime")?;
111            crate::separation::wavesep(
112                inputs.qsm, r2prime, mask, &grid, &config.wavesep, |i, k| progress(i, k),
113            )
114        }
115        SeparationAlgorithm::Decompose => {
116            let magnitude = need(inputs.magnitude_multi, "magnitude_multi")?;
117            let params = DecomposeParams { b0, ..config.decompose.clone() };
118            crate::separation::decompose(
119                inputs.qsm, magnitude, &metadata.echo_times, mask, &params, |i, k| progress(i, k),
120            )
121        }
122        SeparationAlgorithm::HcChisep => {
123            let r2prime = need(inputs.r2prime, "r2prime")?;
124            let magnitude = need(inputs.magnitude_multi, "magnitude_multi")?;
125            let params = HcChisepParams { b0, ..config.hc_chisep.clone() };
126            crate::separation::hc_chisep(
127                inputs.qsm, r2prime, magnitude, &metadata.echo_times,
128                inputs.se_magnitude_multi, mask, &grid, &params, |i, k| progress(i, k),
129            )
130        }
131        SeparationAlgorithm::SusepNet => {
132            let r2prime = need(inputs.r2prime, "r2prime")?;
133            run_susep_net(inputs.local_field_ppm, inputs.qsm, r2prime, mask, &grid)?
134        }
135        SeparationAlgorithm::ChiSepNet => {
136            let r2prime = need(inputs.r2prime, "r2prime")?;
137            run_chi_sepnet(inputs.local_field_ppm, inputs.qsm, r2prime, mask, &grid)?
138        }
139    };
140
141    Ok(SeparationResult { chi_pos, chi_neg, chi_total })
142}
143
144/// Source the SUSEP-Net weights and run inference (requires the `onnx` feature).
145#[cfg(feature = "onnx")]
146fn run_susep_net(
147    local_field_ppm: &[f64],
148    qsm: &[f64],
149    r2prime: &[f64],
150    mask: &[u8],
151    grid: &crate::Grid,
152) -> Result<(Vec<f64>, Vec<f64>, Vec<f64>), PipelineError> {
153    let spec = crate::models::find_model("susep-net")
154        .ok_or_else(|| PipelineError::InvalidConfig("susep-net not in model registry".into()))?;
155    let bytes = crate::models::primary_weight_bytes(spec).map_err(PipelineError::InvalidConfig)?;
156    crate::separation::susep_net(
157        local_field_ppm, qsm, r2prime, mask, grid, &bytes,
158        &crate::separation::SusepNetNorm::default(),
159    )
160    .map_err(|e| PipelineError::AlgorithmError(e.to_string()))
161}
162
163#[cfg(not(feature = "onnx"))]
164fn run_susep_net(
165    _local_field_ppm: &[f64],
166    _qsm: &[f64],
167    _r2prime: &[f64],
168    _mask: &[u8],
169    _grid: &crate::Grid,
170) -> Result<(Vec<f64>, Vec<f64>, Vec<f64>), PipelineError> {
171    Err(PipelineError::InvalidConfig(
172        "SUSEP-Net requires building qsm-core with the 'onnx' feature".into(),
173    ))
174}
175
176/// Source the χ-sepnet weights and run inference (requires the `onnx` feature).
177#[cfg(feature = "onnx")]
178fn run_chi_sepnet(
179    local_field_ppm: &[f64],
180    qsm: &[f64],
181    r2prime: &[f64],
182    mask: &[u8],
183    grid: &crate::Grid,
184) -> Result<(Vec<f64>, Vec<f64>, Vec<f64>), PipelineError> {
185    let bytes = crate::models::primary_weight("chi-sepnet").map_err(PipelineError::InvalidConfig)?;
186    crate::separation::chisepnet(
187        local_field_ppm, qsm, r2prime, mask, grid, &bytes,
188        &crate::separation::ChiSepNetNorm::default(),
189    )
190    .map_err(|e| PipelineError::AlgorithmError(e.to_string()))
191}
192
193#[cfg(not(feature = "onnx"))]
194fn run_chi_sepnet(
195    _local_field_ppm: &[f64],
196    _qsm: &[f64],
197    _r2prime: &[f64],
198    _mask: &[u8],
199    _grid: &crate::Grid,
200) -> Result<(Vec<f64>, Vec<f64>, Vec<f64>), PipelineError> {
201    Err(PipelineError::InvalidConfig(
202        "χ-sepnet requires building qsm-core with the 'onnx' feature".into(),
203    ))
204}
205
206/// Require an optional input, or return an `InvalidInput` error naming it.
207fn need<'a>(opt: Option<&'a [f64]>, what: &str) -> Result<&'a [f64], PipelineError> {
208    opt.ok_or_else(|| PipelineError::InvalidInput(format!("{what} required for this algorithm")))
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214
215    fn meta(nx: usize, ny: usize, nz: usize) -> ScanMetadata {
216        ScanMetadata {
217            dims: (nx, ny, nz),
218            voxel_size: (1.0, 1.0, 1.0),
219            echo_times: vec![0.004, 0.012, 0.020, 0.028],
220            field_strength: 3.0,
221            b0_direction: (0.0, 0.0, 1.0),
222        }
223    }
224
225    fn base_inputs<'a>(qsm: &'a [f64], mask: &'a [u8]) -> SeparationInputs<'a> {
226        SeparationInputs {
227            local_field_ppm: qsm,
228            qsm,
229            mask,
230            r2prime: None,
231            r2star: None,
232            magnitude_rss: None,
233            magnitude_multi: None,
234            se_magnitude_multi: None,
235        }
236    }
237
238    #[test]
239    fn r2star_qsm_dispatch_closed_form() {
240        let (nx, ny, nz) = (6, 6, 6);
241        let n = nx * ny * nz;
242        let qsm = vec![0.02; n];
243        let r2star = vec![10.0; n];
244        let mask = vec![1u8; n];
245        let m = meta(nx, ny, nz);
246        let mut inputs = base_inputs(&qsm, &mask);
247        inputs.r2star = Some(&r2star);
248        let cfg = SeparationConfig { algorithm: SeparationAlgorithm::R2starQsm, ..Default::default() };
249        let r = run_separation(inputs, &m, &cfg, &mut |_, _| {}).unwrap();
250        assert_eq!(r.chi_pos.len(), n);
251        for i in 0..n {
252            assert!(r.chi_pos[i] >= 0.0 && r.chi_neg[i] <= 0.0);
253            assert!((r.chi_total[i] - (r.chi_pos[i] + r.chi_neg[i])).abs() < 1e-9);
254        }
255    }
256
257    #[test]
258    fn wavesep_dispatch_runs() {
259        let (nx, ny, nz) = (8, 8, 8);
260        let n = nx * ny * nz;
261        let qsm = vec![0.01; n];
262        let r2prime = vec![2.0; n];
263        let mask = vec![1u8; n];
264        let m = meta(nx, ny, nz);
265        let mut inputs = base_inputs(&qsm, &mask);
266        inputs.r2prime = Some(&r2prime);
267        let cfg = SeparationConfig { algorithm: SeparationAlgorithm::WaveSep, ..Default::default() };
268        let r = run_separation(inputs, &m, &cfg, &mut |_, _| {}).unwrap();
269        assert_eq!(r.chi_pos.len(), n);
270    }
271
272    #[test]
273    fn chi_sep_ilsqr_and_medi_dispatch() {
274        let (nx, ny, nz) = (6, 6, 6);
275        let n = nx * ny * nz;
276        let field = vec![0.01; n];
277        let r2prime = vec![2.0; n];
278        let mag = vec![1.0; n];
279        let mask = vec![1u8; n];
280        let m = meta(nx, ny, nz);
281        for alg in [SeparationAlgorithm::ChiSepIlsqr, SeparationAlgorithm::ChiSepMedi] {
282            let mut inputs = base_inputs(&field, &mask);
283            inputs.r2prime = Some(&r2prime);
284            inputs.magnitude_rss = Some(&mag);
285            let cfg = SeparationConfig { algorithm: alg, ..Default::default() };
286            let r = run_separation(inputs, &m, &cfg, &mut |_, _| {}).unwrap();
287            assert_eq!(r.chi_pos.len(), n);
288        }
289    }
290
291    #[test]
292    fn decompose_and_hc_dispatch() {
293        let (nx, ny, nz) = (6, 6, 6);
294        let n = nx * ny * nz;
295        let ne = 4;
296        let qsm = vec![0.01; n];
297        let r2prime = vec![2.0; n];
298        let mag_multi = vec![1.0; n * ne];
299        let mask = vec![1u8; n];
300        let m = meta(nx, ny, nz);
301
302        let mut d_inputs = base_inputs(&qsm, &mask);
303        d_inputs.magnitude_multi = Some(&mag_multi);
304        let d_cfg = SeparationConfig { algorithm: SeparationAlgorithm::Decompose, ..Default::default() };
305        let rd = run_separation(d_inputs, &m, &d_cfg, &mut |_, _| {}).unwrap();
306        assert_eq!(rd.chi_pos.len(), n);
307
308        let mut h_inputs = base_inputs(&qsm, &mask);
309        h_inputs.r2prime = Some(&r2prime);
310        h_inputs.magnitude_multi = Some(&mag_multi);
311        let h_cfg = SeparationConfig { algorithm: SeparationAlgorithm::HcChisep, ..Default::default() };
312        let rh = run_separation(h_inputs, &m, &h_cfg, &mut |_, _| {}).unwrap();
313        assert_eq!(rh.chi_pos.len(), n);
314    }
315
316    #[test]
317    fn r2star_from_magnitude_path() {
318        let (nx, ny, nz) = (6, 6, 6);
319        let n = nx * ny * nz;
320        let ne = 4;
321        let qsm = vec![0.02; n];
322        let mag_multi = vec![1.0; n * ne];
323        let mask = vec![1u8; n];
324        let m = meta(nx, ny, nz);
325        // No r2star provided → fit from magnitude_multi.
326        let mut inputs = base_inputs(&qsm, &mask);
327        inputs.magnitude_multi = Some(&mag_multi);
328        let cfg = SeparationConfig { algorithm: SeparationAlgorithm::R2starQsm, ..Default::default() };
329        let r = run_separation(inputs, &m, &cfg, &mut |_, _| {}).unwrap();
330        assert_eq!(r.chi_pos.len(), n);
331    }
332
333    #[test]
334    fn missing_input_errors() {
335        let (nx, ny, nz) = (6, 6, 6);
336        let n = nx * ny * nz;
337        let qsm = vec![0.02; n];
338        let mask = vec![1u8; n];
339        let m = meta(nx, ny, nz);
340        // R2starQsm with neither r2star nor magnitude_multi → InvalidInput.
341        let inputs = base_inputs(&qsm, &mask);
342        let cfg = SeparationConfig { algorithm: SeparationAlgorithm::R2starQsm, ..Default::default() };
343        let err = run_separation(inputs, &m, &cfg, &mut |_, _| {}).unwrap_err();
344        assert!(matches!(err, PipelineError::InvalidInput(_)), "got {err:?}");
345    }
346}