1use crate::inversion::admm::prepare_fansi_spectral;
23use crate::utils::gradient::{bdiv_inplace, fgrad_inplace};
24use crate::utils::{apply_mask_zero, shrink};
25use crate::Grid;
26use num_complex::Complex64;
27
28#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
30#[derive(Clone, Debug)]
31pub struct L1QsmParams {
32 pub alpha1: f64,
34 pub mu1: f64,
36 pub mu2: f64,
38 pub mu3: f64,
40 pub lambda: f64,
42 pub max_iter: usize,
44 pub tol_update: f64,
46 pub tol_delta: f64,
48 pub phase_scale: f64,
50}
51
52impl Default for L1QsmParams {
53 fn default() -> Self {
54 Self {
55 alpha1: 2e-4,
56 mu1: 2e-2,
57 mu2: 1.0,
58 mu3: 1.0,
59 lambda: 1.0,
60 max_iter: 50,
61 tol_update: 1.0,
62 tol_delta: 1e-6,
63 phase_scale: 1.0,
64 }
65 }
66}
67
68fn norm2(v: &[f64]) -> f64 {
70 v.iter().map(|&a| a * a).sum::<f64>().sqrt()
71}
72
73pub fn l1qsm(
86 local_field: &[f64],
87 mask: &[u8],
88 grid: &Grid,
89 bdir: (f64, f64, f64),
90 params: &L1QsmParams,
91 mut progress: impl FnMut(usize, usize),
92) -> Vec<f64> {
93 let n = grid.n_total();
94
95 let (mut fft_ws, k, ee2) = prepare_fansi_spectral(grid, bdir);
96
97 let phase: Vec<f64> = local_field.iter().map(|&f| f * params.phase_scale).collect();
99 let w: Vec<f64> = mask
100 .iter()
101 .map(|&m| if m != 0 { params.lambda } else { 0.0 })
102 .collect();
103
104 let is: Vec<Complex64> = phase
106 .iter()
107 .map(|&p| Complex64::new(p.cos(), p.sin()))
108 .collect();
109
110 let mu1 = params.mu1;
111 let mu2 = params.mu2;
112 let mu3 = params.mu3;
113 let alpha_over_mu = params.alpha1 / mu1;
114
115 let mut x = vec![0.0f64; n];
117 let mut x_prev = vec![0.0f64; n];
118
119 let mut z_dx = vec![0.0f64; n];
121 let mut z_dy = vec![0.0f64; n];
122 let mut z_dz = vec![0.0f64; n];
123 let mut s_dx = vec![0.0f64; n];
124 let mut s_dy = vec![0.0f64; n];
125 let mut s_dz = vec![0.0f64; n];
126
127 let w_max = w.iter().cloned().fold(0.0f64, f64::max);
130 let mut z2 = vec![0.0f64; n];
131 if w_max > 0.0 {
132 for i in 0..n {
133 z2[i] = w[i] * phase[i] / w_max;
134 }
135 }
136 let mut s2 = vec![0.0f64; n];
137
138 let mut z3 = vec![Complex64::new(0.0, 0.0); n];
140 let mut s3 = vec![Complex64::new(0.0, 0.0); n];
141
142 let mut fdiv = vec![Complex64::new(0.0, 0.0); n];
144 let mut fd2 = vec![Complex64::new(0.0, 0.0); n];
145 let mut xhat = vec![Complex64::new(0.0, 0.0); n];
146 let mut fx = vec![Complex64::new(0.0, 0.0); n];
147
148 let mut gxc = vec![0.0f64; n];
149 let mut gyc = vec![0.0f64; n];
150 let mut gzc = vec![0.0f64; n];
151 let mut x_dx = vec![0.0f64; n];
152 let mut x_dy = vec![0.0f64; n];
153 let mut x_dz = vec![0.0f64; n];
154 let mut div = vec![0.0f64; n];
155 let mut dx = vec![0.0f64; n]; let mut rhs_z2 = vec![0.0f64; n];
157 let mut diff = vec![0.0f64; n];
158
159 for t in 0..params.max_iter {
160 progress(t + 1, params.max_iter);
161
162 for i in 0..n {
165 gxc[i] = z_dx[i] - s_dx[i];
166 gyc[i] = z_dy[i] - s_dy[i];
167 gzc[i] = z_dz[i] - s_dz[i];
168 }
169 bdiv_inplace(&mut div, &gxc, &gyc, &gzc, grid);
170 for i in 0..n {
171 fdiv[i] = Complex64::new(div[i], 0.0);
172 }
173 fft_ws.fft3d(&mut fdiv);
174
175 for i in 0..n {
178 fd2[i] = Complex64::new(z2[i] - s2[i], 0.0);
179 }
180 fft_ws.fft3d(&mut fd2);
181
182 for i in 0..n {
184 let num = -mu1 * fdiv[i] + mu2 * k[i] * fd2[i];
187 let den = mu2 * k[i] * k[i] + mu1 * ee2[i];
191 xhat[i] = if den > 1e-20 { num / den } else { Complex64::new(0.0, 0.0) };
192 }
193 fft_ws.ifft3d(&mut xhat);
194 x_prev.copy_from_slice(&x);
195 for i in 0..n {
196 x[i] = xhat[i].re;
197 }
198
199 let xnorm = norm2(&x);
201 if xnorm > 0.0 {
202 for i in 0..n {
203 diff[i] = x[i] - x_prev[i];
204 }
205 let x_update = 100.0 * norm2(&diff) / xnorm;
206 if x_update < params.tol_update || x_update.is_nan() {
207 progress(t + 1, t + 1);
208 break;
209 }
210 }
211
212 if t + 1 >= params.max_iter {
213 break;
214 }
215
216 for i in 0..n {
219 fx[i] = Complex64::new(x[i], 0.0);
220 }
221 fft_ws.fft3d(&mut fx);
222
223 fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
225 for i in 0..n {
226 z_dx[i] = shrink(x_dx[i] + s_dx[i], alpha_over_mu);
227 z_dy[i] = shrink(x_dy[i] + s_dy[i], alpha_over_mu);
228 z_dz[i] = shrink(x_dz[i] + s_dz[i], alpha_over_mu);
229 s_dx[i] += x_dx[i] - z_dx[i];
230 s_dy[i] += x_dy[i] - z_dy[i];
231 s_dz[i] += x_dz[i] - z_dz[i];
232 }
233
234 for i in 0..n {
237 let ez2 = Complex64::new(z2[i].cos(), z2[i].sin());
238 let y3 = ez2 - is[i] + s3[i];
239 let mag = y3.norm();
240 let thr = w[i] / (mu3 + f64::EPSILON);
241 let shr = (mag - thr).max(0.0);
242 z3[i] = if mag > 0.0 {
243 y3 * (shr / mag)
244 } else {
245 Complex64::new(0.0, 0.0)
246 };
247 }
248
249 for i in 0..n {
252 xhat[i] = fx[i] * k[i];
253 }
254 fft_ws.ifft3d(&mut xhat);
255 for i in 0..n {
256 dx[i] = xhat[i].re;
257 rhs_z2[i] = mu2 * (dx[i] + s2[i]);
258 z2[i] = rhs_z2[i] / mu2;
260 }
261
262 let mut yphase = vec![0.0f64; n];
266 let mut cosh_b = vec![0.0f64; n];
267 let mut sinh_b = vec![0.0f64; n];
268 for i in 0..n {
269 let yc = is[i] + z3[i] - s3[i];
270 yphase[i] = yc.arg();
271 let m = yc.norm();
272 if m > 0.0 {
273 cosh_b[i] = 0.5 * (m + 1.0 / m); sinh_b[i] = 0.5 * (m - 1.0 / m); } else {
276 cosh_b[i] = 1.0;
277 sinh_b[i] = 0.0;
278 }
279 }
280
281 let mut delta = f64::INFINITY;
282 let mut inn = 0usize;
283 let mut update = vec![0.0f64; n];
284 while delta > params.tol_delta && inn < 4 {
285 inn += 1;
286 let norm_old = norm2(&z2);
287
288 for i in 0..n {
289 let a = z2[i] - yphase[i];
290 let (sa, ca) = a.sin_cos();
291 let sin_arg = Complex64::new(sa * cosh_b[i], -ca * sinh_b[i]);
293 let cos_arg = Complex64::new(ca * cosh_b[i], sa * sinh_b[i]);
295
296 let temp = mu3 * cos_arg + Complex64::new(mu2 + f64::EPSILON, 0.0);
297 let numer = mu3 * sin_arg + Complex64::new(mu2 * z2[i] - rhs_z2[i], 0.0);
298
299 let tm = temp.norm();
301 let denom = if tm < 0.05 {
302 if tm > 0.0 {
303 temp * (0.05 / tm)
304 } else {
305 Complex64::new(0.05, 0.0)
306 }
307 } else {
308 temp
309 };
310
311 update[i] = (numer / denom).re;
312 z2[i] -= update[i];
313 }
314
315 let delta_new = if norm_old > 0.0 {
316 norm2(&update) / norm_old
317 } else {
318 0.0
319 };
320 if delta_new > delta {
321 break;
322 }
323 delta = delta_new;
324 }
325
326 for i in 0..n {
329 s2[i] += dx[i] - z2[i];
330 }
331 for i in 0..n {
333 let ez2 = Complex64::new(z2[i].cos(), z2[i].sin());
334 s3[i] = ez2 - is[i] + s3[i] - z3[i];
335 }
336 }
337
338 if params.phase_scale != 1.0 {
340 for v in &mut x {
341 *v /= params.phase_scale;
342 }
343 }
344
345 apply_mask_zero(&mut x, mask);
346 x
347}
348
349#[cfg(test)]
350mod tests {
351 use super::*;
352
353 #[test]
354 fn test_l1qsm_zero_field() {
355 let n = 8;
357 let field = vec![0.0; n * n * n];
358 let mask = vec![1u8; n * n * n];
359 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
360 let params = L1QsmParams {
361 max_iter: 10,
362 ..L1QsmParams::default()
363 };
364
365 let chi = l1qsm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
366
367 for &val in chi.iter() {
368 assert!(val.abs() < 1e-6, "Zero field should give ~zero chi, got {}", val);
369 }
370 }
371
372 #[test]
373 fn test_l1qsm_finite() {
374 let n = 8;
376 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
377 let mask = vec![1u8; n * n * n];
378 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
379 let params = L1QsmParams {
380 max_iter: 10,
381 ..L1QsmParams::default()
382 };
383
384 let chi = l1qsm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
385
386 for (i, &val) in chi.iter().enumerate() {
387 assert!(val.is_finite(), "Chi should be finite at index {}", i);
388 }
389 }
390}