Skip to main content

qsm_core/pipeline/
qsmart.rs

1//! QSMART two-stage reconstruction
2//!
3//! Two-stage SDF + iLSQR pipeline with vasculature detection.
4//! Stage 1: whole-ROI reconstruction. Stage 2: tissue-only reconstruction
5//! with vasculature excluded. Final offset adjustment combines both.
6
7use super::config::*;
8
9/// Run QSMART two-stage QSM reconstruction.
10///
11/// # Arguments
12/// * `field_ppm` - Total field in ppm (after field mapping)
13/// * `mask` - Binary brain mask
14/// * `magnitude` - Combined magnitude (for vasculature detection and, when the inner
15///   inversion is MEDI, edge weighting; uniform if None)
16/// * `metadata` - Scan metadata
17/// * `inversion_config` - Inversion configuration. QSMART-specific settings (SDF,
18///   vasculature, iLSQR tolerance) come from `inversion_config.qsmart`; the inner
19///   dipole inversion algorithm is selected by `inversion_config.qsmart.inversion`
20///   and tuned via the matching per-algorithm field on this config.
21/// * `reference` - QSM referencing method
22/// * `progress` - Progress callback (current_step, total_steps)
23///
24/// # Returns
25/// Susceptibility map in ppm (referenced)
26pub fn run_qsmart(
27    field_ppm: &[f64],
28    mask: &[u8],
29    magnitude: Option<&[f64]>,
30    metadata: &ScanMetadata,
31    inversion_config: &InversionConfig,
32    reference: QsmReference,
33    progress: &mut dyn FnMut(usize, usize),
34) -> Result<Vec<f64>, PipelineError> {
35    let (nx, ny, nz) = metadata.dims;
36    let bdir = metadata.b0_direction;
37    let n_voxels = nx * ny * nz;
38    let grid = metadata.grid();
39    let qsmart_params = &inversion_config.qsmart;
40
41    // QSMART runs a standard dipole inversion per stage; the two pipeline-level
42    // algorithms can't be nested here.
43    if matches!(
44        qsmart_params.inversion,
45        InversionAlgorithm::Qsmart | InversionAlgorithm::Tgv
46    ) {
47        return Err(PipelineError::InvalidConfig(
48            "QSMART inversion cannot be Qsmart or Tgv".into(),
49        ));
50    }
51
52    // Inner inversion config: same per-algorithm params as the caller, but with the
53    // QSMART-selected algorithm and iLSQR params pinned to QSMART's own fields so the
54    // default (iLSQR) path is identical to the previous hardcoded call.
55    let mut inner_config = inversion_config.clone();
56    inner_config.algorithm = qsmart_params.inversion;
57    inner_config.ilsqr = crate::inversion::IlsqrParams {
58        tol: qsmart_params.ilsqr_tol,
59        max_iter: qsmart_params.ilsqr_max_iter,
60    };
61
62    // Step 1: Vasculature detection (vasc_mask: 1 = tissue, 0 = vessel)
63    progress(1, 6);
64    let uniform_mag = vec![1.0f64; n_voxels];
65    let mag = magnitude.unwrap_or(&uniform_mag);
66    let vasc_params = crate::utils::VasculatureParams {
67        sphere_radius: qsmart_params.vasc_sphere_radius,
68        frangi_scale_range: qsmart_params.frangi_scale_range,
69        frangi_scale_ratio: qsmart_params.frangi_scale_ratio,
70        frangi_c: qsmart_params.frangi_c,
71    };
72    let vasc_mask = crate::utils::generate_vasculature_mask(
73        mag, mask, &grid, &vasc_params, |_, _| {},
74    );
75
76    // No reliability map at this layer, so the QSMART weighting mask is just the brain mask.
77    let weighted_mask: Vec<f64> = mask.iter().map(|&m| m as f64).collect();
78    let ones_vasc = vec![1.0f64; n_voxels];
79
80    // Step 2: SDF stage 1 — background removal over the whole ROI (vessels included).
81    progress(2, 6);
82    let sdf_params1 = crate::bgremove::SdfParams {
83        sigma1: qsmart_params.sdf_sigma1_stage1,
84        sigma2: qsmart_params.sdf_sigma2_stage1,
85        spatial_radius: qsmart_params.sdf_spatial_radius,
86        lower_lim: qsmart_params.sdf_lower_lim,
87        curv_constant: qsmart_params.sdf_curv_constant,
88        use_curvature: true,
89    };
90    let lfs1 = crate::bgremove::sdf::sdf(
91        field_ppm, &weighted_mask, &ones_vasc, &grid, &sdf_params1, |_, _| {},
92    );
93
94    // Step 3: dipole inversion stage 1 over the whole ROI.
95    progress(3, 6);
96    let mask_stage1: Vec<u8> = weighted_mask.iter().map(|&v| if v > 0.1 { 1 } else { 0 }).collect();
97    let chi1 = super::inversion::run_dipole_inversion(
98        &lfs1, &mask_stage1, metadata, &inner_config, magnitude, &mut |_, _| {},
99    )?;
100
101    // Step 4: SDF stage 2 — tissue-only, vessel-aware (mask-zeroed field input).
102    progress(4, 6);
103    let sdf_params2 = crate::bgremove::SdfParams {
104        sigma1: qsmart_params.sdf_sigma1_stage2,
105        sigma2: qsmart_params.sdf_sigma2_stage2,
106        spatial_radius: qsmart_params.sdf_spatial_radius,
107        lower_lim: qsmart_params.sdf_lower_lim,
108        curv_constant: qsmart_params.sdf_curv_constant,
109        use_curvature: true,
110    };
111    let field_weighted: Vec<f64> = field_ppm.iter()
112        .zip(weighted_mask.iter())
113        .map(|(&f, &m)| f * m)
114        .collect();
115    let lfs2 = crate::bgremove::sdf::sdf(
116        &field_weighted, &weighted_mask, &vasc_mask, &grid, &sdf_params2, |_, _| {},
117    );
118
119    // Step 5: dipole inversion stage 2 over tissue only (in-mask AND not vessel).
120    progress(5, 6);
121    let mask_stage2: Vec<u8> = weighted_mask.iter()
122        .zip(vasc_mask.iter())
123        .map(|(&m, &v)| if m > 0.1 && v > 0.5 { 1 } else { 0 })
124        .collect();
125    let chi2 = super::inversion::run_dipole_inversion(
126        &lfs2, &mask_stage2, metadata, &inner_config, magnitude, &mut |_, _| {},
127    )?;
128
129    // Step 6: offset adjustment (combine stages) and reference.
130    // removed_voxels = vessel regions (in mask but excluded from stage 2).
131    // Everything is in ppm here, so adjust_offset's lfs rescale is the identity (ppm = 1.0).
132    progress(6, 6);
133    let removed_voxels: Vec<f64> = weighted_mask.iter()
134        .zip(vasc_mask.iter())
135        .map(|(&m, &v)| m - v)
136        .collect();
137    let mut chi = crate::utils::adjust_offset(
138        &removed_voxels, &lfs1, &chi1, &chi2, &grid, bdir, 1.0,
139    );
140    // QSMART susceptibility is only defined inside the brain mask.
141    for (c, &m) in chi.iter_mut().zip(mask.iter()) {
142        if m == 0 {
143            *c = 0.0;
144        }
145    }
146
147    Ok(super::referencing::apply_reference(&chi, mask, reference))
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153
154    #[test]
155    fn test_run_qsmart_basic() {
156        let (nx, ny, nz) = (8, 8, 8);
157        let n = nx * ny * nz;
158        let field = vec![0.01; n];
159        let mask = vec![1u8; n];
160        let meta = ScanMetadata {
161            dims: (nx, ny, nz),
162            voxel_size: (1.0, 1.0, 1.0),
163            echo_times: vec![0.005],
164            field_strength: 3.0,
165            b0_direction: (0.0, 0.0, 1.0),
166        };
167        let config = InversionConfig {
168            algorithm: InversionAlgorithm::Qsmart,
169            qsmart: crate::utils::QsmartParams::for_field_strength(3.0),
170            ..Default::default()
171        };
172
173        let result = run_qsmart(
174            &field, &mask, None, &meta, &config, QsmReference::Mean,
175            &mut |_, _| {},
176        );
177        assert!(result.is_ok());
178        let chi = result.unwrap();
179        assert_eq!(chi.len(), n);
180        for &v in &chi {
181            assert!(v.is_finite(), "QSMART output must be finite");
182        }
183    }
184
185    #[test]
186    fn test_run_qsmart_swappable_inversion() {
187        // QSMART with a non-default inner inversion (TKD) should run end to end.
188        let (nx, ny, nz) = (8, 8, 8);
189        let n = nx * ny * nz;
190        let field = vec![0.01; n];
191        let mask = vec![1u8; n];
192        let meta = ScanMetadata {
193            dims: (nx, ny, nz),
194            voxel_size: (1.0, 1.0, 1.0),
195            echo_times: vec![0.005],
196            field_strength: 3.0,
197            b0_direction: (0.0, 0.0, 1.0),
198        };
199        let mut qsmart = crate::utils::QsmartParams::for_field_strength(3.0);
200        qsmart.inversion = InversionAlgorithm::Tkd;
201        let config = InversionConfig {
202            algorithm: InversionAlgorithm::Qsmart,
203            qsmart,
204            ..Default::default()
205        };
206
207        let chi = run_qsmart(
208            &field, &mask, None, &meta, &config, QsmReference::Mean,
209            &mut |_, _| {},
210        )
211        .unwrap();
212        assert_eq!(chi.len(), n);
213        for &v in &chi {
214            assert!(v.is_finite(), "QSMART output must be finite");
215        }
216    }
217
218    #[test]
219    fn test_run_qsmart_rejects_nested_pipeline_inversion() {
220        let (nx, ny, nz) = (4, 4, 4);
221        let n = nx * ny * nz;
222        let meta = ScanMetadata {
223            dims: (nx, ny, nz),
224            voxel_size: (1.0, 1.0, 1.0),
225            echo_times: vec![0.005],
226            field_strength: 3.0,
227            b0_direction: (0.0, 0.0, 1.0),
228        };
229        let mut qsmart = crate::utils::QsmartParams::for_field_strength(3.0);
230        qsmart.inversion = InversionAlgorithm::Qsmart;
231        let config = InversionConfig {
232            algorithm: InversionAlgorithm::Qsmart,
233            qsmart,
234            ..Default::default()
235        };
236
237        let result = run_qsmart(
238            &vec![0.0; n], &vec![1u8; n], None, &meta, &config,
239            QsmReference::Mean, &mut |_, _| {},
240        );
241        assert!(result.is_err());
242    }
243}