1use num_complex::Complex32;
25use crate::Grid;
26use crate::fft::Fft3dWorkspaceF32;
27use crate::kernels::dipole::dipole_kernel_f32;
28use crate::inversion::medi::{
29 gradient_mask_f32,
30 fgrad_periodic_inplace_f32,
31 bdiv_periodic_inplace_f32,
32};
33use crate::utils::padding::{next_fast_fft_size, pad3d, unpad3d};
34use crate::utils::simd_ops::{
35 dot_product_f32, norm_squared_f32, axpy_f32, xpby_f32,
36 apply_gradient_weights_f32, compute_p_weights_f32,
37};
38
39struct ChiSepWorkspace {
41 n: usize,
42 nx: usize, ny: usize, nz: usize,
43 vsx: f32, vsy: f32, vsz: f32,
44
45 fft_ws: Fft3dWorkspaceF32,
46
47 gx: Vec<f32>,
48 gy: Vec<f32>,
49 gz: Vec<f32>,
50
51 reg_x: Vec<f32>,
52 reg_y: Vec<f32>,
53 reg_z: Vec<f32>,
54
55 div_buf: Vec<f32>,
56
57 complex_buf: Vec<Complex32>,
58 dipole_buf: Vec<f32>,
59}
60
61impl ChiSepWorkspace {
62 fn new(nx: usize, ny: usize, nz: usize, vsx: f32, vsy: f32, vsz: f32) -> Self {
63 let n = nx * ny * nz;
64 Self {
65 n, nx, ny, nz, vsx, vsy, vsz,
66 fft_ws: Fft3dWorkspaceF32::new(nx, ny, nz),
67 gx: vec![0.0; n],
68 gy: vec![0.0; n],
69 gz: vec![0.0; n],
70 reg_x: vec![0.0; n],
71 reg_y: vec![0.0; n],
72 reg_z: vec![0.0; n],
73 div_buf: vec![0.0; n],
74 complex_buf: vec![Complex32::new(0.0, 0.0); n],
75 dipole_buf: vec![0.0; n],
76 }
77 }
78}
79
80#[cfg_attr(feature = "introspection", derive(serde::Serialize))]
86#[derive(Clone, Debug)]
87pub struct ChiSepParams {
88 pub cf: f64,
90 pub lambda_para: f64,
92 pub lambda_dia: f64,
94 pub lambda_cpl: f64,
96 pub dr_pos: f64,
98 pub dr_neg: f64,
100 pub percentage: f64,
102 pub cg_tol: f64,
104 pub cg_max_iter: usize,
106 pub max_iter: usize,
108 pub tol: f64,
110}
111
112impl Default for ChiSepParams {
113 fn default() -> Self {
114 Self {
115 cf: 123.2e6,
116 lambda_para: 1000.0,
117 lambda_dia: 1000.0,
118 lambda_cpl: 100.0,
119 dr_pos: 114.0,
120 dr_neg: 30.0,
121 percentage: 0.3,
122 cg_tol: 0.01,
123 cg_max_iter: 100,
124 max_iter: 10,
125 tol: 0.1,
126 }
127 }
128}
129
130pub fn chi_sep_medi<F>(
150 local_field: &[f64],
151 r2prime: &[f64],
152 magnitude: &[f64],
153 mask: &[u8],
154 grid: &Grid,
155 bdir: (f64, f64, f64),
156 params: &ChiSepParams,
157 progress: F,
158) -> (Vec<f64>, Vec<f64>, Vec<f64>)
159where
160 F: FnMut(usize, usize),
161{
162 let dims = grid.dims;
163 let fast = (
164 next_fast_fft_size(dims.0),
165 next_fast_fft_size(dims.1),
166 next_fast_fft_size(dims.2),
167 );
168 if fast == dims {
169 return chi_sep_medi_core(local_field, r2prime, magnitude, mask, grid, bdir, params, progress);
170 }
171 let (vsx, vsy, vsz) = grid.voxel_size;
172 let pgrid = Grid::new(fast.0, fast.1, fast.2, vsx, vsy, vsz);
173 let (chi_pos, chi_neg, chi_total) = chi_sep_medi_core(
174 &pad3d(local_field, dims, fast),
175 &pad3d(r2prime, dims, fast),
176 &pad3d(magnitude, dims, fast),
177 &pad3d(mask, dims, fast),
178 &pgrid,
179 bdir,
180 params,
181 progress,
182 );
183 (
184 unpad3d(&chi_pos, fast, dims),
185 unpad3d(&chi_neg, fast, dims),
186 unpad3d(&chi_total, fast, dims),
187 )
188}
189
190#[allow(clippy::too_many_arguments)]
191fn chi_sep_medi_core<F>(
192 local_field: &[f64],
193 r2prime: &[f64],
194 magnitude: &[f64],
195 mask: &[u8],
196 grid: &Grid,
197 bdir: (f64, f64, f64),
198 params: &ChiSepParams,
199 mut progress: F,
200) -> (Vec<f64>, Vec<f64>, Vec<f64>)
201where
202 F: FnMut(usize, usize),
203{
204 let cf = params.cf;
205 let lambda_para = params.lambda_para;
206 let lambda_dia = params.lambda_dia;
207 let lambda_cpl = params.lambda_cpl;
208 let dr_pos = params.dr_pos;
209 let dr_neg = params.dr_neg;
210 let percentage = params.percentage;
211 let cg_tol = params.cg_tol;
212 let cg_max_iter = params.cg_max_iter;
213 let max_iter = params.max_iter;
214 let tol = params.tol;
215 let (nx, ny, nz) = grid.dims;
216 let (vsx, vsy, vsz) = grid.voxel_size;
217 let n = nx * ny * nz;
218 let ppm_factor = (1.0e6 / cf) as f32;
219
220 let vsx_f32 = vsx as f32;
221 let vsy_f32 = vsy as f32;
222 let vsz_f32 = vsz as f32;
223 let bdir_f32 = (bdir.0 as f32, bdir.1 as f32, bdir.2 as f32);
224 let lambda_para_f32 = lambda_para as f32;
225 let lambda_dia_f32 = lambda_dia as f32;
226 let lambda_cpl_f32 = lambda_cpl as f32;
227 let cg_tol_f32 = cg_tol as f32;
228 let tol_f32 = tol as f32;
229
230 let dr_p_eff = ppm_factor * dr_pos as f32;
233 let dr_q_eff = ppm_factor * dr_neg as f32;
234
235 let dr_sum = dr_p_eff + dr_q_eff;
248 let target_eig = 10.0 * lambda_para_f32.max(lambda_dia_f32);
249 let r2_scale = (target_eig / (lambda_cpl_f32 * dr_sum * dr_sum)).sqrt();
250 let dr_p_use = dr_p_eff * r2_scale;
251 let dr_q_use = dr_q_eff * r2_scale;
252
253 let field_f32: Vec<f32> = local_field.iter()
256 .zip(mask.iter())
257 .map(|(&v, &m)| if m != 0 { (v * cf * 1.0e-6) as f32 } else { 0.0 })
258 .collect();
259 let r2p_f32: Vec<f32> = r2prime.iter()
260 .zip(mask.iter())
261 .map(|(&v, &m)| if m != 0 { (v as f32) * r2_scale } else { 0.0 })
262 .collect();
263 let mag_f32: Vec<f32> = magnitude.iter().map(|&v| v as f32).collect();
264
265 let mut ws = ChiSepWorkspace::new(nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
266 let d_kernel = dipole_kernel_f32(grid, bdir_f32);
267 let (mx, my, mz) = gradient_mask_f32(
268 &mag_f32, mask, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32, percentage as f32,
269 );
270
271 let mut chi_pos = vec![0.0f32; n];
273 let mut chi_neg = vec![0.0f32; n];
274
275 let mut vr_pos = vec![0.0f32; n];
276 let mut vr_neg = vec![0.0f32; n];
277 let mut vr_sum = vec![0.0f32; n];
278 let n2 = 2 * n;
279 let mut dx = vec![0.0f32; n2];
280 let mut rhs = vec![0.0f32; n2];
281 let mut chi_sum_buf = vec![0.0f32; n];
282 let mut field_residual = vec![0.0f32; n];
283 let mut r2_residual = vec![0.0f32; n];
284 let mut stage = vec![0.0f32; n];
286
287 let lambda_sum_f32 = 0.0_f32;
290
291 let eps = 1.0e-6_f32;
292
293 for iter in 0..max_iter {
294 progress(iter + 1, max_iter);
295
296 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
298 &chi_pos, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
299 compute_p_weights_f32(&mut vr_pos, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
300
301 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
302 &chi_neg, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
303 compute_p_weights_f32(&mut vr_neg, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
304
305 for i in 0..n {
307 chi_sum_buf[i] = chi_pos[i] + chi_neg[i];
308 }
309 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
310 &chi_sum_buf, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
311 compute_p_weights_f32(&mut vr_sum, &mx, &my, &mz, &ws.gx, &ws.gy, &ws.gz, eps);
312
313 ws.fft_ws.apply_dipole_inplace(&chi_sum_buf, &d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
316 for i in 0..n {
317 field_residual[i] = field_f32[i] - ws.dipole_buf[i];
318 }
319
320 for i in 0..n {
323 r2_residual[i] = r2p_f32[i] - dr_p_use * chi_pos[i] + dr_q_use * chi_neg[i];
324 }
325
326 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
346 &chi_pos, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
347 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
348 &mx, &my, &mz, &vr_pos, &ws.gx, &ws.gy, &ws.gz);
349 bdiv_periodic_inplace_f32(&mut ws.div_buf,
350 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
351 for i in 0..n {
352 rhs[i] = lambda_para_f32 * ws.div_buf[i];
353 }
354
355 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
357 &chi_neg, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
358 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
359 &mx, &my, &mz, &vr_neg, &ws.gx, &ws.gy, &ws.gz);
360 bdiv_periodic_inplace_f32(&mut ws.div_buf,
361 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
362 for i in 0..n {
363 rhs[n + i] = lambda_dia_f32 * ws.div_buf[i];
364 }
365
366 for i in 0..n {
368 chi_sum_buf[i] = chi_pos[i] + chi_neg[i];
369 }
370 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
371 &chi_sum_buf, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
372 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
373 &mx, &my, &mz, &vr_sum, &ws.gx, &ws.gy, &ws.gz);
374 bdiv_periodic_inplace_f32(&mut ws.div_buf,
375 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx_f32, vsy_f32, vsz_f32);
376 for i in 0..n {
377 let tv_sum = lambda_sum_f32 * ws.div_buf[i];
378 rhs[i] += tv_sum;
379 rhs[n + i] += tv_sum;
380 }
381
382 ws.fft_ws.apply_dipole_inplace(&field_residual, &d_kernel,
384 &mut ws.dipole_buf, &mut ws.complex_buf);
385 for i in 0..n {
386 let fg = lambda_cpl_f32 * ws.dipole_buf[i];
387 rhs[i] -= fg;
388 rhs[n + i] -= fg;
389 }
390
391 for i in 0..n {
395 if mask[i] == 0 { continue; }
396 rhs[i] -= lambda_cpl_f32 * r2_residual[i] * dr_p_use;
397 rhs[n + i] += lambda_cpl_f32 * r2_residual[i] * dr_q_use;
398 }
399
400 for v in rhs.iter_mut() {
402 *v = -*v;
403 }
404
405 cg_solve_chisep(
407 &mut ws, &d_kernel,
408 &mx, &my, &mz,
409 &vr_pos, &vr_neg, &vr_sum,
410 lambda_para_f32, lambda_dia_f32, lambda_sum_f32, lambda_cpl_f32,
411 dr_p_use, dr_q_use,
412 mask,
413 &rhs, &mut dx, &mut stage,
414 cg_tol_f32, cg_max_iter,
415 );
416
417 for i in 0..n {
419 chi_pos[i] += 0.5 * dx[i];
420 chi_neg[i] += 0.5 * dx[n + i];
421 }
422
423 for i in 0..n {
425 if mask[i] == 0 {
426 chi_pos[i] = 0.0;
427 chi_neg[i] = 0.0;
428 } else {
429 chi_pos[i] = chi_pos[i].max(0.0);
430 chi_neg[i] = chi_neg[i].min(0.0);
431 }
432 }
433
434 let update_norm = norm_squared_f32(&dx).sqrt();
436 let sol_norm = (norm_squared_f32(&chi_pos) + norm_squared_f32(&chi_neg)).sqrt();
437 let ratio = update_norm / (sol_norm + 1e-6);
438
439 if ratio < tol_f32 {
440 break;
441 }
442 }
443
444 let chi_pos_out: Vec<f64> = chi_pos.iter()
446 .zip(mask.iter())
447 .map(|(&v, &m)| if m == 0 { 0.0 } else { (v * ppm_factor) as f64 })
448 .collect();
449 let chi_neg_out: Vec<f64> = chi_neg.iter()
450 .zip(mask.iter())
451 .map(|(&v, &m)| if m == 0 { 0.0 } else { (v * ppm_factor) as f64 })
452 .collect();
453 let chi_total: Vec<f64> = chi_pos_out.iter()
454 .zip(chi_neg_out.iter())
455 .map(|(&p, &n)| p + n)
456 .collect();
457
458 (chi_pos_out, chi_neg_out, chi_total)
459}
460
461#[allow(clippy::too_many_arguments)]
472fn apply_chisep_operator(
473 ws: &mut ChiSepWorkspace,
474 d_kernel: &[f32],
475 mx: &[f32], my: &[f32], mz: &[f32],
476 vr_pos: &[f32], vr_neg: &[f32], vr_sum: &[f32],
477 lambda_para: f32, lambda_dia: f32, lambda_sum: f32, lambda_cpl: f32,
478 dr_p: f32, dr_q: f32,
479 mask: &[u8],
480 dx: &[f32],
481 out: &mut [f32],
482 stage: &mut [f32],
483) {
484 let n = ws.n;
485 let (nx, ny, nz) = (ws.nx, ws.ny, ws.nz);
486 let (vsx, vsy, vsz) = (ws.vsx, ws.vsy, ws.vsz);
487
488 let d_pos = &dx[..n];
489 let d_neg = &dx[n..];
490
491 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
493 d_pos, nx, ny, nz, vsx, vsy, vsz);
494 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
495 mx, my, mz, vr_pos, &ws.gx, &ws.gy, &ws.gz);
496 bdiv_periodic_inplace_f32(&mut ws.div_buf,
497 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
498 for i in 0..n {
499 out[i] = lambda_para * ws.div_buf[i];
500 }
501
502 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
504 d_neg, nx, ny, nz, vsx, vsy, vsz);
505 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
506 mx, my, mz, vr_neg, &ws.gx, &ws.gy, &ws.gz);
507 bdiv_periodic_inplace_f32(&mut ws.div_buf,
508 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
509 for i in 0..n {
510 out[n + i] = lambda_dia * ws.div_buf[i];
511 }
512
513 for i in 0..n {
517 stage[i] = d_pos[i] + d_neg[i];
518 }
519 fgrad_periodic_inplace_f32(&mut ws.gx, &mut ws.gy, &mut ws.gz,
520 stage, nx, ny, nz, vsx, vsy, vsz);
521 apply_gradient_weights_f32(&mut ws.reg_x, &mut ws.reg_y, &mut ws.reg_z,
522 mx, my, mz, vr_sum, &ws.gx, &ws.gy, &ws.gz);
523 bdiv_periodic_inplace_f32(&mut ws.div_buf,
524 &ws.reg_x, &ws.reg_y, &ws.reg_z, nx, ny, nz, vsx, vsy, vsz);
525 for i in 0..n {
526 let tv_s = lambda_sum * ws.div_buf[i];
527 out[i] += tv_s;
528 out[n + i] += tv_s;
529 }
530
531 ws.fft_ws.apply_dipole_inplace(stage, d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
533 stage.copy_from_slice(&ws.dipole_buf);
534 ws.fft_ws.apply_dipole_inplace(stage, d_kernel, &mut ws.dipole_buf, &mut ws.complex_buf);
535
536 for i in 0..n {
537 let ff = lambda_cpl * ws.dipole_buf[i];
538 out[i] += ff;
539 out[n + i] += ff;
540 }
541
542 for i in 0..n {
547 if mask[i] == 0 { continue; }
548 let r2_lin = dr_p * d_pos[i] - dr_q * d_neg[i];
549 out[i] += lambda_cpl * dr_p * r2_lin;
550 out[n + i] -= lambda_cpl * dr_q * r2_lin;
551 }
552}
553
554#[allow(clippy::too_many_arguments)]
556fn cg_solve_chisep(
557 ws: &mut ChiSepWorkspace,
558 d_kernel: &[f32],
559 mx: &[f32], my: &[f32], mz: &[f32],
560 vr_pos: &[f32], vr_neg: &[f32], vr_sum: &[f32],
561 lambda_para: f32, lambda_dia: f32, lambda_sum: f32, lambda_cpl: f32,
562 dr_p: f32, dr_q: f32,
563 mask: &[u8],
564 b: &[f32],
565 x: &mut [f32],
566 stage: &mut [f32],
567 tol: f32,
568 max_iter: usize,
569) {
570 let n2 = 2 * ws.n;
571 x.fill(0.0);
572
573 let mut cg_r = vec![0.0f32; n2];
574 let mut cg_p = vec![0.0f32; n2];
575 let mut cg_ap = vec![0.0f32; n2];
576
577 cg_r.copy_from_slice(&b[..n2]);
578 cg_p.copy_from_slice(&cg_r);
579
580 let mut rsold = dot_product_f32(&cg_r, &cg_r);
581 let b_norm = dot_product_f32(b, b).sqrt();
582
583 if b_norm < 1e-10 {
584 return;
585 }
586
587 for _cg_iter in 0..max_iter {
588 apply_chisep_operator(
589 ws, d_kernel, mx, my, mz,
590 vr_pos, vr_neg, vr_sum,
591 lambda_para, lambda_dia, lambda_sum, lambda_cpl,
592 dr_p, dr_q,
593 mask,
594 &cg_p, &mut cg_ap, stage,
595 );
596
597 let pap = dot_product_f32(&cg_p, &cg_ap);
598 if pap.abs() < 1e-15 {
599 break;
600 }
601
602 let alpha = rsold / pap;
603 axpy_f32(x, alpha, &cg_p);
604 axpy_f32(&mut cg_r, -alpha, &cg_ap);
605
606 let rsnew = dot_product_f32(&cg_r, &cg_r);
607 if rsnew.sqrt() < tol * b_norm {
608 break;
609 }
610
611 let beta_cg = rsnew / rsold;
612 xpby_f32(&mut cg_p, &cg_r, beta_cg);
613 rsold = rsnew;
614 }
615}
616
617#[cfg(test)]
618mod tests {
619 use super::*;
620 use crate::Grid;
621 use crate::kernels::dipole::dipole_kernel;
622 use crate::fft::{fft3d_real, ifft3d_real};
623
624 fn make_sphere(nx: usize, ny: usize, nz: usize, cx: f64, cy: f64, cz: f64, r: f64) -> Vec<f64> {
625 let mut vol = vec![0.0; nx * ny * nz];
626 for k in 0..nz {
627 for j in 0..ny {
628 for i in 0..nx {
629 let dx = i as f64 - cx;
630 let dy = j as f64 - cy;
631 let dz = k as f64 - cz;
632 if dx * dx + dy * dy + dz * dz <= r * r {
633 vol[i + j * nx + k * nx * ny] = 1.0;
634 }
635 }
636 }
637 }
638 vol
639 }
640
641 #[test]
642 fn test_chi_sep_medi_basic() {
643 let (nx, ny, nz) = (32, 32, 32);
644 let n = nx * ny * nz;
645 let grid = Grid::new(nx, ny, nz, 1.0, 1.0, 1.0);
646 let bdir = (0.0, 0.0, 1.0);
647 let cf: f64 = 123.2e6; let chi_pos_true_ppm = 0.05;
650 let chi_neg_true_ppm = -0.03;
651
652 let sphere_inner = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 4.0);
653 let sphere_outer = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 8.0);
654 let brain_mask = make_sphere(nx, ny, nz, 16.0, 16.0, 16.0, 12.0);
655
656 let mut chi_pos_ppm = vec![0.0f64; n];
657 let mut chi_neg_ppm = vec![0.0f64; n];
658 for i in 0..n {
659 if sphere_inner[i] > 0.5 {
660 chi_pos_ppm[i] = chi_pos_true_ppm;
661 }
662 if sphere_outer[i] > 0.5 && sphere_inner[i] < 0.5 {
663 chi_neg_ppm[i] = chi_neg_true_ppm;
664 }
665 }
666
667 let chi_total_ppm: Vec<f64> = chi_pos_ppm.iter()
669 .zip(chi_neg_ppm.iter())
670 .map(|(&p, &n)| p + n)
671 .collect();
672 let d = dipole_kernel(&grid, bdir);
673 let chi_fft = fft3d_real(&chi_total_ppm, nx, ny, nz);
674 let field_fft: Vec<_> = chi_fft.iter()
675 .zip(d.iter())
676 .map(|(&c, &dk)| c * dk)
677 .collect();
678 let local_field = ifft3d_real(&field_fft, nx, ny, nz);
679
680 let dr_pos: f64 = 114.0;
682 let dr_neg: f64 = 30.0;
683 let r2prime: Vec<f64> = (0..n).map(|i| {
684 dr_pos * chi_pos_ppm[i].abs() + dr_neg * chi_neg_ppm[i].abs()
685 }).collect();
686
687 let mask: Vec<u8> = brain_mask.iter()
688 .map(|&v| if v > 0.5 { 1 } else { 0 })
689 .collect();
690
691 let magnitude: Vec<f64> = (0..n).map(|i| {
692 if mask[i] == 0 { return 0.0; }
693 let base = 100.0;
694 if sphere_inner[i] > 0.5 {
695 base * 1.5
696 } else if sphere_outer[i] > 0.5 {
697 base * 0.7
698 } else {
699 base
700 }
701 }).collect();
702
703 let params = ChiSepParams {
704 cf,
705 lambda_para: 1000.0, lambda_dia: 1000.0, lambda_cpl: 100.0,
706 dr_pos, dr_neg,
707 percentage: 0.3, cg_tol: 0.01, cg_max_iter: 100, max_iter: 10, tol: 0.1,
708 };
709 let (chi_pos_out, chi_neg_out, chi_total_out) = chi_sep_medi(
710 &local_field, &r2prime, &magnitude, &mask,
711 &grid, bdir, ¶ms,
712 |_, _| {},
713 );
714
715 for i in 0..n {
717 if mask[i] != 0 {
718 assert!(chi_pos_out[i] >= -1e-10,
719 "chi+ should be non-negative at voxel {}, got {}", i, chi_pos_out[i]);
720 assert!(chi_neg_out[i] <= 1e-10,
721 "chi- should be non-positive at voxel {}, got {}", i, chi_neg_out[i]);
722 }
723 }
724
725 for i in 0..n {
726 let diff = (chi_total_out[i] - chi_pos_out[i] - chi_neg_out[i]).abs();
727 assert!(diff < 1e-10, "chi_total != chi+ + chi- at voxel {}", i);
728 }
729
730 let pos_max = chi_pos_out.iter().cloned().fold(0.0_f64, f64::max);
731 let neg_min = chi_neg_out.iter().cloned().fold(0.0_f64, f64::min);
732 assert!(pos_max > 0.0, "chi+ should have positive values, max={}", pos_max);
733 assert!(neg_min < 0.0, "chi- should have negative values, min={}", neg_min);
734 }
735}