1use super::config::*;
13use crate::separation::{
14 ChiSepIlsqrParams, ChiSepParams, DecomposeParams, HcChisepParams, R2starQsmParams,
15};
16
17const GAMMA_HZ_PER_T: f64 = 42.576e6;
19
20#[derive(Clone, Copy)]
23pub struct SeparationInputs<'a> {
24 pub local_field_ppm: &'a [f64],
26 pub qsm: &'a [f64],
29 pub mask: &'a [u8],
31 pub r2prime: Option<&'a [f64]>,
33 pub r2star: Option<&'a [f64]>,
35 pub magnitude_rss: Option<&'a [f64]>,
37 pub magnitude_multi: Option<&'a [f64]>,
40 pub se_magnitude_multi: Option<&'a [f64]>,
42}
43
44#[derive(Clone, Debug)]
46pub struct SeparationResult {
47 pub chi_pos: Vec<f64>,
49 pub chi_neg: Vec<f64>,
51 pub chi_total: Vec<f64>,
53}
54
55pub 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, ¶ms, |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, ¶ms, |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, ¶ms)
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, ¶ms,
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, ¶ms, |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, ¶ms, |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#[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#[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
206fn 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 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 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}