1use crate::inversion::admm::prepare_fansi_spectral;
29use crate::utils::gradient::{bdiv_inplace, fgrad_inplace};
30use crate::utils::{apply_mask_zero, shrink};
31use crate::Grid;
32use num_complex::Complex64;
33use std::f64::consts::PI;
34
35#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
37#[derive(Clone, Debug)]
38pub struct FansiParams {
39 pub alpha1: f64,
41 pub mu1: f64,
43 pub mu2: f64,
45 pub alpha0: f64,
47 pub mu0: f64,
49 pub max_iter: usize,
51 pub tol_update: f64,
53 pub tol_delta: f64,
55 pub phase_scale: f64,
58 pub is_tgv: bool,
60}
61
62impl Default for FansiParams {
63 fn default() -> Self {
64 Self {
65 alpha1: 2e-4,
66 mu1: 2e-2,
67 mu2: 1.0,
68 alpha0: 4e-4,
69 mu0: 4e-2,
70 max_iter: 150,
71 tol_update: 0.1,
72 tol_delta: 1e-6,
73 phase_scale: 1.0,
74 is_tgv: false,
75 }
76 }
77}
78
79fn norm2(v: &[f64]) -> f64 {
81 v.iter().map(|&a| a * a).sum::<f64>().sqrt()
82}
83
84pub fn fansi(
97 local_field: &[f64],
98 mask: &[u8],
99 grid: &Grid,
100 bdir: (f64, f64, f64),
101 params: &FansiParams,
102 progress: impl FnMut(usize, usize),
103) -> Vec<f64> {
104 if params.is_tgv {
105 nltgv(local_field, mask, grid, bdir, params, progress)
106 } else {
107 nltv(local_field, mask, grid, bdir, params, progress)
108 }
109}
110
111fn nltv(
113 local_field: &[f64],
114 mask: &[u8],
115 grid: &Grid,
116 bdir: (f64, f64, f64),
117 params: &FansiParams,
118 mut progress: impl FnMut(usize, usize),
119) -> Vec<f64> {
120 let n = grid.n_total();
121
122 let (mut fft_ws, k, ee2) = prepare_fansi_spectral(grid, bdir);
123
124 let phase: Vec<f64> = local_field.iter().map(|&f| f * params.phase_scale).collect();
126 let w: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
127
128 let mu1 = params.mu1;
129 let mu2 = params.mu2;
130 let alpha_over_mu = params.alpha1 / mu1;
131
132 let mut x = vec![0.0f64; n];
134 let mut x_prev = vec![0.0f64; n];
135
136 let mut z_dx = vec![0.0f64; n];
138 let mut z_dy = vec![0.0f64; n];
139 let mut z_dz = vec![0.0f64; n];
140 let mut s_dx = vec![0.0f64; n];
141 let mut s_dy = vec![0.0f64; n];
142 let mut s_dz = vec![0.0f64; n];
143
144 let mut z2 = vec![0.0f64; n];
146 for i in 0..n {
147 let den = w[i] + mu2;
148 z2[i] = if den != 0.0 { w[i] * phase[i] / den } else { 0.0 };
149 }
150 let mut s2 = vec![0.0f64; n];
151
152 let mut fdiv = vec![Complex64::new(0.0, 0.0); n];
154 let mut fd2 = vec![Complex64::new(0.0, 0.0); n];
155 let mut xhat = vec![Complex64::new(0.0, 0.0); n];
156 let mut fx = vec![Complex64::new(0.0, 0.0); n];
157
158 let mut gxc = vec![0.0f64; n];
159 let mut gyc = vec![0.0f64; n];
160 let mut gzc = vec![0.0f64; n];
161 let mut x_dx = vec![0.0f64; n];
162 let mut x_dy = vec![0.0f64; n];
163 let mut x_dz = vec![0.0f64; n];
164 let mut div = vec![0.0f64; n];
165 let mut dx = vec![0.0f64; n]; let mut rhs_z2 = vec![0.0f64; n];
167 let mut diff = vec![0.0f64; n];
168 let mut update = vec![0.0f64; n];
169
170 for t in 0..params.max_iter {
171 progress(t + 1, params.max_iter);
172
173 for i in 0..n {
176 gxc[i] = z_dx[i] - s_dx[i];
177 gyc[i] = z_dy[i] - s_dy[i];
178 gzc[i] = z_dz[i] - s_dz[i];
179 }
180 bdiv_inplace(&mut div, &gxc, &gyc, &gzc, grid);
181 for i in 0..n {
182 fdiv[i] = Complex64::new(div[i], 0.0);
183 }
184 fft_ws.fft3d(&mut fdiv);
185
186 for i in 0..n {
188 fd2[i] = Complex64::new(z2[i] - s2[i], 0.0);
189 }
190 fft_ws.fft3d(&mut fd2);
191
192 for i in 0..n {
193 let num = -mu1 * fdiv[i] + mu2 * k[i] * fd2[i];
200 let den = mu2 * k[i] * k[i] + mu1 * ee2[i];
201 xhat[i] = if den > 1e-20 { num / den } else { Complex64::new(0.0, 0.0) };
206 }
207 fft_ws.ifft3d(&mut xhat);
208 x_prev.copy_from_slice(&x);
209 for i in 0..n {
210 x[i] = xhat[i].re;
211 }
212
213 let xnorm = norm2(&x);
215 if xnorm > 0.0 {
216 for i in 0..n {
217 diff[i] = x[i] - x_prev[i];
218 }
219 let x_update = 100.0 * norm2(&diff) / xnorm;
220 if x_update < params.tol_update || x_update.is_nan() {
221 progress(t + 1, t + 1);
222 break;
223 }
224 }
225
226 if t + 1 >= params.max_iter {
227 break;
228 }
229
230 for i in 0..n {
233 fx[i] = Complex64::new(x[i], 0.0);
234 }
235 fft_ws.fft3d(&mut fx);
236
237 fgrad_inplace(&mut x_dx, &mut x_dy, &mut x_dz, &x, grid);
238 for i in 0..n {
239 z_dx[i] = shrink(x_dx[i] + s_dx[i], alpha_over_mu);
240 z_dy[i] = shrink(x_dy[i] + s_dy[i], alpha_over_mu);
241 z_dz[i] = shrink(x_dz[i] + s_dz[i], alpha_over_mu);
242 s_dx[i] += x_dx[i] - z_dx[i];
243 s_dy[i] += x_dy[i] - z_dy[i];
244 s_dz[i] += x_dz[i] - z_dz[i];
245 }
246
247 for i in 0..n {
250 xhat[i] = fx[i] * k[i];
251 }
252 fft_ws.ifft3d(&mut xhat);
253 for i in 0..n {
254 dx[i] = xhat[i].re;
255 rhs_z2[i] = mu2 * (dx[i] + s2[i]);
256 z2[i] = rhs_z2[i] / mu2;
257 }
258
259 let mut delta = f64::INFINITY;
260 let mut inn = 0usize;
261 while delta > params.tol_delta && inn < 10 {
262 inn += 1;
263 let norm_old = norm2(&z2);
264 for i in 0..n {
265 let a = z2[i] - phase[i];
266 let numer = w[i] * a.sin() + mu2 * z2[i] - rhs_z2[i];
267 let denom = w[i] * a.cos() + mu2;
268 update[i] = numer / denom;
269 z2[i] -= update[i];
270 }
271 delta = if norm_old > 0.0 {
272 norm2(&update) / norm_old
273 } else {
274 0.0
275 };
276 }
277
278 for i in 0..n {
281 s2[i] += dx[i] - z2[i];
282 }
283 }
284
285 if params.phase_scale != 1.0 {
286 for v in &mut x {
287 *v /= params.phase_scale;
288 }
289 }
290
291 apply_mask_zero(&mut x, mask);
292 x
293}
294
295#[inline]
297fn spectral_mul_assign(dst: &mut [Complex64], m: &[Complex64]) {
298 for (d, &mm) in dst.iter_mut().zip(m.iter()) {
299 *d *= mm;
300 }
301}
302
303#[allow(clippy::too_many_lines)]
309fn nltgv(
310 local_field: &[f64],
311 mask: &[u8],
312 grid: &Grid,
313 bdir: (f64, f64, f64),
314 params: &FansiParams,
315 mut progress: impl FnMut(usize, usize),
316) -> Vec<f64> {
317 let n = grid.n_total();
318 let (nx, ny, nz) = (grid.nx(), grid.ny(), grid.nz());
319
320 let (mut fft_ws, k, _ee2) = prepare_fansi_spectral(grid, bdir);
322
323 let mut e1 = vec![Complex64::new(0.0, 0.0); n];
325 let mut e2 = vec![Complex64::new(0.0, 0.0); n];
326 let mut e3 = vec![Complex64::new(0.0, 0.0); n];
327 let two_pi = 2.0 * PI;
328 for kk in 0..nz {
329 let ez = Complex64::new(0.0, 1.0) * (two_pi * (kk as f64) / (nz as f64));
330 let e3v = Complex64::new(1.0, 0.0) - ez.exp();
331 for jj in 0..ny {
332 let ey = Complex64::new(0.0, 1.0) * (two_pi * (jj as f64) / (ny as f64));
333 let e2v = Complex64::new(1.0, 0.0) - ey.exp();
334 for ii in 0..nx {
335 let ex = Complex64::new(0.0, 1.0) * (two_pi * (ii as f64) / (nx as f64));
336 let e1v = Complex64::new(1.0, 0.0) - ex.exp();
337 let idx = ii + jj * nx + kk * nx * ny;
338 e1[idx] = e1v;
339 e2[idx] = e2v;
340 e3[idx] = e3v;
341 }
342 }
343 }
344
345 let phase: Vec<f64> = local_field.iter().map(|&f| f * params.phase_scale).collect();
346 let w: Vec<f64> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
347
348 let mu0 = params.mu0;
349 let mu1 = params.mu1;
350 let mu2 = params.mu2;
351
352 let mut d11 = vec![Complex64::new(0.0, 0.0); n];
354 let mut d21 = d11.clone();
355 let mut d31 = d11.clone();
356 let mut d41 = d11.clone();
357 let mut d12 = d11.clone();
358 let mut d22 = d11.clone();
359 let mut d32 = d11.clone();
360 let mut d42 = d11.clone();
361 let mut d13 = d11.clone();
362 let mut d23 = d11.clone();
363 let mut d33 = d11.clone();
364 let mut d43 = d11.clone();
365 let mut d14 = d11.clone();
366 let mut d24 = d11.clone();
367 let mut d34 = d11.clone();
368 let mut d44 = d11.clone();
369 let mut det_ainv = d11.clone();
370
371 let half = 0.5;
372 for i in 0..n {
373 let e1i = e1[i];
374 let e2i = e2[i];
375 let e3i = e3[i];
376 let et1 = e1i.conj();
377 let et2 = e2i.conj();
378 let et3 = e3i.conj();
379
380 let e1te1 = et1 * e1i;
381 let e2te2 = et2 * e2i;
382 let e3te3 = et3 * e3i;
383 let mu0h_e1te2 = (mu0 * half) * et1 * e2i;
384 let mu0h_e1te3 = (mu0 * half) * et1 * e3i;
385 let mu0h_e2te3 = (mu0 * half) * et2 * e3i;
386
387 let a0 = Complex64::new(mu2 * k[i] * k[i], 0.0);
389 let a1 = a0 + mu1 * (e1te1 + e2te2 + e3te3);
390 let a2 = Complex64::new(mu1, 0.0) + mu0 * (e1te1 + (e2te2 + e3te3) * half);
391 let a3 = Complex64::new(mu1, 0.0) + mu0 * (e1te1 * half + e2te2 + e3te3 * half);
392 let a4 = Complex64::new(mu1, 0.0) + mu0 * ((e1te1 + e2te2) * half + e3te3);
393 let a5 = -mu1 * e1i;
394 let a6 = -mu1 * e2i;
395 let a7 = mu0h_e1te2;
396 let a8 = -mu1 * e3i;
397 let a9 = mu0h_e1te3;
398 let a10 = mu0h_e2te3;
399 let a5t = a5.conj();
400 let a6t = a6.conj();
401 let a7t = a7.conj();
402 let a8t = a8.conj();
403 let a9t = a9.conj();
404 let a10t = a10.conj();
405
406 let c11 = a2 * a3 * a4 + a7t * a9 * a10t + a7 * a9t * a10
407 - a3 * a9 * a9t
408 - a2 * a10 * a10t
409 - a4 * a7 * a7t;
410 let c21 = a3 * a4 * a5t + a6t * a9 * a10t + a7 * a8t * a10
411 - a3 * a8t * a9
412 - a5t * a10 * a10t
413 - a4 * a6t * a7;
414 let c31 = a4 * a5t * a7t + a6t * a9 * a9t + a2 * a8t * a10
415 - a7t * a8t * a9
416 - a5t * a9t * a10
417 - a2 * a4 * a6t;
418 let c41 = a5t * a7t * a10t + a6t * a7 * a9t + a2 * a3 * a8t
419 - a7 * a7t * a8t
420 - a3 * a5t * a9t
421 - a2 * a6t * a10t;
422 let c12 = a3 * a4 * a5 + a7t * a8 * a10t + a6 * a9t * a10
423 - a3 * a8 * a9t
424 - a5 * a10 * a10t
425 - a4 * a6 * a7t;
426 let c22 = a1 * a3 * a4 + a6t * a8 * a10t + a6 * a8t * a10
427 - a3 * a8 * a8t
428 - a1 * a10 * a10t
429 - a4 * a6 * a6t;
430 let c32 = a1 * a4 * a7t + a6t * a8 * a9t + a5 * a8t * a10
431 - a7t * a8 * a8t
432 - a1 * a9t * a10
433 - a4 * a5 * a6t;
434 let c42 = a1 * a7t * a10t + a6 * a6t * a9t + a3 * a5 * a8t
435 - a6 * a7t * a8t
436 - a1 * a3 * a9t
437 - a5 * a6t * a10t;
438 let c13 = a4 * a5 * a7 + a2 * a8 * a10t + a6 * a9 * a9t
439 - a7 * a8 * a9t
440 - a5 * a9 * a10t
441 - a2 * a4 * a6;
442 let c23 = a1 * a4 * a7 + a5t * a8 * a10t + a6 * a8t * a9
443 - a7 * a8 * a8t
444 - a1 * a9 * a10t
445 - a4 * a5t * a6;
446 let c33 = a1 * a2 * a4 + a5t * a8 * a9t + a5 * a8t * a9
447 - a2 * a8 * a8t
448 - a1 * a9 * a9t
449 - a4 * a5 * a5t;
450 let c43 = a1 * a2 * a10t + a5t * a6 * a9t + a5 * a7 * a8t
451 - a2 * a6 * a8t
452 - a1 * a7 * a9t
453 - a5 * a5t * a10t;
454 let c14 = a5 * a7 * a10 + a2 * a3 * a8 + a6 * a7t * a9
455 - a7 * a7t * a8
456 - a3 * a5 * a9
457 - a2 * a6 * a10;
458 let c24 = a1 * a7 * a10 + a3 * a5t * a8 + a6 * a6t * a9
459 - a6t * a7 * a8
460 - a1 * a3 * a9
461 - a5t * a6 * a10;
462 let c34 = a1 * a2 * a10 + a5t * a7t * a8 + a5 * a6t * a9
463 - a2 * a6t * a8
464 - a1 * a7t * a9
465 - a5 * a5t * a10;
466 let c44 = a1 * a2 * a3 + a5t * a6 * a7t + a5 * a6t * a7
467 - a2 * a6 * a6t
468 - a1 * a7 * a7t
469 - a3 * a5 * a5t;
470
471 let det_a = a1 * c11 - a5 * c21 + a6 * c31 - a8 * c41;
472
473 d11[i] = c11;
474 d21[i] = c21;
475 d31[i] = c31;
476 d41[i] = c41;
477 d12[i] = c12;
478 d22[i] = c22;
479 d32[i] = c32;
480 d42[i] = c42;
481 d13[i] = c13;
482 d23[i] = c23;
483 d33[i] = c33;
484 d43[i] = c43;
485 d14[i] = c14;
486 d24[i] = c24;
487 d34[i] = c34;
488 d44[i] = c44;
489 det_ainv[i] = if det_a.norm() > 1e-20 {
493 Complex64::new(1.0, 0.0) / det_a
494 } else {
495 Complex64::new(0.0, 0.0)
496 };
497 }
498
499 let et1: Vec<Complex64> = e1.iter().map(|c| c.conj()).collect();
501 let et2: Vec<Complex64> = e2.iter().map(|c| c.conj()).collect();
502 let et3: Vec<Complex64> = e3.iter().map(|c| c.conj()).collect();
503
504 let mut x = vec![0.0f64; n];
506 let mut x_prev = vec![0.0f64; n];
507 let mut v1 = vec![0.0f64; n];
508 let mut v2 = vec![0.0f64; n];
509 let mut v3 = vec![0.0f64; n];
510
511 let mut z1_1 = vec![0.0f64; n];
513 let mut z1_2 = vec![0.0f64; n];
514 let mut z1_3 = vec![0.0f64; n];
515 let mut s1_1 = vec![0.0f64; n];
516 let mut s1_2 = vec![0.0f64; n];
517 let mut s1_3 = vec![0.0f64; n];
518
519 let mut z0_1 = vec![0.0f64; n];
521 let mut z0_2 = vec![0.0f64; n];
522 let mut z0_3 = vec![0.0f64; n];
523 let mut z0_4 = vec![0.0f64; n];
524 let mut z0_5 = vec![0.0f64; n];
525 let mut z0_6 = vec![0.0f64; n];
526 let mut s0_1 = vec![0.0f64; n];
527 let mut s0_2 = vec![0.0f64; n];
528 let mut s0_3 = vec![0.0f64; n];
529 let mut s0_4 = vec![0.0f64; n];
530 let mut s0_5 = vec![0.0f64; n];
531 let mut s0_6 = vec![0.0f64; n];
532
533 let mut z2 = vec![0.0f64; n];
535 for i in 0..n {
536 let den = w[i] + mu2;
537 z2[i] = if den != 0.0 { w[i] * phase[i] / den } else { 0.0 };
538 }
539 let mut s2 = vec![0.0f64; n];
540
541 let alpha1_over_mu1 = params.alpha1 / mu1;
542 let alpha0_over_mu0 = params.alpha0 / mu0;
543
544 let mut rhs1 = vec![Complex64::new(0.0, 0.0); n];
546 let mut rhs2 = rhs1.clone();
547 let mut rhs3 = rhs1.clone();
548 let mut rhs4 = rhs1.clone();
549 let mut t1 = rhs1.clone();
550 let mut t2 = rhs1.clone();
551 let mut t3 = rhs1.clone();
552 let mut fx = rhs1.clone();
553 let mut fv1 = rhs1.clone();
554 let mut fv2 = rhs1.clone();
555 let mut fv3 = rhs1.clone();
556 let mut cbuf = rhs1.clone();
557
558 let mut dx1 = vec![0.0f64; n];
559 let mut dx2 = vec![0.0f64; n];
560 let mut dx3 = vec![0.0f64; n];
561 let mut ev1 = vec![0.0f64; n];
562 let mut ev2 = vec![0.0f64; n];
563 let mut ev3 = vec![0.0f64; n];
564 let mut ev4 = vec![0.0f64; n];
565 let mut ev5 = vec![0.0f64; n];
566 let mut ev6 = vec![0.0f64; n];
567 let mut dx = vec![0.0f64; n];
568 let mut rhs_z2 = vec![0.0f64; n];
569 let mut diff = vec![0.0f64; n];
570 let mut update = vec![0.0f64; n];
571
572 let mut rbuf = vec![0.0f64; n];
574
575 macro_rules! fft_real {
577 ($dst:expr, $src:expr) => {{
578 for i in 0..n {
579 $dst[i] = Complex64::new($src[i], 0.0);
580 }
581 fft_ws.fft3d(&mut $dst);
582 }};
583 }
584 macro_rules! fft_real_diff {
586 ($dst:expr, $a:expr, $b:expr) => {{
587 for i in 0..n {
588 rbuf[i] = $a[i] - $b[i];
589 }
590 fft_real!($dst, rbuf);
591 }};
592 }
593
594 for t in 0..params.max_iter {
595 progress(t + 1, params.max_iter);
596
597 for i in 0..n {
600 cbuf[i] = Complex64::new(z2[i] - s2[i], 0.0);
601 }
602 fft_ws.fft3d(&mut cbuf);
603 for i in 0..n {
604 rhs1[i] = mu2 * k[i] * cbuf[i];
605 }
606 fft_real_diff!(t1, z1_1, s1_1);
607 fft_real_diff!(t2, z1_2, s1_2);
608 fft_real_diff!(t3, z1_3, s1_3);
609 for i in 0..n {
610 rhs1[i] += mu1 * (et1[i] * t1[i] + et2[i] * t2[i] + et3[i] * t3[i]);
611 }
612
613 for i in 0..n {
616 rhs2[i] = -mu1 * t1[i];
617 rhs3[i] = -mu1 * t2[i];
618 rhs4[i] = -mu1 * t3[i];
619 }
620 fft_real_diff!(t1, z0_1, s0_1);
622 fft_real_diff!(t2, z0_4, s0_4);
623 fft_real_diff!(t3, z0_5, s0_5);
624 for i in 0..n {
625 rhs2[i] += mu0 * (et1[i] * t1[i] + et2[i] * t2[i] + et3[i] * t3[i]);
626 }
627 fft_real_diff!(t1, z0_2, s0_2);
629 fft_real_diff!(t2, z0_4, s0_4);
630 fft_real_diff!(t3, z0_6, s0_6);
631 for i in 0..n {
632 rhs3[i] += mu0 * (et2[i] * t1[i] + et1[i] * t2[i] + et3[i] * t3[i]);
633 }
634 fft_real_diff!(t1, z0_3, s0_3);
636 fft_real_diff!(t2, z0_5, s0_5);
637 fft_real_diff!(t3, z0_6, s0_6);
638 for i in 0..n {
639 rhs4[i] += mu0 * (et3[i] * t1[i] + et1[i] * t2[i] + et2[i] * t3[i]);
640 }
641
642 for i in 0..n {
644 let r1 = rhs1[i];
645 let r2 = rhs2[i];
646 let r3 = rhs3[i];
647 let r4 = rhs4[i];
648 let da = det_ainv[i];
649 fx[i] = (r1 * d11[i] - r2 * d21[i] + r3 * d31[i] - r4 * d41[i]) * da;
650 fv1[i] = (-r1 * d12[i] + r2 * d22[i] - r3 * d32[i] + r4 * d42[i]) * da;
651 fv2[i] = (r1 * d13[i] - r2 * d23[i] + r3 * d33[i] - r4 * d43[i]) * da;
652 fv3[i] = (-r1 * d14[i] + r2 * d24[i] - r3 * d34[i] + r4 * d44[i]) * da;
653 }
654 t1.copy_from_slice(&fx);
657 fft_ws.ifft3d(&mut t1);
658 x_prev.copy_from_slice(&x);
659 for i in 0..n {
660 x[i] = t1[i].re;
661 }
662 t1.copy_from_slice(&fv1);
663 fft_ws.ifft3d(&mut t1);
664 for i in 0..n {
665 v1[i] = t1[i].re;
666 }
667 t1.copy_from_slice(&fv2);
668 fft_ws.ifft3d(&mut t1);
669 for i in 0..n {
670 v2[i] = t1[i].re;
671 }
672 t1.copy_from_slice(&fv3);
673 fft_ws.ifft3d(&mut t1);
674 for i in 0..n {
675 v3[i] = t1[i].re;
676 }
677
678 let xnorm = norm2(&x);
680 if xnorm > 0.0 {
681 for i in 0..n {
682 diff[i] = x[i] - x_prev[i];
683 }
684 let x_update = 100.0 * norm2(&diff) / xnorm;
685 if x_update < params.tol_update || x_update.is_nan() {
686 progress(t + 1, t + 1);
687 break;
688 }
689 }
690
691 if t + 1 >= params.max_iter {
692 break;
693 }
694
695 fft_real!(fx, x);
697 fft_real!(fv1, v1);
698 fft_real!(fv2, v2);
699 fft_real!(fv3, v3);
700
701 cbuf.copy_from_slice(&fx);
703 spectral_mul_assign(&mut cbuf, &e1);
704 fft_ws.ifft3d(&mut cbuf);
705 for i in 0..n {
706 dx1[i] = cbuf[i].re;
707 }
708 cbuf.copy_from_slice(&fx);
709 spectral_mul_assign(&mut cbuf, &e2);
710 fft_ws.ifft3d(&mut cbuf);
711 for i in 0..n {
712 dx2[i] = cbuf[i].re;
713 }
714 cbuf.copy_from_slice(&fx);
715 spectral_mul_assign(&mut cbuf, &e3);
716 fft_ws.ifft3d(&mut cbuf);
717 for i in 0..n {
718 dx3[i] = cbuf[i].re;
719 }
720
721 cbuf.copy_from_slice(&fv1);
723 spectral_mul_assign(&mut cbuf, &e1);
724 fft_ws.ifft3d(&mut cbuf);
725 for i in 0..n {
726 ev1[i] = cbuf[i].re;
727 }
728 cbuf.copy_from_slice(&fv2);
729 spectral_mul_assign(&mut cbuf, &e2);
730 fft_ws.ifft3d(&mut cbuf);
731 for i in 0..n {
732 ev2[i] = cbuf[i].re;
733 }
734 cbuf.copy_from_slice(&fv3);
735 spectral_mul_assign(&mut cbuf, &e3);
736 fft_ws.ifft3d(&mut cbuf);
737 for i in 0..n {
738 ev3[i] = cbuf[i].re;
739 }
740
741 for i in 0..n {
743 cbuf[i] = e1[i] * fv2[i] + e2[i] * fv1[i];
744 }
745 fft_ws.ifft3d(&mut cbuf);
746 for i in 0..n {
747 ev4[i] = cbuf[i].re * 0.5;
748 }
749 for i in 0..n {
750 cbuf[i] = e1[i] * fv3[i] + e3[i] * fv1[i];
751 }
752 fft_ws.ifft3d(&mut cbuf);
753 for i in 0..n {
754 ev5[i] = cbuf[i].re * 0.5;
755 }
756 for i in 0..n {
757 cbuf[i] = e2[i] * fv3[i] + e3[i] * fv2[i];
758 }
759 fft_ws.ifft3d(&mut cbuf);
760 for i in 0..n {
761 ev6[i] = cbuf[i].re * 0.5;
762 }
763
764 for i in 0..n {
766 z0_1[i] = shrink(ev1[i] + s0_1[i], alpha0_over_mu0);
767 z0_2[i] = shrink(ev2[i] + s0_2[i], alpha0_over_mu0);
768 z0_3[i] = shrink(ev3[i] + s0_3[i], alpha0_over_mu0);
769 z0_4[i] = shrink(ev4[i] + s0_4[i], alpha0_over_mu0);
770 z0_5[i] = shrink(ev5[i] + s0_5[i], alpha0_over_mu0);
771 z0_6[i] = shrink(ev6[i] + s0_6[i], alpha0_over_mu0);
772 }
773
774 for i in 0..n {
776 z1_1[i] = shrink(dx1[i] - v1[i] + s1_1[i], alpha1_over_mu1);
777 z1_2[i] = shrink(dx2[i] - v2[i] + s1_2[i], alpha1_over_mu1);
778 z1_3[i] = shrink(dx3[i] - v3[i] + s1_3[i], alpha1_over_mu1);
779 }
780
781 for i in 0..n {
784 cbuf[i] = fx[i] * k[i];
785 }
786 fft_ws.ifft3d(&mut cbuf);
787 for i in 0..n {
788 dx[i] = cbuf[i].re;
789 rhs_z2[i] = mu2 * (dx[i] + s2[i]);
790 z2[i] = rhs_z2[i] / mu2;
791 }
792 let mut delta = f64::INFINITY;
793 let mut inn = 0usize;
794 while delta > params.tol_delta && inn < 50 {
795 inn += 1;
796 let norm_old = norm2(&z2);
797 for i in 0..n {
798 let a = z2[i] - phase[i];
799 let numer = w[i] * a.sin() + mu2 * z2[i] - rhs_z2[i];
800 let denom = w[i] * a.cos() + mu2;
801 update[i] = numer / denom;
802 z2[i] -= update[i];
803 }
804 delta = if norm_old > 0.0 {
805 norm2(&update) / norm_old
806 } else {
807 0.0
808 };
809 }
810
811 for i in 0..n {
813 s0_1[i] += ev1[i] - z0_1[i];
814 s0_2[i] += ev2[i] - z0_2[i];
815 s0_3[i] += ev3[i] - z0_3[i];
816 s0_4[i] += ev4[i] - z0_4[i];
817 s0_5[i] += ev5[i] - z0_5[i];
818 s0_6[i] += ev6[i] - z0_6[i];
819 s1_1[i] += dx1[i] - v1[i] - z1_1[i];
820 s1_2[i] += dx2[i] - v2[i] - z1_2[i];
821 s1_3[i] += dx3[i] - v3[i] - z1_3[i];
822 s2[i] += dx[i] - z2[i];
823 }
824 }
825
826 if params.phase_scale != 1.0 {
827 for v in &mut x {
828 *v /= params.phase_scale;
829 }
830 }
831
832 apply_mask_zero(&mut x, mask);
833 x
834}
835
836#[cfg(test)]
837mod tests {
838 use super::*;
839
840 #[test]
841 fn test_fansi_nltv_zero_field() {
842 let n = 8;
844 let field = vec![0.0; n * n * n];
845 let mask = vec![1u8; n * n * n];
846 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
847 let params = FansiParams {
848 max_iter: 10,
849 is_tgv: false,
850 ..FansiParams::default()
851 };
852
853 let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
854
855 for &val in chi.iter() {
856 assert!(val.abs() < 1e-6, "Zero field should give ~zero chi, got {}", val);
857 }
858 }
859
860 #[test]
861 fn test_fansi_nltv_finite() {
862 let n = 8;
864 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
865 let mask = vec![1u8; n * n * n];
866 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
867 let params = FansiParams {
868 max_iter: 10,
869 is_tgv: false,
870 ..FansiParams::default()
871 };
872
873 let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
874
875 for (i, &val) in chi.iter().enumerate() {
876 assert!(val.is_finite(), "Chi should be finite at index {}", i);
877 }
878 }
879
880 #[test]
881 fn test_fansi_nltgv_zero_field() {
882 let n = 8;
884 let field = vec![0.0; n * n * n];
885 let mask = vec![1u8; n * n * n];
886 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
887 let params = FansiParams {
888 max_iter: 10,
889 is_tgv: true,
890 ..FansiParams::default()
891 };
892
893 let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
894
895 for &val in chi.iter() {
896 assert!(val.abs() < 1e-6, "Zero field should give ~zero chi, got {}", val);
897 }
898 }
899
900 #[test]
901 fn test_fansi_nltgv_finite() {
902 let n = 8;
904 let field: Vec<f64> = (0..n * n * n).map(|i| (i as f64) * 0.001).collect();
905 let mask = vec![1u8; n * n * n];
906 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
907 let params = FansiParams {
908 max_iter: 10,
909 is_tgv: true,
910 ..FansiParams::default()
911 };
912
913 let chi = fansi(&field, &mask, &grid, (0.0, 0.0, 1.0), ¶ms, |_, _| {});
914
915 for (i, &val) in chi.iter().enumerate() {
916 assert!(val.is_finite(), "Chi should be finite at index {}", i);
917 }
918 }
919}