1use num_complex::Complex64;
22use crate::inversion::admm::prepare_fansi_spectral;
23use crate::utils::gradient::{bdiv_inplace, fgrad_inplace};
24use crate::utils::{apply_mask_zero, shrink};
25use crate::Grid;
26
27#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
29#[derive(Clone, Debug)]
30pub struct HdQsmParams {
31 pub alpha_l2: f64,
33 pub mu1_l2: f64,
41 pub mu2: f64,
43 pub max_iter_l1: usize,
45 pub max_iter_l2: usize,
47 pub tol_update: f64,
49 }
51
52impl Default for HdQsmParams {
53 fn default() -> Self {
54 Self {
55 alpha_l2: 1e-4,
56 mu1_l2: 1e-2,
57 mu2: 1.0,
58 max_iter_l1: 20,
59 max_iter_l2: 80,
60 tol_update: 1.0,
61 }
62 }
63}
64
65#[inline]
67fn norm2(v: &[f64]) -> f64 {
68 v.iter().map(|&x| x * x).sum::<f64>().sqrt()
69}
70
71#[inline]
73fn percent_update(x: &[f64], x_prev: &[f64]) -> f64 {
74 let diff: f64 = x.iter().zip(x_prev).map(|(&a, &b)| (a - b) * (a - b)).sum::<f64>().sqrt();
75 let nx = norm2(x);
76 if nx > 0.0 { 100.0 * diff / nx } else { 0.0 }
77}
78
79fn apply_dipole(
83 fft_ws: &mut crate::fft::Fft3dWorkspace,
84 k: &[f64],
85 x: &[f64],
86 cbuf: &mut [Complex64],
87 out: &mut [f64],
88) {
89 for (c, &xv) in cbuf.iter_mut().zip(x.iter()) {
90 *c = Complex64::new(xv, 0.0);
91 }
92 fft_ws.fft3d(cbuf);
93 for (c, &kv) in cbuf.iter_mut().zip(k.iter()) {
94 *c *= kv;
95 }
96 fft_ws.ifft3d(cbuf);
97 for (o, c) in out.iter_mut().zip(cbuf.iter()) {
98 *o = c.re;
99 }
100}
101
102pub fn hdqsm(
115 local_field: &[f64],
116 mask: &[u8],
117 grid: &Grid,
118 bdir: (f64, f64, f64),
119 params: &HdQsmParams,
120 mut progress: impl FnMut(usize, usize),
121) -> Vec<f64> {
122 let n = grid.n_total();
123 let (mut fft_ws, k, ee2) = prepare_fansi_spectral(grid, bdir);
124 let denom_base: Vec<f64> = k.iter()
126 .map(|&kk| 1e-30 + params.mu2 * kk * kk)
127 .collect();
128
129 let weight_mask: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
131
132 let total_iter = params.max_iter_l1 + params.max_iter_l2;
133 let mut global_iter = 0usize;
134
135 let mut x = vec![0.0f64; n];
137 let mut x_prev = vec![0.0f64; n];
138
139 let mut z_dx = vec![0.0f64; n];
140 let mut z_dy = vec![0.0f64; n];
141 let mut z_dz = vec![0.0f64; n];
142 let mut s_dx = vec![0.0f64; n];
143 let mut s_dy = vec![0.0f64; n];
144 let mut s_dz = vec![0.0f64; n];
145 let mut x_dx = vec![0.0f64; n];
146 let mut x_dy = vec![0.0f64; n];
147 let mut x_dz = vec![0.0f64; n];
148
149 let mut z2 = vec![0.0f64; n];
150 let mut s2 = vec![0.0f64; n];
151 let mut dx = vec![0.0f64; n]; let mut div = vec![0.0f64; n];
154 let mut gx = vec![0.0f64; n];
155 let mut gy = vec![0.0f64; n];
156 let mut gz = vec![0.0f64; n];
157
158 let mut cbuf = vec![Complex64::new(0.0, 0.0); n];
159 let mut fdiv = vec![Complex64::new(0.0, 0.0); n];
160 let mut fdt = vec![Complex64::new(0.0, 0.0); n];
161
162 {
169 let mu = params.mu1_l2.sqrt();
170 let alpha = params.alpha_l2.sqrt();
171 let mu2 = params.mu2;
172 let ll = alpha / mu;
173
174 let denom: Vec<f64> = (0..n).map(|i| denom_base[i] + mu * ee2[i]).collect();
176
177 let wy = local_field;
179
180 for t in 0..params.max_iter_l1 {
181 x_prev.copy_from_slice(&x);
182
183 for i in 0..n {
185 gx[i] = z2[i] - s2[i] + wy[i]; }
187 for (c, &v) in fdt.iter_mut().zip(gx.iter()) {
189 *c = Complex64::new(v, 0.0);
190 }
191 fft_ws.fft3d(&mut fdt);
192
193 for i in 0..n {
195 gx[i] = z_dx[i] - s_dx[i];
196 gy[i] = z_dy[i] - s_dy[i];
197 gz[i] = z_dz[i] - s_dz[i];
198 }
199 bdiv_inplace(&mut div, &gx, &gy, &gz, grid);
200 for (c, &v) in fdiv.iter_mut().zip(div.iter()) {
201 *c = Complex64::new(v, 0.0);
202 }
203 fft_ws.fft3d(&mut fdiv);
204
205 for i in 0..n {
210 let num = -fdiv[i] * mu + fdt[i] * (mu2 * k[i]);
213 cbuf[i] = if denom[i] > 1e-20 { num / denom[i] } else { Complex64::new(0.0, 0.0) };
214 }
215 fft_ws.ifft3d(&mut cbuf);
216 for i in 0..n {
217 x[i] = cbuf[i].re;
218 }
219
220 global_iter += 1;
221 progress(global_iter, total_iter);
222
223 if t < params.max_iter_l1 - 1 {
224 fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
226 for i in 0..n {
228 let ax = x_dx[i] + s_dx[i];
229 let ay = x_dy[i] + s_dy[i];
230 let az = x_dz[i] + s_dz[i];
231 z_dx[i] = shrink(ax, ll);
232 z_dy[i] = shrink(ay, ll);
233 z_dz[i] = shrink(az, ll);
234 s_dx[i] += x_dx[i] - z_dx[i];
235 s_dy[i] += x_dy[i] - z_dy[i];
236 s_dz[i] += x_dz[i] - z_dz[i];
237 }
238
239 apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
241 for i in 0..n {
243 let z2_inner = dx[i] + s2[i] - wy[i];
244 z2[i] = shrink(z2_inner, weight_mask[i] / mu2);
245 s2[i] = z2_inner - z2[i];
246 }
247 }
248 }
249 }
250
251 let mut dphi = vec![0.0f64; n];
254 {
255 apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
257 let mut maxv = 0.0f64;
258 for i in 0..n {
259 let v = (local_field[i] - dx[i]).abs() * (if mask[i] != 0 { 1.0 } else { 0.0 });
260 dphi[i] = v;
261 if v > maxv {
262 maxv = v;
263 }
264 }
265 if maxv > 0.0 {
266 for v in dphi.iter_mut() {
267 *v /= maxv;
268 }
269 }
270 }
271
272 {
276 let mu = params.mu1_l2;
277 let alpha = params.alpha_l2;
278 let mu2 = params.mu2;
279 let ll = alpha / mu;
280
281 let denom: Vec<f64> = (0..n).map(|i| denom_base[i] + mu * ee2[i]).collect();
283
284 let weight: Vec<f64> = (0..n).map(|i| {
286 let w = weight_mask[i];
287 let d = 1.0 - dphi[i];
288 w * w * d * d
289 }).collect();
290 let wy: Vec<f64> = (0..n).map(|i| weight[i] * local_field[i] / (weight[i] + mu2)).collect();
292
293 for i in 0..n {
295 z_dx[i] = 0.0; z_dy[i] = 0.0; z_dz[i] = 0.0;
296 s_dx[i] = 0.0; s_dy[i] = 0.0; s_dz[i] = 0.0;
297 s2[i] = 0.0;
298 }
299
300 for _t in 0..params.max_iter_l2 {
301 x_prev.copy_from_slice(&x);
302
303 fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
305 for i in 0..n {
306 let ax = x_dx[i] + s_dx[i];
307 let ay = x_dy[i] + s_dy[i];
308 let az = x_dz[i] + s_dz[i];
309 z_dx[i] = shrink(ax, ll);
310 z_dy[i] = shrink(ay, ll);
311 z_dz[i] = shrink(az, ll);
312 s_dx[i] += x_dx[i] - z_dx[i];
313 s_dy[i] += x_dy[i] - z_dy[i];
314 s_dz[i] += x_dz[i] - z_dz[i];
315 }
316
317 apply_dipole(&mut fft_ws, &k, &x, &mut cbuf, &mut dx);
319 for i in 0..n {
321 z2[i] = wy[i] + mu2 * (dx[i] + s2[i]) / (weight[i] + mu2);
322 s2[i] = s2[i] + dx[i] - z2[i];
323 }
324
325 for i in 0..n {
327 gx[i] = z2[i] - s2[i];
328 }
329 for (c, &v) in fdt.iter_mut().zip(gx.iter()) {
330 *c = Complex64::new(v, 0.0);
331 }
332 fft_ws.fft3d(&mut fdt);
333
334 for i in 0..n {
335 gx[i] = z_dx[i] - s_dx[i];
336 gy[i] = z_dy[i] - s_dy[i];
337 gz[i] = z_dz[i] - s_dz[i];
338 }
339 bdiv_inplace(&mut div, &gx, &gy, &gz, grid);
340 for (c, &v) in fdiv.iter_mut().zip(div.iter()) {
341 *c = Complex64::new(v, 0.0);
342 }
343 fft_ws.fft3d(&mut fdiv);
344
345 for i in 0..n {
346 let num = -fdiv[i] * mu + fdt[i] * (mu2 * k[i]);
349 cbuf[i] = if denom[i] > 1e-20 { num / denom[i] } else { Complex64::new(0.0, 0.0) };
351 }
352 fft_ws.ifft3d(&mut cbuf);
353 for i in 0..n {
354 x[i] = cbuf[i].re;
355 }
356
357 global_iter += 1;
358 progress(global_iter, total_iter);
359
360 let upd = percent_update(&x, &x_prev);
361 if upd < params.tol_update {
362 break;
363 }
364 }
365 }
366
367 if global_iter < total_iter {
369 progress(total_iter, total_iter);
370 }
371
372 apply_mask_zero(&mut x, mask);
373 x
374}
375
376#[cfg(test)]
377mod tests {
378 use super::*;
379
380 #[test]
381 fn test_hdqsm_zero_field() {
382 let n = 8;
383 let field = vec![0.0; n * n * n];
384 let mask = vec![1u8; n * n * n];
385 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
386 let params = HdQsmParams {
387 max_iter_l1: 5,
388 max_iter_l2: 10,
389 ..HdQsmParams::default()
390 };
391
392 let chi = hdqsm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
393
394 for &val in chi.iter() {
395 assert!(val.abs() < 1e-6, "Zero field should give zero chi, got {}", val);
396 }
397 }
398
399 #[test]
400 fn test_hdqsm_finite() {
401 let n = 8;
402 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
403 let mask = vec![1u8; n * n * n];
404 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
405 let params = HdQsmParams {
406 max_iter_l1: 5,
407 max_iter_l2: 10,
408 ..HdQsmParams::default()
409 };
410
411 let chi = hdqsm(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
412
413 for (i, &val) in chi.iter().enumerate() {
414 assert!(val.is_finite(), "Chi should be finite at index {}", i);
415 }
416 }
417}