1use num_complex::Complex32;
32#[cfg(feature = "parallel")]
33use rayon::prelude::*;
34use crate::fft::Fft3dWorkspaceF32;
35use crate::kernels::dipole::dipole_kernel_f32;
36use crate::kernels::smv::smv_kernel_f32;
37use crate::utils::simd_ops::{
38 dot_product_f32, norm_squared_f32, axpy_f32, xpby_f32,
39 apply_gradient_weights_f32, compute_p_weights_f32, combine_terms_f32, negate_f32,
40};
41use crate::Grid;
42#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
49#[derive(Clone, Debug)]
50pub struct MediParams {
51 pub lambda: f64,
53 pub merit: bool,
55 pub smv: bool,
57 pub smv_radius: f64,
59 pub data_weighting: i32,
61 pub percentage: f64,
63 pub cg_tol: f64,
65 pub cg_max_iter: usize,
67 pub max_iter: usize,
69 pub tol: f64,
71}
72
73impl Default for MediParams {
74 fn default() -> Self {
75 Self {
76 lambda: 1e-3,
81 merit: false,
82 smv: true,
83 smv_radius: 5.0,
84 data_weighting: 1,
85 percentage: 0.3,
86 cg_tol: 0.01,
87 cg_max_iter: 10,
88 max_iter: 30,
89 tol: 0.1,
90 }
91 }
92}
93
94pub struct MediWorkspace {
97 pub n_total: usize,
98 pub nx: usize,
99 pub ny: usize,
100 pub nz: usize,
101 pub vsx: f32,
102 pub vsy: f32,
103 pub vsz: f32,
104
105 pub fft_ws: Fft3dWorkspaceF32,
107
108 pub gx: Vec<f32>,
110 pub gy: Vec<f32>,
111 pub gz: Vec<f32>,
112
113 pub reg_x: Vec<f32>,
115 pub reg_y: Vec<f32>,
116 pub reg_z: Vec<f32>,
117
118 pub div_buf: Vec<f32>,
120
121 pub complex_buf: Vec<Complex32>,
123 pub complex_buf2: Vec<Complex32>,
124
125 pub dipole_buf: Vec<f32>,
127
128 pub cg_r: Vec<f32>,
130 pub cg_p: Vec<f32>,
131 pub cg_ap: Vec<f32>,
132}
133
134impl MediWorkspace {
135 pub fn new(grid: &Grid) -> Self {
137 let (nx, ny, nz) = grid.dims;
138 let (vsx, vsy, vsz) = (grid.vsx() as f32, grid.vsy() as f32, grid.vsz() as f32);
139 let n_total = grid.n_total();
140
141 Self {
142 n_total,
143 nx, ny, nz,
144 vsx, vsy, vsz,
145 fft_ws: Fft3dWorkspaceF32::new(nx, ny, nz),
146 gx: vec![0.0; n_total],
147 gy: vec![0.0; n_total],
148 gz: vec![0.0; n_total],
149 reg_x: vec![0.0; n_total],
150 reg_y: vec![0.0; n_total],
151 reg_z: vec![0.0; n_total],
152 div_buf: vec![0.0; n_total],
153 complex_buf: vec![Complex32::new(0.0, 0.0); n_total],
154 complex_buf2: vec![Complex32::new(0.0, 0.0); n_total],
155 dipole_buf: vec![0.0; n_total],
156 cg_r: vec![0.0; n_total],
157 cg_p: vec![0.0; n_total],
158 cg_ap: vec![0.0; n_total],
159 }
160 }
161}
162
163#[inline]
165pub(crate) fn apply_dipole_conv(
166 fft_ws: &mut Fft3dWorkspaceF32,
167 x: &[f32],
168 d_kernel: &[f32],
169 out: &mut [f32],
170 complex_buf: &mut [Complex32],
171) {
172 fft_ws.apply_dipole_inplace(x, d_kernel, out, complex_buf);
173}
174
175pub(crate) struct MediOpBuffers<'a> {
177 pub gx: &'a mut [f32],
178 pub gy: &'a mut [f32],
179 pub gz: &'a mut [f32],
180 pub reg_x: &'a mut [f32],
181 pub reg_y: &'a mut [f32],
182 pub reg_z: &'a mut [f32],
183 pub div_buf: &'a mut [f32],
184 pub dipole_buf: &'a mut [f32],
185 pub complex_buf: &'a mut [Complex32],
186 pub complex_buf2: &'a mut [Complex32],
187}
188
189#[inline]
194pub(crate) fn apply_medi_operator_core(
195 fft_ws: &mut Fft3dWorkspaceF32,
196 bufs: &mut MediOpBuffers,
197 n: usize,
198 nx: usize, ny: usize, nz: usize,
199 vsx: f32, vsy: f32, vsz: f32,
200 dx: &[f32],
201 w: &[Complex32],
202 d_kernel: &[f32],
203 mx: &[f32], my: &[f32], mz: &[f32], vr: &[f32],
207 lambda: f32,
208 out: &mut [f32],
209) {
210 fgrad_periodic_inplace_f32(bufs.gx, bufs.gy, bufs.gz, dx, nx, ny, nz, vsx, vsy, vsz);
212
213 apply_gradient_weights_f32(
216 bufs.reg_x, bufs.reg_y, bufs.reg_z,
217 mx, my, mz, vr,
218 bufs.gx, bufs.gy, bufs.gz,
219 );
220
221 bdiv_periodic_inplace_f32(bufs.div_buf, bufs.reg_x, bufs.reg_y, bufs.reg_z, nx, ny, nz, vsx, vsy, vsz);
223
224 apply_dipole_conv(fft_ws, dx, d_kernel, bufs.dipole_buf, bufs.complex_buf);
226
227 for i in 0..n {
229 let w_mag_sq = w[i].norm_sqr();
230 bufs.complex_buf2[i] = Complex32::new(bufs.dipole_buf[i] * w_mag_sq, 0.0);
231 }
232
233 fft_ws.fft3d(bufs.complex_buf2);
235 for i in 0..n {
236 bufs.complex_buf2[i] *= d_kernel[i];
237 }
238 fft_ws.ifft3d(bufs.complex_buf2);
239
240 for i in 0..n {
243 bufs.dipole_buf[i] = bufs.complex_buf2[i].re;
244 }
245 combine_terms_f32(out, bufs.div_buf, bufs.dipole_buf, lambda);
246}
247
248#[inline]
254fn cg_solve_medi<F>(
255 ws: &mut MediWorkspace,
256 w: &[Complex32],
257 d_kernel: &[f32],
258 mx: &[f32], my: &[f32], mz: &[f32], vr: &[f32],
262 lambda: f32,
263 b: &[f32],
264 x: &mut [f32],
265 tol: f32,
266 max_iter: usize,
267 mut progress_callback: F,
268) where
269 F: FnMut(usize, usize),
270{
271 let n = ws.n_total;
272 let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
273 let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
274
275 x.fill(0.0);
277
278 ws.cg_r.copy_from_slice(b);
280
281 ws.cg_p.copy_from_slice(&ws.cg_r);
283
284 let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
286
287 let b_norm: f32 = norm_squared_f32(b).sqrt();
289 if b_norm < 1e-10 {
290 return; }
292
293 let mut p_copy = vec![0.0f32; n];
295
296 for cg_iter in 0..max_iter {
297 progress_callback(cg_iter + 1, max_iter);
299
300 p_copy.copy_from_slice(&ws.cg_p);
302
303 {
305 let mut bufs = MediOpBuffers {
306 gx: &mut ws.gx,
307 gy: &mut ws.gy,
308 gz: &mut ws.gz,
309 reg_x: &mut ws.reg_x,
310 reg_y: &mut ws.reg_y,
311 reg_z: &mut ws.reg_z,
312 div_buf: &mut ws.div_buf,
313 dipole_buf: &mut ws.dipole_buf,
314 complex_buf: &mut ws.complex_buf,
315 complex_buf2: &mut ws.complex_buf2,
316 };
317 apply_medi_operator_core(
318 &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
319 &p_copy, w, d_kernel, mx, my, mz, vr, lambda, &mut ws.cg_ap
320 );
321 }
322
323 let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
325
326 if pap.abs() < 1e-15 {
327 break;
328 }
329
330 let alpha = rsold / pap;
331
332 axpy_f32(x, alpha, &ws.cg_p);
334
335 axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
337
338 let rsnew: f32 = norm_squared_f32(&ws.cg_r);
340 let residual = rsnew.sqrt();
341
342 if residual < tol * b_norm {
344 break;
345 }
346
347 let beta = rsnew / rsold;
348
349 xpby_f32(&mut ws.cg_p, &ws.cg_r, beta);
351
352 rsold = rsnew;
353 }
354}
355
356pub fn medi(
371 local_field: &[f64],
372 n_std: &[f64],
373 magnitude: &[f64],
374 mask: &[u8],
375 grid: &Grid,
376 bdir: (f64, f64, f64),
377 params: &MediParams,
378 mut progress: impl FnMut(usize, usize),
379) -> Vec<f64> {
380 let (nx, ny, nz) = grid.dims;
381 let n_total = grid.n_total();
382
383 let vsx_f32 = grid.vsx() as f32;
385 let vsy_f32 = grid.vsy() as f32;
386 let vsz_f32 = grid.vsz() as f32;
387 let lambda_f32 = params.lambda as f32;
388 let bdir_f32 = (bdir.0 as f32, bdir.1 as f32, bdir.2 as f32);
389 let smv_radius_f32 = params.smv_radius as f32;
390 let percentage_f32 = params.percentage as f32;
391 let cg_tol_f32 = params.cg_tol as f32;
392 let tol_f32 = params.tol as f32;
393 let max_iter = params.max_iter;
394 let cg_max_iter = params.cg_max_iter;
395 let data_weighting = params.data_weighting;
396
397 let local_field_f32: Vec<f32> = local_field.iter().map(|&v| v as f32).collect();
399 let n_std_f32: Vec<f32> = n_std.iter().map(|&v| v as f32).collect();
400 let magnitude_f32: Vec<f32> = magnitude.iter().map(|&v| v as f32).collect();
401
402 let mut ws = MediWorkspace::new(grid);
404
405 let mut rdf: Vec<f32> = local_field_f32.clone();
407 let mut work_mask: Vec<u8> = mask.to_vec();
408 let mut tempn: Vec<f32> = n_std_f32.clone();
409
410 for i in 0..n_total {
412 if mask[i] == 0 {
413 tempn[i] = 0.0;
414 }
415 }
416
417 let mut d_kernel = dipole_kernel_f32(grid, bdir_f32);
419
420 let sphere_k = if params.smv {
422 let sk = smv_kernel_f32(grid, smv_radius_f32);
423
424 let mut sk_fft: Vec<Complex32> = sk.iter()
426 .map(|&v| Complex32::new(v, 0.0))
427 .collect();
428 ws.fft_ws.fft3d(&mut sk_fft);
429
430 let mask_f32: Vec<f32> = work_mask.iter().map(|&m| m as f32).collect();
432 let smv_mask = apply_smv_kernel_ws(&mask_f32, &sk_fft, &mut ws);
433 for i in 0..n_total {
434 work_mask[i] = if smv_mask[i] > 0.999 { 1 } else { 0 };
435 }
436
437 for i in 0..n_total {
439 d_kernel[i] *= 1.0 - sk[i];
440 }
441
442 let smv_rdf = apply_smv_kernel_ws(&rdf, &sk_fft, &mut ws);
444 for i in 0..n_total {
445 rdf[i] -= smv_rdf[i];
446 if work_mask[i] == 0 {
447 rdf[i] = 0.0;
448 }
449 }
450
451 let tempn_sq: Vec<f32> = tempn.iter().map(|&t| t * t).collect();
453 let smv_tempn_sq = apply_smv_kernel_ws(&tempn_sq, &sk_fft, &mut ws);
454 for i in 0..n_total {
455 tempn[i] = (smv_tempn_sq[i] + tempn_sq[i]).sqrt();
456 }
457
458 Some(sk_fft)
459 } else {
460 None
461 };
462
463 let mut m = dataterm_mask_f32(data_weighting, &tempn, &work_mask);
465
466 let mut b0: Vec<Complex32> = rdf.iter()
468 .zip(m.iter())
469 .map(|(&f, &mi)| {
470 let phase = Complex32::new(0.0, f);
471 mi * phase.exp()
472 })
473 .collect();
474
475 let (w_gx, w_gy, w_gz) = gradient_mask_f32(&magnitude_f32, &work_mask, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32, percentage_f32);
478
479 let w_gx = if w_gx.iter().any(|&v| v != 0.0) { w_gx } else { magnitude_f32.clone() };
481 let w_gy = if w_gy.iter().any(|&v| v != 0.0) { w_gy } else { magnitude_f32.clone() };
482 let w_gz = if w_gz.iter().any(|&v| v != 0.0) { w_gz } else { magnitude_f32.clone() };
483
484 let mut chi = vec![0.0f32; n_total];
486 let mut dx = vec![0.0f32; n_total]; let mut rhs = vec![0.0f32; n_total]; let mut vr = vec![0.0f32; n_total]; let mut w: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); n_total]; let mut chi_prev = vec![0.0f32; n_total]; let mut badpoint = vec![0.0f32; n_total];
492 let mut n_std_work: Vec<f32> = n_std_f32.clone();
493
494 let beta = 1.49e-8_f32;
498
499 let total_steps = max_iter * cg_max_iter;
501
502 for iter in 0..max_iter {
504 chi_prev.copy_from_slice(&chi);
506
507 fgrad_periodic_inplace_f32(
512 &mut ws.gx, &mut ws.gy, &mut ws.gz,
513 &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
514 );
515
516 compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
517
518 apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
520 for i in 0..n_total {
521 let phase = Complex32::new(0.0, ws.dipole_buf[i]);
522 w[i] = m[i] * phase.exp();
523 }
524
525 compute_rhs_inplace(&chi, &w, &b0, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda_f32, &mut rhs, &mut ws);
527
528 negate_f32(&mut rhs);
530
531 let gn_iter = iter;
533 cg_solve_medi(
534 &mut ws, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda_f32, &rhs, &mut dx, cg_tol_f32, cg_max_iter,
535 |cg_iter, cg_total| {
536 let current = gn_iter * cg_total + cg_iter;
537 progress(current, total_steps);
538 }
539 );
540
541 axpy_f32(&mut chi, 1.0, &dx);
543
544 let norm_dx_sq = norm_squared_f32(&dx);
546 let norm_chi_sq = norm_squared_f32(&chi_prev);
547 let rel_change = norm_dx_sq.sqrt() / (norm_chi_sq.sqrt() + 1e-6);
548
549 if params.merit {
551 apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
553 let mut wres: Vec<Complex32> = ws.dipole_buf.iter()
554 .zip(m.iter())
555 .zip(b0.iter())
556 .map(|((&dc, &mi), &b0i)| {
557 let phase = Complex32::new(0.0, dc);
558 mi * phase.exp() - b0i
559 })
560 .collect();
561
562 let mask_count = work_mask.iter().filter(|&&m| m != 0).count() as f32;
564 if mask_count > 0.0 {
565 let mean_wres: Complex32 = wres.iter()
566 .zip(work_mask.iter())
567 .filter(|(_, &m)| m != 0)
568 .map(|(w, _)| w)
569 .sum::<Complex32>() / mask_count;
570
571 for i in 0..n_total {
572 if work_mask[i] != 0 {
573 wres[i] -= mean_wres;
574 }
575 }
576 }
577
578 let abs_wres: Vec<f32> = wres.iter()
580 .zip(work_mask.iter())
581 .filter(|(_, &m)| m != 0)
582 .map(|(w, _)| w.norm())
583 .collect();
584
585 if !abs_wres.is_empty() {
586 let mean_abs: f32 = abs_wres.iter().sum::<f32>() / abs_wres.len() as f32;
587 let var: f32 = abs_wres.iter()
588 .map(|&v| (v - mean_abs).powi(2))
589 .sum::<f32>() / abs_wres.len() as f32;
590 let factor = var.sqrt() * 6.0;
591
592 if factor > 1e-10 {
593 let mut wres_norm: Vec<f32> = wres.iter()
595 .map(|w| w.norm() / factor)
596 .collect();
597
598 for v in wres_norm.iter_mut() {
600 if *v < 1.0 {
601 *v = 1.0;
602 }
603 }
604
605 for i in 0..n_total {
607 if wres_norm[i] > 1.0 {
608 badpoint[i] = 1.0;
609 }
610 if work_mask[i] != 0 {
611 n_std_work[i] *= wres_norm[i].powi(2);
612 }
613 }
614
615 tempn = n_std_work.clone();
617 if let Some(ref sk_fft) = sphere_k {
618 let tempn_sq: Vec<f32> = tempn.iter().map(|&t| t * t).collect();
619 let smv_tempn_sq = apply_smv_kernel_ws(&tempn_sq, sk_fft, &mut ws);
620 for i in 0..n_total {
621 tempn[i] = (smv_tempn_sq[i] + tempn_sq[i]).sqrt();
622 }
623 }
624
625 m = dataterm_mask_f32(data_weighting, &tempn, &work_mask);
627 b0 = rdf.iter()
628 .zip(m.iter())
629 .map(|(&f, &mi)| {
630 let phase = Complex32::new(0.0, f);
631 mi * phase.exp()
632 })
633 .collect();
634 }
635 }
636 }
637
638 if rel_change < tol_f32 {
639 progress(total_steps, total_steps);
641 break;
642 }
643 }
644
645 let _ = badpoint;
647
648 chi.iter()
650 .zip(mask.iter())
651 .map(|(&c, &m)| if m == 0 { 0.0 } else { c as f64 })
652 .collect()
653}
654
655fn apply_smv_kernel_ws(
657 x: &[f32],
658 sk_fft: &[Complex32],
659 ws: &mut MediWorkspace,
660) -> Vec<f32> {
661 let n_total = ws.n_total;
662
663 for (c, &r) in ws.complex_buf.iter_mut().zip(x.iter()) {
665 *c = Complex32::new(r, 0.0);
666 }
667
668 ws.fft_ws.fft3d(&mut ws.complex_buf);
669
670 for i in 0..n_total {
671 ws.complex_buf[i] *= sk_fft[i];
672 }
673
674 ws.fft_ws.ifft3d(&mut ws.complex_buf);
675
676 ws.complex_buf.iter().map(|c| c.re).collect()
677}
678
679pub(crate) fn compute_rhs_inplace(
683 chi: &[f32],
684 w: &[Complex32],
685 b0: &[Complex32],
686 d_kernel: &[f32],
687 mx: &[f32], my: &[f32], mz: &[f32], vr: &[f32],
691 lambda: f32,
692 rhs: &mut [f32],
693 ws: &mut MediWorkspace,
694) {
695 let n = ws.n_total;
696
697 fgrad_periodic_inplace_f32(
702 &mut ws.gx, &mut ws.gy, &mut ws.gz,
703 chi, ws.nx, ws.ny, ws.nz,
704 ws.vsx, ws.vsy, ws.vsz,
705 );
706
707 apply_gradient_weights_f32(
709 &mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
710 mx, my, mz, vr,
711 &ws.gx, &ws.gy, &ws.gz,
712 );
713
714 bdiv_periodic_inplace_f32(
715 &mut ws.div_buf,
716 &ws.reg_x, &ws.reg_y, &ws.reg_z,
717 ws.nx, ws.ny, ws.nz,
718 ws.vsx, ws.vsy, ws.vsz,
719 );
720
721 for i in 0..n {
724 let diff = w[i] - b0[i];
725 let conj_w = w[i].conj();
726 let neg_i = Complex32::new(0.0, -1.0);
727 ws.complex_buf2[i] = conj_w * neg_i * diff;
728 }
729
730 ws.fft_ws.fft3d(&mut ws.complex_buf2);
732
733 for i in 0..n {
734 ws.complex_buf2[i] *= d_kernel[i];
735 }
736
737 ws.fft_ws.ifft3d(&mut ws.complex_buf2);
738
739 for i in 0..n {
741 ws.dipole_buf[i] = ws.complex_buf2[i].re;
742 }
743
744 combine_terms_f32(rhs, &ws.div_buf, &ws.dipole_buf, lambda);
746}
747
748pub(crate) fn dataterm_mask_f32(mode: i32, n_std: &[f32], mask: &[u8]) -> Vec<f32> {
755 let n = n_std.len();
756
757 if mode == 0 {
758 mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect()
760 } else {
761 let mut w: Vec<f32> = n_std.iter()
763 .zip(mask.iter())
764 .map(|(&n, &m)| {
765 if m != 0 && n > 1e-10 {
766 1.0 / n
767 } else {
768 0.0
769 }
770 })
771 .collect();
772
773 let mask_count = mask.iter().filter(|&&m| m != 0).count() as f32;
775 if mask_count > 0.0 {
776 let sum: f32 = w.iter()
777 .zip(mask.iter())
778 .filter(|(_, &m)| m != 0)
779 .map(|(&wi, _)| wi)
780 .sum();
781 let mean = sum / mask_count;
782
783 if mean > 1e-10 {
784 for i in 0..n {
786 w[i] /= mean;
787 }
788 }
789 }
790
791 for i in 0..n {
793 if mask[i] == 0 {
794 w[i] = 0.0;
795 }
796 }
797
798 w
799 }
800}
801
802pub(crate) fn gradient_mask_f32(
818 magnitude: &[f32],
819 mask: &[u8],
820 nx: usize, ny: usize, nz: usize,
821 vsx: f32, vsy: f32, vsz: f32,
822 percentage: f32,
823) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
824 let n_total = nx * ny * nz;
825
826 let mag_max = magnitude.iter()
828 .zip(mask.iter())
829 .filter(|(_, &m)| m != 0)
830 .map(|(&v, _)| v.abs())
831 .fold(0.0_f32, f32::max);
832
833 let mag_normalized: Vec<f32> = magnitude.iter()
834 .zip(mask.iter())
835 .map(|(&m, &msk)| {
836 if msk != 0 && mag_max > 1e-10 {
837 m / mag_max
838 } else {
839 0.0
840 }
841 })
842 .collect();
843
844 let (gx, gy, gz) = fgrad_linext_f32(&mag_normalized, nx, ny, nz, vsx, vsy, vsz);
846
847 let abs_gx: Vec<f32> = gx.iter().map(|&v| v.abs()).collect();
849 let abs_gy: Vec<f32> = gy.iter().map(|&v| v.abs()).collect();
850 let abs_gz: Vec<f32> = gz.iter().map(|&v| v.abs()).collect();
851
852 let mut all_grads: Vec<f32> = Vec::with_capacity(3 * n_total);
854 for i in 0..n_total {
855 if mask[i] != 0 {
856 all_grads.push(abs_gx[i]);
857 all_grads.push(abs_gy[i]);
858 all_grads.push(abs_gz[i]);
859 }
860 }
861
862 if all_grads.is_empty() {
863 return (vec![1.0; n_total], vec![1.0; n_total], vec![1.0; n_total]);
864 }
865
866 all_grads.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
869 let percentile_idx = ((1.0 - percentage) * (all_grads.len() - 1) as f32) as usize;
870 let threshold = all_grads[percentile_idx.min(all_grads.len() - 1)];
871
872 let mx: Vec<f32> = abs_gx.iter()
875 .zip(mask.iter())
876 .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
877 .collect();
878
879 let my: Vec<f32> = abs_gy.iter()
880 .zip(mask.iter())
881 .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
882 .collect();
883
884 let mz: Vec<f32> = abs_gz.iter()
885 .zip(mask.iter())
886 .map(|(&g, &m)| if m != 0 && g < threshold { 1.0 } else { 0.0 })
887 .collect();
888
889 (mx, my, mz)
890}
891
892pub(crate) fn fgrad_linext_f32(
895 x: &[f32],
896 nx: usize, ny: usize, nz: usize,
897 vsx: f32, vsy: f32, vsz: f32,
898) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
899 let n_total = nx * ny * nz;
900 let mut gx = vec![0.0f32; n_total];
901 let mut gy = vec![0.0f32; n_total];
902 let mut gz = vec![0.0f32; n_total];
903 fgrad_linext_inplace_f32(&mut gx, &mut gy, &mut gz, x, nx, ny, nz, vsx, vsy, vsz);
904 (gx, gy, gz)
905}
906
907#[inline]
910pub(crate) fn fgrad_linext_inplace_f32(
911 gx: &mut [f32], gy: &mut [f32], gz: &mut [f32],
912 x: &[f32],
913 nx: usize, ny: usize, nz: usize,
914 vsx: f32, vsy: f32, vsz: f32,
915) {
916 let hx = 1.0 / vsx;
917 let hy = 1.0 / vsy;
918 let hz = 1.0 / vsz;
919
920 for k in 0..nz {
921 let k_offset = k * nx * ny;
922
923 for j in 0..ny {
924 let j_offset = j * nx;
925
926 for i in 0..nx {
927 let idx = i + j_offset + k_offset;
928 let x_val = x[idx];
929
930 if i + 1 < nx {
933 gx[idx] = (x[idx + 1] - x_val) * hx;
934 } else if i > 0 {
935 gx[idx] = gx[idx - 1];
937 } else {
938 gx[idx] = 0.0;
939 }
940
941 if j + 1 < ny {
942 gy[idx] = (x[i + (j + 1) * nx + k_offset] - x_val) * hy;
943 } else if j > 0 {
944 gy[idx] = gy[i + (j - 1) * nx + k_offset];
945 } else {
946 gy[idx] = 0.0;
947 }
948
949 if k + 1 < nz {
950 gz[idx] = (x[i + j_offset + (k + 1) * nx * ny] - x_val) * hz;
951 } else if k > 0 {
952 gz[idx] = gz[i + j_offset + (k - 1) * nx * ny];
953 } else {
954 gz[idx] = 0.0;
955 }
956 }
957 }
958 }
959}
960
961
962pub(crate) fn fgrad_periodic_inplace_f32(
969 gx: &mut [f32], gy: &mut [f32], gz: &mut [f32],
970 x: &[f32],
971 nx: usize, ny: usize, nz: usize,
972 vsx: f32, vsy: f32, vsz: f32,
973) {
974 let hx = 1.0 / vsx;
975 let hy = 1.0 / vsy;
976 let hz = 1.0 / vsz;
977 let nxny = nx * ny;
978
979 maybe_par_chunks_mut!(gx, nxny)
980 .zip(maybe_par_chunks_mut!(gy, nxny))
981 .zip(maybe_par_chunks_mut!(gz, nxny))
982 .enumerate()
983 .for_each(|(k, ((gx_s, gy_s), gz_s))| {
984 let k_offset = k * nxny;
985
986 for j in 0..ny {
987 let j_offset = j * nx;
988
989 for i in 0..nx {
990 let local = i + j_offset;
991 let idx = local + k_offset;
992 let x_val = x[idx];
993
994 let x_next = if i + 1 < nx { x[idx + 1] } else { x[j_offset + k_offset] };
996 gx_s[local] = (x_next - x_val) * hx;
997
998 let y_next = if j + 1 < ny { x[i + (j + 1) * nx + k_offset] } else { x[i + k_offset] };
1000 gy_s[local] = (y_next - x_val) * hy;
1001
1002 let z_next = if k + 1 < nz { x[i + j_offset + (k + 1) * nxny] } else { x[i + j_offset] };
1004 gz_s[local] = (z_next - x_val) * hz;
1005 }
1006 }
1007 });
1008}
1009
1010pub(crate) fn bdiv_periodic_inplace_f32(
1017 div: &mut [f32],
1018 gx: &[f32], gy: &[f32], gz: &[f32],
1019 nx: usize, ny: usize, nz: usize,
1020 vsx: f32, vsy: f32, vsz: f32,
1021) {
1022 let hx = -1.0 / vsx; let hy = -1.0 / vsy;
1024 let hz = -1.0 / vsz;
1025 let nxny = nx * ny;
1026
1027 maybe_par_chunks_mut!(div, nxny).enumerate().for_each(|(k, div_s)| {
1028 let k_offset = k * nxny;
1029
1030 for j in 0..ny {
1031 let j_offset = j * nx;
1032
1033 for i in 0..nx {
1034 let local = i + j_offset;
1035 let idx = local + k_offset;
1036
1037 let gx_prev = if i > 0 { gx[idx - 1] } else { gx[(nx - 1) + j_offset + k_offset] };
1039 let gx_term = (gx[idx] - gx_prev) * hx;
1040
1041 let gy_prev = if j > 0 { gy[i + (j - 1) * nx + k_offset] } else { gy[i + (ny - 1) * nx + k_offset] };
1043 let gy_term = (gy[idx] - gy_prev) * hy;
1044
1045 let gz_prev = if k > 0 { gz[i + j_offset + (k - 1) * nxny] } else { gz[i + j_offset + (nz - 1) * nxny] };
1047 let gz_term = (gz[idx] - gz_prev) * hz;
1048
1049 div_s[local] = gx_term + gy_term + gz_term;
1050 }
1051 }
1052 });
1053}
1054
1055
1056#[cfg(test)]
1057mod tests {
1058 use super::*;
1059
1060 #[test]
1061 fn test_dataterm_mask_uniform() {
1062 let n_std = vec![1.0f32; 27];
1063 let mask = vec![1u8; 27];
1064
1065 let w = dataterm_mask_f32(0, &n_std, &mask);
1066
1067 for &wi in w.iter() {
1068 assert!((wi - 1.0).abs() < 1e-10);
1069 }
1070 }
1071
1072 #[test]
1073 fn test_dataterm_mask_snr() {
1074 let n_std = vec![2.0f32; 27];
1075 let mask = vec![1u8; 27];
1076
1077 let w = dataterm_mask_f32(1, &n_std, &mask);
1078
1079 let mean: f32 = w.iter().sum::<f32>() / 27.0;
1081 assert!((mean - 1.0).abs() < 1e-5);
1082 }
1083
1084 #[test]
1085 fn test_gradient_mask_constant() {
1086 let mag = vec![1.0f32; 8 * 8 * 8];
1088 let mask = vec![1u8; 8 * 8 * 8];
1089
1090 let (mx, my, mz) = gradient_mask_f32(&mag, &mask, 8, 8, 8, 1.0, 1.0, 1.0, 0.3);
1091
1092 for i in 0..(8 * 8 * 8) {
1094 assert!(mx[i] == 0.0 || mx[i] == 1.0, "mx should be binary, got {}", mx[i]);
1095 assert!(my[i] == 0.0 || my[i] == 1.0, "my should be binary, got {}", my[i]);
1096 assert!(mz[i] == 0.0 || mz[i] == 1.0, "mz should be binary, got {}", mz[i]);
1097 }
1098 }
1099
1100 fn test_medi_params() -> MediParams {
1101 MediParams {
1102 lambda: 1000.0, percentage: 0.9, cg_tol: 0.1,
1103 cg_max_iter: 10, max_iter: 3, tol: 0.1,
1104 ..MediParams::default()
1105 }
1106 }
1107
1108 #[test]
1109 fn test_medi_zero_field() {
1110 let n = 8;
1111 let field = vec![0.0; n * n * n];
1112 let mask = vec![1u8; n * n * n];
1113 let mag = vec![1.0; n * n * n];
1114 let n_std = vec![1.0; n * n * n];
1115 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1116
1117 let chi = medi(
1118 &field, &n_std, &mag, &mask, &grid,
1119 (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1120 );
1121
1122 for &val in chi.iter() {
1123 assert!(val.abs() < 1e-4, "Zero field should give near-zero chi, got {}", val);
1124 }
1125 }
1126
1127 #[test]
1128 fn test_medi_finite() {
1129 let n = 8;
1130 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1131 let mask = vec![1u8; n * n * n];
1132 let mag = vec![1.0; n * n * n];
1133 let n_std = vec![1.0; n * n * n];
1134 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1135
1136 let chi = medi(
1137 &field, &n_std, &mag, &mask, &grid,
1138 (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1139 );
1140
1141 for (i, &val) in chi.iter().enumerate() {
1142 assert!(val.is_finite(), "Chi should be finite at index {}", i);
1143 }
1144 }
1145
1146 #[test]
1147 fn test_medi_with_smv() {
1148 let n = 8;
1149 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1150 let mask = vec![1u8; n * n * n];
1151 let mag = vec![1.0; n * n * n];
1152 let n_std = vec![1.0; n * n * n];
1153 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1154
1155 let params = MediParams { smv: true, smv_radius: 2.0, ..test_medi_params() };
1157 let chi = medi(
1158 &field, &n_std, &mag, &mask, &grid,
1159 (0.0, 0.0, 1.0), ¶ms, |_, _| {},
1160 );
1161
1162 for (i, &val) in chi.iter().enumerate() {
1163 assert!(val.is_finite(), "Chi with SMV should be finite at index {}", i);
1164 }
1165 }
1166
1167 #[test]
1168 fn test_medi_mask() {
1169 let n = 8;
1170 let field: Vec<f64> = (0..n*n*n).map(|i| (i as f64) * 0.001).collect();
1171 let mut mask = vec![1u8; n * n * n];
1172 let mag = vec![1.0; n * n * n];
1173 let n_std = vec![1.0; n * n * n];
1174 mask[0] = 0;
1175 mask[10] = 0;
1176 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1177
1178 let chi = medi(
1179 &field, &n_std, &mag, &mask, &grid,
1180 (0.0, 0.0, 1.0), &test_medi_params(), |_, _| {},
1181 );
1182
1183 assert_eq!(chi[0], 0.0, "Masked voxel should be zero");
1184 assert_eq!(chi[10], 0.0, "Masked voxel should be zero");
1185 }
1186
1187 #[test]
1191 #[ignore]
1192 fn test_medi_debug() {
1193 let data_path = "/home/ashley/OUT/2bgRemoved.nii";
1194 if !std::path::Path::new(data_path).exists() {
1195 eprintln!("Skipping: {} not found", data_path);
1196 return;
1197 }
1198
1199 let outdir = "/home/ashley/OUT/debug";
1200 std::fs::create_dir_all(outdir).ok();
1201
1202 let bytes = std::fs::read(data_path).unwrap();
1204 let nifti_data = crate::io::load_nifti(&bytes).unwrap();
1205 let (nx, ny, nz) = nifti_data.dims;
1206 let (vsx, vsy, vsz) = nifti_data.voxel_size;
1207
1208 let n_total = nx * ny * nz;
1209 let vsx_f32 = vsx as f32;
1210 let vsy_f32 = vsy as f32;
1211 let vsz_f32 = vsz as f32;
1212
1213 eprintln!("Data: {}x{}x{}, voxel: {}x{}x{}", nx, ny, nz, vsx, vsy, vsz);
1214
1215 let local_field: Vec<f32> = nifti_data.data.iter().map(|&v| v as f32).collect();
1217
1218 let mask: Vec<u8> = local_field.iter()
1220 .map(|&v| if v.abs() > 1e-10 { 1 } else { 0 })
1221 .collect();
1222 let mask_count: usize = mask.iter().filter(|&&m| m != 0).count();
1223 eprintln!("Mask voxels: {} / {}", mask_count, n_total);
1224
1225 save_f32_raw(&local_field, &format!("{}/f_rust.raw", outdir));
1227 let mask_f32: Vec<f32> = mask.iter().map(|&m| m as f32).collect();
1228 save_f32_raw(&mask_f32, &format!("{}/mask_rust.raw", outdir));
1229
1230 let lambda: f32 = 7.5e-5;
1232 let beta: f32 = 1.49e-8;
1233 let bdir = (0.0f32, 0.0f32, 1.0f32);
1234 let cg_tol: f32 = 0.1; let cg_max_iter: usize = 10;
1236
1237 let debug_grid = Grid::new(nx, ny, nz, vsx, vsy, vsz);
1239
1240 let d_kernel = crate::kernels::dipole::dipole_kernel_f32(
1242 &debug_grid, bdir,
1243 );
1244 save_f32_raw(&d_kernel, &format!("{}/D_rust.raw", outdir));
1245 eprintln!("D: min={} max={} D[0]={}", fmin(&d_kernel), fmax(&d_kernel), d_kernel[0]);
1246 let mut ws = MediWorkspace::new(&debug_grid);
1247
1248 let m: Vec<f32> = mask.iter().map(|&m| if m != 0 { 1.0 } else { 0.0 }).collect();
1250
1251 let b0: Vec<Complex32> = local_field.iter().zip(m.iter())
1253 .map(|(&f, &mi)| {
1254 let phase = Complex32::new(0.0, f);
1255 mi * phase.exp()
1256 })
1257 .collect();
1258
1259 let w_gx: Vec<f32> = m.clone();
1261 let w_gy: Vec<f32> = m.clone();
1262 let w_gz: Vec<f32> = m.clone();
1263
1264 eprintln!("\n=== Iteration 1 ===");
1266 let mut chi = vec![0.0f32; n_total];
1267 let mut dx = vec![0.0f32; n_total];
1268 let mut rhs = vec![0.0f32; n_total];
1269 let mut vr = vec![0.0f32; n_total];
1270 let mut w: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); n_total];
1271
1272 fgrad_periodic_inplace_f32(
1274 &mut ws.gx, &mut ws.gy, &mut ws.gz,
1275 &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
1276 );
1277 compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
1278 save_f32_raw(&vr, &format!("{}/P1_rust.raw", outdir));
1279 eprintln!("P1: min={} max={} mean={}", fmin(&vr), fmax(&vr), fmean(&vr));
1280
1281 apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
1283 for i in 0..n_total {
1284 let phase = Complex32::new(0.0, ws.dipole_buf[i]);
1285 w[i] = m[i] * phase.exp();
1286 }
1287
1288 compute_rhs_inplace(
1290 &chi, &w, &b0, &d_kernel,
1291 &w_gx, &w_gy, &w_gz, &vr, lambda,
1292 &mut rhs, &mut ws,
1293 );
1294 save_f32_raw(&rhs, &format!("{}/rhs1_rust.raw", outdir));
1295 eprintln!("RHS1: min={} max={} norm={}", fmin(&rhs), fmax(&rhs), fnorm(&rhs));
1296
1297 negate_f32(&mut rhs);
1299
1300 let mut cg_residuals: Vec<f32> = Vec::new();
1302 {
1303 let n = ws.n_total;
1305 let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
1306 let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
1307
1308 dx.fill(0.0);
1309 ws.cg_r.copy_from_slice(&rhs);
1310 ws.cg_p.copy_from_slice(&ws.cg_r);
1311 let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
1312 let b_norm: f32 = norm_squared_f32(&rhs).sqrt();
1313
1314 let mut p_copy = vec![0.0f32; n];
1315 let mut prev_residual = rsold.sqrt();
1316
1317 for cg_iter in 0..cg_max_iter {
1318 let residual_before = rsold.sqrt();
1319 cg_residuals.push(residual_before);
1320
1321 p_copy.copy_from_slice(&ws.cg_p);
1322 {
1323 let mut bufs = MediOpBuffers {
1324 gx: &mut ws.gx, gy: &mut ws.gy, gz: &mut ws.gz,
1325 reg_x: &mut ws.reg_x, reg_y: &mut ws.reg_y, reg_z: &mut ws.reg_z,
1326 div_buf: &mut ws.div_buf, dipole_buf: &mut ws.dipole_buf,
1327 complex_buf: &mut ws.complex_buf, complex_buf2: &mut ws.complex_buf2,
1328 };
1329 apply_medi_operator_core(
1330 &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
1331 &p_copy, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda, &mut ws.cg_ap,
1332 );
1333 }
1334
1335 let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
1336 if pap.abs() < 1e-15 { break; }
1337 let alpha = rsold / pap;
1338
1339 axpy_f32(&mut dx, alpha, &ws.cg_p);
1340 axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
1341
1342 let rsnew: f32 = norm_squared_f32(&ws.cg_r);
1343 let residual = rsnew.sqrt();
1344
1345 eprintln!(" CG iter {}: res={:.6e}, alpha={:.6e}, pap={:.6e}",
1346 cg_iter + 1, residual, alpha, pap);
1347
1348 if residual < cg_tol * b_norm { break; }
1349
1350 let beta_cg = rsnew / rsold;
1353 xpby_f32(&mut ws.cg_p, &ws.cg_r, beta_cg);
1354 rsold = rsnew;
1355 prev_residual = residual;
1356 }
1357 }
1358 save_f32_raw(&dx, &format!("{}/dx1_rust.raw", outdir));
1359 eprintln!("dx1: min={} max={} norm={}", fmin(&dx), fmax(&dx), fnorm(&dx));
1360
1361 axpy_f32(&mut chi, 1.0, &dx);
1363 save_f32_raw(&chi, &format!("{}/chi1_rust.raw", outdir));
1364 eprintln!("chi1: min={} max={} norm={}", fmin(&chi), fmax(&chi), fnorm(&chi));
1365
1366 eprintln!("\n=== Iteration 2 ===");
1368
1369 fgrad_periodic_inplace_f32(
1371 &mut ws.gx, &mut ws.gy, &mut ws.gz,
1372 &chi, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32,
1373 );
1374 compute_p_weights_f32(&mut vr, &w_gx, &w_gy, &w_gz, &ws.gx, &ws.gy, &ws.gz, beta);
1375 save_f32_raw(&vr, &format!("{}/P2_rust.raw", outdir));
1376 eprintln!("P2: min={} max={} mean={}", fmin(&vr), fmax(&vr), fmean(&vr));
1377
1378 apply_dipole_conv(&mut ws.fft_ws, &chi, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
1380 for i in 0..n_total {
1381 let phase = Complex32::new(0.0, ws.dipole_buf[i]);
1382 w[i] = m[i] * phase.exp();
1383 }
1384
1385 compute_rhs_inplace(
1387 &chi, &w, &b0, &d_kernel,
1388 &w_gx, &w_gy, &w_gz, &vr, lambda,
1389 &mut rhs, &mut ws,
1390 );
1391 save_f32_raw(&rhs, &format!("{}/rhs2_rust.raw", outdir));
1392 eprintln!("RHS2: min={} max={} norm={}", fmin(&rhs), fmax(&rhs), fnorm(&rhs));
1393
1394 negate_f32(&mut rhs);
1395
1396 {
1398 let n = ws.n_total;
1399 let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
1400 let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
1401
1402 dx.fill(0.0);
1403 ws.cg_r.copy_from_slice(&rhs);
1404 ws.cg_p.copy_from_slice(&ws.cg_r);
1405 let mut rsold: f32 = norm_squared_f32(&ws.cg_r);
1406 let b_norm: f32 = norm_squared_f32(&rhs).sqrt();
1407 let mut p_copy = vec![0.0f32; n];
1408
1409 for cg_iter in 0..cg_max_iter {
1410 p_copy.copy_from_slice(&ws.cg_p);
1411 {
1412 let mut bufs = MediOpBuffers {
1413 gx: &mut ws.gx, gy: &mut ws.gy, gz: &mut ws.gz,
1414 reg_x: &mut ws.reg_x, reg_y: &mut ws.reg_y, reg_z: &mut ws.reg_z,
1415 div_buf: &mut ws.div_buf, dipole_buf: &mut ws.dipole_buf,
1416 complex_buf: &mut ws.complex_buf, complex_buf2: &mut ws.complex_buf2,
1417 };
1418 apply_medi_operator_core(
1419 &mut ws.fft_ws, &mut bufs, n, nx, ny, nz, vsx, vsy, vsz,
1420 &p_copy, &w, &d_kernel, &w_gx, &w_gy, &w_gz, &vr, lambda, &mut ws.cg_ap,
1421 );
1422 }
1423 let pap: f32 = dot_product_f32(&ws.cg_p, &ws.cg_ap);
1424 if pap.abs() < 1e-15 { break; }
1425 let alpha = rsold / pap;
1426 axpy_f32(&mut dx, alpha, &ws.cg_p);
1427 axpy_f32(&mut ws.cg_r, -alpha, &ws.cg_ap);
1428 let rsnew: f32 = norm_squared_f32(&ws.cg_r);
1429 let residual = rsnew.sqrt();
1430 eprintln!(" CG iter {}: res={:.6e}", cg_iter + 1, residual);
1431 if residual < cg_tol * b_norm { break; }
1432 let beta_cg = rsnew / rsold;
1433 xpby_f32(&mut ws.cg_p, &ws.cg_r, beta_cg);
1434 rsold = rsnew;
1435 }
1436 }
1437 save_f32_raw(&dx, &format!("{}/dx2_rust.raw", outdir));
1438 eprintln!("dx2: min={} max={} norm={}", fmin(&dx), fmax(&dx), fnorm(&dx));
1439
1440 axpy_f32(&mut chi, 1.0, &dx);
1441 save_f32_raw(&chi, &format!("{}/chi2_rust.raw", outdir));
1442 eprintln!("chi2: min={} max={} norm={}", fmin(&chi), fmax(&chi), fnorm(&chi));
1443
1444 eprintln!("\nDone. Intermediates saved to {}", outdir);
1445 }
1446
1447 fn save_f32_raw(data: &[f32], path: &str) {
1448 use std::io::Write;
1449 let mut file = std::fs::File::create(path).unwrap();
1450 for &val in data {
1451 file.write_all(&val.to_le_bytes()).unwrap();
1452 }
1453 }
1454
1455 fn fmin(data: &[f32]) -> f32 { data.iter().cloned().fold(f32::MAX, f32::min) }
1456 fn fmax(data: &[f32]) -> f32 { data.iter().cloned().fold(f32::MIN, f32::max) }
1457 fn fmean(data: &[f32]) -> f32 { data.iter().sum::<f32>() / data.len() as f32 }
1458 fn fnorm(data: &[f32]) -> f32 { data.iter().map(|&v| v * v).sum::<f32>().sqrt() }
1459
1460 #[test]
1461 fn test_medi_small() {
1462 let n = 8;
1464 let n_total = n * n * n;
1465
1466 let mut field = vec![0.0f64; n_total];
1468 let center = n / 2;
1469 for z in 0..n {
1470 for y in 0..n {
1471 for x in 0..n {
1472 let idx = x + y * n + z * n * n;
1473 let dx = (x as f64) - (center as f64);
1474 let dy = (y as f64) - (center as f64);
1475 let dz = (z as f64) - (center as f64);
1476 let r2 = dx*dx + dy*dy + dz*dz;
1477 if r2 > 1.0 {
1478 let r = r2.sqrt();
1480 let cos_theta = dz / r;
1481 field[idx] = (3.0 * cos_theta * cos_theta - 1.0) / (r * r * r) * 0.01;
1482 }
1483 }
1484 }
1485 }
1486
1487 let mask = vec![1u8; n_total];
1488 let mag = vec![1.0f64; n_total];
1489 let n_std = vec![1.0f64; n_total];
1490
1491 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1492
1493 let params = MediParams {
1495 lambda: 1e-4, percentage: 0.3,
1496 cg_tol: 0.01, cg_max_iter: 10, max_iter: 5, tol: 0.1,
1497 ..MediParams::default()
1498 };
1499 let chi = medi(
1500 &field, &n_std, &mag, &mask, &grid,
1501 (0.0, 0.0, 1.0), ¶ms, |_, _| {},
1502 );
1503
1504 assert_eq!(chi.len(), n_total);
1506 for (i, &val) in chi.iter().enumerate() {
1507 assert!(val.is_finite(), "MEDI L1 chi should be finite at index {}", i);
1508 }
1509
1510 let chi_norm: f64 = chi.iter().map(|&v| v * v).sum::<f64>().sqrt();
1512 assert!(chi_norm > 1e-10, "MEDI L1 should produce non-zero susceptibility for dipole field, got norm={}", chi_norm);
1513 }
1514
1515 #[test]
1516 fn test_medi_weight_types() {
1517 let n = 8;
1518 let n_total = n * n * n;
1519
1520 let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1521 let mask = vec![1u8; n_total];
1522 let mag = vec![1.0f64; n_total];
1523 let n_std = vec![1.0f64; n_total];
1524
1525 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1526
1527 let params_uniform = MediParams { data_weighting: 0, ..test_medi_params() };
1529 let chi_uniform = medi(
1530 &field, &n_std, &mag, &mask, &grid,
1531 (0.0, 0.0, 1.0), ¶ms_uniform, |_, _| {},
1532 );
1533
1534 let n_std_varying: Vec<f64> = (0..n_total).map(|i| 0.5 + (i as f64) * 0.01).collect();
1536 let params_snr = MediParams { data_weighting: 1, ..test_medi_params() };
1537 let chi_snr = medi(
1538 &field, &n_std_varying, &mag, &mask, &grid,
1539 (0.0, 0.0, 1.0), ¶ms_snr, |_, _| {},
1540 );
1541
1542 for (i, &val) in chi_uniform.iter().enumerate() {
1544 assert!(val.is_finite(), "Uniform weighting chi should be finite at {}", i);
1545 }
1546 for (i, &val) in chi_snr.iter().enumerate() {
1547 assert!(val.is_finite(), "SNR weighting chi should be finite at {}", i);
1548 }
1549
1550 let diff_norm: f64 = chi_uniform.iter()
1552 .zip(chi_snr.iter())
1553 .map(|(&a, &b)| (a - b).powi(2))
1554 .sum::<f64>()
1555 .sqrt();
1556 assert!(diff_norm.is_finite(), "Difference between weight modes should be finite");
1558 }
1559
1560 #[test]
1561 fn test_medi_with_progress_small() {
1562 let n = 8;
1563 let n_total = n * n * n;
1564
1565 let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1566 let mask = vec![1u8; n_total];
1567 let mag = vec![1.0f64; n_total];
1568 let n_std = vec![1.0f64; n_total];
1569 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1570
1571 let mut progress_calls = 0usize;
1572 let chi = medi(
1573 &field, &n_std, &mag, &mask, &grid,
1574 (0.0, 0.0, 1.0), &test_medi_params(),
1575 |_iter, _max| { progress_calls += 1; },
1576 );
1577
1578 assert_eq!(chi.len(), n_total);
1579 for &val in &chi {
1580 assert!(val.is_finite(), "medi with progress output should be finite");
1581 }
1582 assert!(progress_calls > 0, "progress callback should be called at least once");
1583 }
1584
1585 #[test]
1586 fn test_medi_with_merit() {
1587 let n = 8;
1588 let n_total = n * n * n;
1589
1590 let mut field = vec![0.0f64; n_total];
1591 let center = n / 2;
1592 for z in 0..n {
1593 for y in 0..n {
1594 for x in 0..n {
1595 let idx = x + y * n + z * n * n;
1596 let dx = (x as f64) - (center as f64);
1597 let dy = (y as f64) - (center as f64);
1598 let dz = (z as f64) - (center as f64);
1599 let r2 = dx * dx + dy * dy + dz * dz;
1600 if r2 > 1.0 {
1601 let r = r2.sqrt();
1602 let cos_theta = dz / r;
1603 field[idx] = (3.0 * cos_theta * cos_theta - 1.0) / (r * r * r) * 0.01;
1604 }
1605 }
1606 }
1607 }
1608
1609 let mask = vec![1u8; n_total];
1610 let mag = vec![1.0f64; n_total];
1611 let n_std = vec![1.0f64; n_total];
1612
1613 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1614 let params = MediParams {
1616 lambda: 1e-4, merit: true, percentage: 0.3,
1617 cg_tol: 0.01, cg_max_iter: 10, max_iter: 5, tol: 0.1,
1618 ..MediParams::default()
1619 };
1620 let chi = medi(
1621 &field, &n_std, &mag, &mask, &grid,
1622 (0.0, 0.0, 1.0), ¶ms, |_, _| {},
1623 );
1624
1625 assert_eq!(chi.len(), n_total);
1626 for &val in &chi {
1627 assert!(val.is_finite(), "MEDI with merit should produce finite results");
1628 }
1629 }
1630
1631 #[test]
1632 fn test_medi_with_smv_final() {
1633 let n = 8;
1634 let n_total = n * n * n;
1635
1636 let field: Vec<f64> = (0..n_total).map(|i| (i as f64) * 0.001).collect();
1637 let mask = vec![1u8; n_total];
1638 let mag = vec![1.0f64; n_total];
1639 let n_std = vec![1.0f64; n_total];
1640 let grid = Grid::new(n, n, n, 1.0, 1.0, 1.0);
1641
1642 let params = MediParams {
1644 smv: true, smv_radius: 3.0, percentage: 0.3,
1645 cg_tol: 0.01, cg_max_iter: 10, max_iter: 3, tol: 0.1,
1646 ..MediParams::default()
1647 };
1648 let chi = medi(
1649 &field, &n_std, &mag, &mask, &grid,
1650 (0.0, 0.0, 1.0), ¶ms, |_, _| {},
1651 );
1652
1653 assert_eq!(chi.len(), n_total);
1654 for &val in &chi {
1655 assert!(val.is_finite(), "MEDI with SMV should produce finite results");
1656 }
1657 }
1658}