qsm_core/separation/
susep_net.rs1use crate::grid::Grid;
18use crate::models::onnx::{OnnxModel, OnnxError, Tensor};
19
20#[derive(Clone, Copy, Debug)]
24pub struct SusepNetNorm {
25 pub qsm: (f64, f64),
26 pub lfs: (f64, f64),
27 pub r2prime: (f64, f64),
28 pub chi_pos: (f64, f64),
29 pub chi_neg: (f64, f64),
30}
31
32impl Default for SusepNetNorm {
33 fn default() -> Self {
35 Self {
36 qsm: (-6.0663105e-05, 0.023533047),
37 lfs: (-4.8702253e-05, 0.012554166),
38 r2prime: (4.7629275, 10.889079),
39 chi_pos: (0.0089528897, 0.025519046),
40 chi_neg: (0.0090135528, 0.019603666),
41 }
42 }
43}
44
45#[allow(clippy::too_many_arguments)]
55pub fn susep_net(
56 local_field_ppm: &[f64],
57 qsm: &[f64],
58 r2prime: &[f64],
59 mask: &[u8],
60 grid: &Grid,
61 model_onnx: &[u8],
62 norm: &SusepNetNorm,
63) -> Result<(Vec<f64>, Vec<f64>, Vec<f64>), OnnxError> {
64 let (nx, ny, nz) = grid.dims;
65 let n = nx * ny * nz;
66 for (name, v) in [("field", local_field_ppm), ("qsm", qsm), ("r2prime", r2prime)] {
67 assert_eq!(v.len(), n, "{name} length must match grid");
68 }
69 assert_eq!(mask.len(), n, "mask length must match grid");
70
71 let (px, py, pz) = (nx.div_ceil(8) * 8, ny.div_ceil(8) * 8, nz.div_ceil(8) * 8);
73
74 let pack = |src: &[f64], (mean, std): (f64, f64)| -> Tensor {
76 let inv = 1.0 / std;
77 let mut buf = vec![0.0f32; px * py * pz];
78 for z in 0..nz {
79 for y in 0..ny {
80 for x in 0..nx {
81 let i = x + nx * (y + ny * z);
82 if mask[i] != 0 {
83 let dst = (x * py + y) * pz + z;
84 buf[dst] = ((src[i] - mean) * inv) as f32;
85 }
86 }
87 }
88 }
89 Tensor::new(vec![1, 1, px, py, pz], buf)
90 };
91
92 let inputs = [
94 pack(qsm, norm.qsm),
95 pack(r2prime, norm.r2prime),
96 pack(local_field_ppm, norm.lfs),
97 ];
98 let model = OnnxModel::load(model_onnx)?;
99 let outs = model.run(&inputs)?;
100 if outs.len() < 2 {
101 return Err(OnnxError::Run(format!("expected 2 outputs, got {}", outs.len())));
102 }
103
104 let mut chi_pos = vec![0.0f64; n];
106 let mut chi_neg = vec![0.0f64; n];
107 let mut chi_total = vec![0.0f64; n];
108 let (pm, ps) = norm.chi_pos;
109 let (nm, ns) = norm.chi_neg;
110 for z in 0..nz {
111 for y in 0..ny {
112 for x in 0..nx {
113 let i = x + nx * (y + ny * z);
114 if mask[i] != 0 {
115 let src = (x * py + y) * pz + z;
116 let pos = outs[0].data[src] as f64 * ps + pm;
117 let neg_mag = outs[1].data[src] as f64 * ns + nm;
118 chi_pos[i] = pos;
119 chi_neg[i] = -neg_mag;
120 chi_total[i] = pos - neg_mag;
121 }
122 }
123 }
124 }
125 Ok((chi_pos, chi_neg, chi_total))
126}