1use crate::Grid;
40#[cfg(feature = "parallel")]
41use rayon::prelude::*;
42
43use std::f64::consts::PI;
44
45const GAMMA_BAR: f64 = 42.577e6;
47
48const CHI_I: f64 = -0.06e-6; const CHI_A: f64 = -0.10e-6; const E_EXCH: f64 = 0.02e-6; const G_RATIO: f64 = 0.7;
53const T2_M: f64 = 10e-3; const T2_A: f64 = 64e-3; const T2_E: f64 = 48e-3; const F_AXON: f64 = 0.55; const CHI_NEG_REF: f64 = 0.038; const MWF_REF: f64 = 0.12;
61const MWF_MIN: f64 = 0.03;
62const MWF_MAX: f64 = 0.25;
63
64const SMOOTH_CLS: [f64; 2] = [0.75, 1.0]; const SMOOTH_FIT: f64 = 1.0;
67const W_CENTER: f64 = 1.5;
68const W_SCALE: f64 = 0.3;
69const CHI_GATE_C: f64 = 0.01;
70const CHI_GATE_S: f64 = 0.01;
71const LAM: f64 = 0.7;
72const NCONV_SIGMA: f64 = 1.5;
73const W_SUPPORT: f64 = 0.02;
74const NO_BEAT_FRAC: f64 = 0.01;
75
76const K_CHI: f64 = CHI_NEG_REF / MWF_REF;
78
79const NT: usize = 61; const NM: usize = 45; #[cfg_attr(feature = "introspection", derive(serde::Serialize))]
85#[derive(Clone, Debug)]
86pub struct HcChisepParams {
87 pub b0: f64,
89 pub se_echo_times: Vec<f64>,
91 pub dr_pos_3t: f64,
93 pub bin_hz: f64,
95}
96
97impl Default for HcChisepParams {
98 fn default() -> Self {
99 Self {
100 b0: 7.0,
101 se_echo_times: Vec::new(),
102 dr_pos_3t: 137.0,
103 bin_hz: 0.25,
104 }
105 }
106}
107
108#[allow(clippy::too_many_arguments)]
129pub fn hc_chisep(
130 chi_total: &[f64],
131 r2prime: &[f64],
132 magnitude: &[f64],
133 echo_times: &[f64],
134 se_magnitude: Option<&[f64]>,
135 mask: &[u8],
136 grid: &Grid,
137 params: &HcChisepParams,
138 mut progress: impl FnMut(usize, usize),
139) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
140 let dims = grid.dims;
141 let n = dims.0 * dims.1 * dims.2;
142 let ne = echo_times.len();
143 assert_eq!(chi_total.len(), n, "chi_total length must match grid");
144 assert_eq!(r2prime.len(), n, "r2prime length must match grid");
145 assert_eq!(
146 magnitude.len(),
147 n * ne,
148 "magnitude must be n_voxels * n_echoes"
149 );
150 assert_eq!(mask.len(), n, "mask length must match grid");
151
152 let b0 = params.b0;
153
154 let mag_em = to_echo_major(magnitude, n, ne);
157 let se_e_len = params.se_echo_times.len();
158 let se_em: Option<Vec<f64>> = match se_magnitude {
159 Some(s) if se_e_len >= 2 && s.len() == n * se_e_len => Some(to_echo_major(s, n, se_e_len)),
160 _ => None,
161 };
162
163 let n_stages = 5;
164 let mut stage = 0usize;
165 let mut tick = |s: &mut usize, p: &mut dyn FnMut(usize, usize)| {
166 *s += 1;
167 p(*s, n_stages);
168 };
169
170 let dr_default = params.dr_pos_3t * b0 / 3.0;
172 let dr_pos = {
173 let ratios: Vec<f64> = (0..n)
174 .filter(|&i| mask[i] != 0 && chi_total[i] > 0.02)
175 .map(|i| r2prime[i] / chi_total[i])
176 .collect();
177 if ratios.len() > 5000 {
178 let p5 = percentile(&ratios, 5.0);
179 if (0.3 * dr_default..=3.0 * dr_default).contains(&p5) {
180 p5
181 } else {
182 dr_default
183 }
184 } else {
185 dr_default
186 }
187 };
188 tick(&mut stage, &mut progress);
189
190 let mut cf_pos = vec![0.0_f64; n];
195 let mut cf_neg = vec![0.0_f64; n];
196 for i in 0..n {
197 let p = (r2prime[i] / dr_pos).max(0.0);
198 let neg = p - chi_total[i];
199 if neg < 0.0 {
200 cf_pos[i] = chi_total[i].max(0.0);
201 } else {
202 cf_pos[i] = p;
203 }
204 cf_neg[i] = (cf_pos[i] - chi_total[i]).max(0.0);
205 }
206 tick(&mut stage, &mut progress);
207
208 let fitter = AnchoredFitter::new(echo_times, b0, ¶ms.se_echo_times);
210 let have_se = se_em.is_some() && fitter.seml.is_some();
211 let se_e = params.se_echo_times.len();
212
213 let bii: Vec<usize> = (0..n).filter(|&i| mask[i] != 0).collect();
215 if bii.is_empty() {
216 return finish(&cf_pos, &cf_neg, mask, n);
217 }
218 let rho_b: Vec<f64> = bii.iter().map(|&i| r2prime[i]).collect();
219
220 let maskf: Vec<f64> = mask.iter().map(|&m| m as f64).collect();
224 let mut lr_best: Option<Vec<f64>> = None;
225 for &s in &SMOOTH_CLS {
226 let m_s = smooth_stack(&mag_em, dims, ne, s);
227 let sn = gather_normalised(&m_s, &bii, ne);
228 let (_t, _m, sse_hc) = fitter.fit(&sn, &rho_b, None, None, params.bin_hz);
229 let mut num = sse_hc;
230 let mut den: Vec<f64> = (0..bii.len())
231 .map(|v| mono_sse(&sn[v * ne..v * ne + ne], echo_times))
232 .collect();
233 if have_se {
234 let se_s = smooth_stack(se_em.as_ref().unwrap(), dims, se_e, s);
235 let sen = gather_normalised(&se_s, &bii, se_e);
236 let seml = fitter.seml.as_ref().unwrap();
237 for v in 0..bii.len() {
238 let mut best = f64::INFINITY;
240 for mi in 0..NM {
241 let mut acc = 0.0;
242 for e in 0..se_e {
243 let d = sen[v * se_e + e] - seml[mi * se_e + e];
244 acc += d * d;
245 }
246 best = best.min(acc);
247 }
248 num[v] += best;
249 den[v] += mono_sse(&sen[v * se_e..v * se_e + se_e], ¶ms.se_echo_times);
250 }
251 }
252 let lr: Vec<f64> = (0..bii.len())
253 .map(|v| (num[v] / den[v].max(1e-12)).max(1e-6).log10())
254 .collect();
255 lr_best = Some(match lr_best {
256 None => lr,
257 Some(prev) => (0..bii.len()).map(|v| prev[v].min(lr[v])).collect(),
258 });
259 }
260 let lr_best = lr_best.unwrap();
261
262 let lr_vol = volmap(&lr_best, &bii, n, 10.0);
264 let lr_med = median_filter_3(&lr_vol, dims);
265
266 let chi_masked: Vec<f64> = (0..n).map(|i| chi_total[i] * maskf[i]).collect();
268 let chs_num = gaussian_filter_3d(&chi_masked, dims, 1.0);
269 let chs_den = gaussian_filter_3d(&maskf, dims, 1.0);
270 let chs: Vec<f64> = (0..n).map(|i| chs_num[i] / chs_den[i].max(1e-6)).collect();
271
272 let mut w = vec![0.0_f64; n];
274 for i in 0..n {
275 let s1 = 1.0 / (1.0 + ((lr_med[i] - W_CENTER) / W_SCALE).exp());
276 let s2 = 1.0 / (1.0 + ((chs[i] - CHI_GATE_C) / CHI_GATE_S).exp());
277 w[i] = s1 * s2 * maskf[i];
278 }
279 let beat_frac = {
280 let cnt = bii.iter().filter(|&&i| w[i] > 0.5).count();
281 cnt as f64 / bii.len() as f64
282 };
283 tick(&mut stage, &mut progress);
284 if beat_frac < NO_BEAT_FRAC {
285 return finish(&cf_pos, &cf_neg, mask, n);
287 }
288
289 let sii: Vec<usize> = bii.iter().cloned().filter(|&i| w[i] > W_SUPPORT).collect();
291 let rho_s: Vec<f64> = sii.iter().map(|&i| r2prime[i]).collect();
292
293 let m_fit = smooth_stack(&mag_em, dims, ne, SMOOTH_FIT);
294 let sn_sm = gather_normalised(&m_fit, &sii, ne);
295 let sn_raw = gather_normalised(&mag_em, &sii, ne);
296 let (sen_sm, sen_raw) = if have_se {
297 let se = se_em.as_ref().unwrap();
298 let se_fit = smooth_stack(se, dims, se_e, SMOOTH_FIT);
299 (
300 Some(gather_normalised(&se_fit, &sii, se_e)),
301 Some(gather_normalised(se, &sii, se_e)),
302 )
303 } else {
304 (None, None)
305 };
306
307 let (th_sm, _m_sm, _s_sm) = fitter.fit(&sn_sm, &rho_s, sen_sm.as_deref(), None, params.bin_hz);
309 let th_fill = volmap(&th_sm, &sii, n, 45.0);
311 let th_reg = median_filter_3(&th_fill, dims);
312 let th_pin: Vec<f64> = sii.iter().map(|&i| th_reg[i]).collect();
313
314 let (_t1, mw_raw, _s1) = fitter.fit(
316 &sn_raw,
317 &rho_s,
318 sen_raw.as_deref(),
319 Some(&th_pin),
320 params.bin_hz,
321 );
322 let (_t2, mw_sm, _s2) = fitter.fit(
323 &sn_sm,
324 &rho_s,
325 sen_sm.as_deref(),
326 Some(&th_pin),
327 params.bin_hz,
328 );
329 let mw_est: Vec<f64> = (0..sii.len())
330 .map(|v| {
331 let bnd = (-chi_total[sii[v]] / K_CHI).clamp(MWF_MIN, MWF_MAX);
332 0.5 * (mw_raw[v].max(bnd) + mw_sm[v].max(bnd))
333 })
334 .collect();
335 let mw_vol = volmap(&mw_est, &sii, n, 0.0);
336
337 let wv: Vec<f64> = w.iter().map(|&x| x.clamp(0.0, 1.0)).collect();
339 let mw_w: Vec<f64> = (0..n).map(|i| mw_vol[i] * wv[i]).collect();
340 let mw_num = gaussian_filter_3d(&mw_w, dims, NCONV_SIGMA);
341 let mw_den = gaussian_filter_3d(&wv, dims, NCONV_SIGMA);
342 let mw_reg: Vec<f64> = (0..n).map(|i| mw_num[i] / mw_den[i].max(1e-3)).collect();
343 tick(&mut stage, &mut progress);
344
345 let route: Vec<f64> = (0..n).map(|i| chi_total[i] + K_CHI * mw_reg[i]).collect();
347 let conf: Vec<usize> = (0..n).filter(|&i| w[i] > 0.5).collect();
348 let c0 = if conf.len() > 100 {
349 let vals: Vec<f64> = conf.iter().map(|&i| route[i].max(0.0)).collect();
350 median(&vals)
351 } else {
352 0.005
353 };
354 let mut chi_pos = vec![0.0_f64; n];
355 let mut chi_neg = vec![0.0_f64; n];
356 let mut chi_out = vec![0.0_f64; n];
357 for i in 0..n {
358 if mask[i] == 0 {
359 continue;
360 }
361 let wm_pos = (LAM * route[i] + (1.0 - LAM) * c0).max(0.0);
362 let wm_neg = (wm_pos - chi_total[i]).max(0.0);
363 let pos = w[i] * wm_pos + (1.0 - w[i]) * cf_pos[i];
364 let neg = w[i] * wm_neg + (1.0 - w[i]) * cf_neg[i];
365 chi_pos[i] = pos;
366 chi_neg[i] = -neg; chi_out[i] = pos - neg;
368 }
369 tick(&mut stage, &mut progress);
370 (chi_pos, chi_neg, chi_out)
371}
372
373fn finish(cf_pos: &[f64], cf_neg: &[f64], mask: &[u8], n: usize) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
375 let mut chi_pos = vec![0.0_f64; n];
376 let mut chi_neg = vec![0.0_f64; n];
377 let mut chi_out = vec![0.0_f64; n];
378 for i in 0..n {
379 if mask[i] != 0 {
380 chi_pos[i] = cf_pos[i];
381 chi_neg[i] = -cf_neg[i];
382 chi_out[i] = cf_pos[i] - cf_neg[i];
383 }
384 }
385 (chi_pos, chi_neg, chi_out)
386}
387
388fn hc_compartment_freqs(theta: f64, b0: f64) -> (f64, f64) {
395 let s2 = theta.sin() * theta.sin();
396 let w0 = GAMMA_BAR * b0;
397 let ln_term = 0.75 * CHI_A * (1.0 / G_RATIO).ln() * s2;
398 let f_my = w0
399 * (CHI_I * (2.0 / 3.0 - s2) / 2.0
400 + CHI_A * (1.0 / 12.0 - 5.0 / 12.0 * s2)
401 + ln_term
402 + E_EXCH);
403 let f_ax = w0 * ln_term;
404 (f_my, f_ax)
405}
406
407fn hc_wm_signal_mag(te: f64, theta: f64, b0: f64, mwf: f64, r2p_meso: f64) -> f64 {
410 let f_m = mwf;
411 let rest = 1.0 - f_m;
412 let f_a = rest * F_AXON;
413 let f_e = rest * (1.0 - F_AXON);
414 let (dfm, dfa) = hc_compartment_freqs(theta, b0);
415 let (mut re, mut im) = (0.0, 0.0);
416 let terms = [(f_m, T2_M, dfm), (f_a, T2_A, dfa), (f_e, T2_E, 0.0)];
417 for (f, t2, df) in terms {
418 let amp = f * (-te / t2).exp();
419 let ph = 2.0 * PI * df * te;
420 re += amp * ph.cos();
421 im += amp * ph.sin();
422 }
423 let mag = (re * re + im * im).sqrt();
424 mag * (-r2p_meso * te).exp()
425}
426
427fn hc_wm_se_signal(te: f64, mwf: f64) -> f64 {
429 let rest = 1.0 - mwf;
430 mwf * (-te / T2_M).exp()
431 + rest * F_AXON * (-te / T2_A).exp()
432 + rest * (1.0 - F_AXON) * (-te / T2_E).exp()
433}
434
435fn hc_wm_r2prime(theta: f64, mwf: f64, tes: &[f64], b0: f64) -> f64 {
438 let tbar = tes.iter().sum::<f64>() / tes.len() as f64;
439 let denom: f64 = tes.iter().map(|&t| (t - tbar) * (t - tbar)).sum();
440 let mut acc = 0.0;
441 for &te in tes {
442 let gre = hc_wm_signal_mag(te, theta, b0, mwf, 0.0);
443 let se = hc_wm_se_signal(te, mwf).max(1e-12);
444 let ratio = (gre / se).max(1e-12);
445 acc += ratio.ln() * (te - tbar);
446 }
447 (-acc / denom).max(0.0)
448}
449
450struct AnchoredFitter {
452 e: usize,
453 pn: Vec<f64>,
455 h: Vec<f64>,
457 dte: Vec<f64>,
459 seml: Option<Vec<f64>>,
461}
462
463impl AnchoredFitter {
464 fn new(tes: &[f64], b0: f64, se_tes: &[f64]) -> Self {
465 let e = tes.len();
466 let mut pn = vec![0.0_f64; NT * NM * e];
467 let mut h = vec![0.0_f64; NT * NM];
468 for ti in 0..NT {
469 let theta = (ti as f64 * 1.5).to_radians();
470 for mi in 0..NM {
471 let mwf = MWF_MIN + mi as f64 * 0.005;
472 let base = (ti * NM + mi) * e;
473 let first = hc_wm_signal_mag(tes[0], theta, b0, mwf, 0.0).max(1e-30);
474 for (k, &te) in tes.iter().enumerate() {
475 pn[base + k] = hc_wm_signal_mag(te, theta, b0, mwf, 0.0) / first;
476 }
477 h[ti * NM + mi] = hc_wm_r2prime(theta, mwf, tes, b0);
478 }
479 }
480 let dte: Vec<f64> = tes.iter().map(|&t| t - tes[0]).collect();
481 let seml = if se_tes.len() >= 2 {
482 let se_e = se_tes.len();
483 let mut m = vec![0.0_f64; NM * se_e];
484 for mi in 0..NM {
485 let mwf = MWF_MIN + mi as f64 * 0.005;
486 let first = hc_wm_se_signal(se_tes[0], mwf).max(1e-30);
487 for (k, &te) in se_tes.iter().enumerate() {
488 m[mi * se_e + k] = hc_wm_se_signal(te, mwf) / first;
489 }
490 }
491 Some(m)
492 } else {
493 None
494 };
495 Self {
496 e,
497 pn,
498 h,
499 dte,
500 seml,
501 }
502 }
503
504 fn fit(
507 &self,
508 sig_n: &[f64],
509 rho: &[f64],
510 se_n: Option<&[f64]>,
511 theta_pin: Option<&[f64]>,
512 bin_hz: f64,
513 ) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
514 let e = self.e;
515 let n = rho.len();
516 let l = NT * NM;
517
518 let se_e = se_n.map(|s| s.len() / n).unwrap_or(0);
520 let sse_se: Option<Vec<f64>> = match (se_n, &self.seml) {
521 (Some(sn), Some(seml)) => {
522 let mut out = vec![0.0_f64; n * NM];
523 for v in 0..n {
524 for mi in 0..NM {
525 let mut acc = 0.0;
526 for k in 0..se_e {
527 let d = sn[v * se_e + k] - seml[mi * se_e + k];
528 acc += d * d;
529 }
530 out[v * NM + mi] = acc;
531 }
532 }
533 Some(out)
534 }
535 _ => None,
536 };
537
538 use std::collections::HashMap;
540 let mut groups: HashMap<i64, Vec<usize>> = HashMap::new();
541 for (v, &r) in rho.iter().enumerate() {
542 let b = (r.clamp(0.0, 200.0) / bin_hz).round() as i64;
543 groups.entry(b).or_default().push(v);
544 }
545
546 let group_vec: Vec<(i64, Vec<usize>)> = groups.into_iter().collect();
547 let per_bin: Vec<Vec<(usize, f64, f64, f64)>> = maybe_par_iter!(group_vec)
549 .map(|(b, vs)| {
550 let mut flat = vec![0.0_f64; l * e];
552 let mut m2 = vec![0.0_f64; l];
553 let mut bad = vec![false; l];
554 let rho_bin = *b as f64 * bin_hz;
555 for li in 0..l {
556 let meso = rho_bin - self.h[li];
557 bad[li] = meso < -1.0;
558 let decay = meso.max(0.0);
559 let mut acc = 0.0;
560 for k in 0..e {
561 let val = self.pn[li * e + k] * (-decay * self.dte[k]).exp();
562 flat[li * e + k] = val;
563 acc += val * val;
564 }
565 m2[li] = 0.5 * acc;
566 }
567 let all_bad = bad.iter().all(|&x| x);
568
569 let mut out = Vec::with_capacity(vs.len());
570 for &v in vs {
571 let sig = &sig_n[v * e..v * e + e];
572 let se_row = sse_se.as_ref().map(|s| &s[v * NM..v * NM + NM]);
573 let pin_row = theta_pin.map(|tp| {
574 ((tp[v] / 1.5).round() as isize).clamp(0, NT as isize - 1) as usize
575 });
576
577 let (mut best_score, mut best_l) = (f64::NEG_INFINITY, 0usize);
578 let (lo, hi) = match pin_row {
579 Some(r) => (r * NM, r * NM + NM),
580 None => (0, l),
581 };
582 for li in lo..hi {
583 if bad[li] && !all_bad {
584 continue;
585 }
586 let mut s = -m2[li];
587 for k in 0..e {
588 s += sig[k] * flat[li * e + k];
589 }
590 if let Some(se_row) = se_row {
591 s -= 0.5 * se_row[li % NM];
592 }
593 if s > best_score {
594 best_score = s;
595 best_l = li;
596 }
597 }
598 let mut sse = 0.0;
600 for k in 0..e {
601 let d = sig[k] - flat[best_l * e + k];
602 sse += d * d;
603 }
604 let theta_deg = (best_l / NM) as f64 * 1.5;
605 let mwf = MWF_MIN + (best_l % NM) as f64 * 0.005;
606 out.push((v, theta_deg, mwf, sse));
607 }
608 out
609 })
610 .collect();
611
612 let mut out_t = vec![0.0_f64; n];
613 let mut out_m = vec![MWF_REF; n];
614 let mut out_s = vec![f64::INFINITY; n];
615 for bin in per_bin {
616 for (v, t, m, s) in bin {
617 out_t[v] = t;
618 out_m[v] = m;
619 out_s[v] = s;
620 }
621 }
622 (out_t, out_m, out_s)
623 }
624}
625
626fn mono_sse(sig: &[f64], tes: &[f64]) -> f64 {
628 let e = tes.len();
629 let tbar = tes.iter().sum::<f64>() / e as f64;
630 let logs: Vec<f64> = sig.iter().map(|&s| s.max(1e-9).ln()).collect();
631 let den: f64 = tes.iter().map(|&t| (t - tbar) * (t - tbar)).sum();
632 let b: f64 = (0..e).map(|k| logs[k] * (tes[k] - tbar)).sum::<f64>() / den;
633 let a: f64 = logs.iter().sum::<f64>() / e as f64;
634 (0..e)
635 .map(|k| {
636 let pred = (a + b * (tes[k] - tbar)).exp();
637 let d = sig[k] - pred;
638 d * d
639 })
640 .sum()
641}
642
643fn to_echo_major(vm: &[f64], n: usize, ne: usize) -> Vec<f64> {
646 let mut out = vec![0.0_f64; n * ne];
647 for i in 0..n {
648 for e in 0..ne {
649 out[e * n + i] = vm[i * ne + e];
650 }
651 }
652 out
653}
654
655fn gather_normalised(stack: &[f64], idx: &[usize], ne: usize) -> Vec<f64> {
658 let n = stack.len() / ne;
659 let mut out = vec![0.0_f64; idx.len() * ne];
660 for (k, &i) in idx.iter().enumerate() {
661 let first = stack[i].max(1e-9); for e in 0..ne {
663 out[k * ne + e] = stack[e * n + i] / first;
664 }
665 }
666 out
667}
668
669fn volmap(vals: &[f64], idx: &[usize], n: usize, fill: f64) -> Vec<f64> {
671 let mut out = vec![fill; n];
672 for (k, &i) in idx.iter().enumerate() {
673 out[i] = vals[k];
674 }
675 out
676}
677
678#[inline]
684fn reflect_index(mut p: isize, len: usize) -> usize {
685 let l = len as isize;
686 loop {
687 if p < 0 {
688 p = -p - 1;
689 } else if p >= l {
690 p = 2 * l - p - 1;
691 } else {
692 return p as usize;
693 }
694 }
695}
696
697fn smooth_stack(stack: &[f64], dims: (usize, usize, usize), ne: usize, sigma: f64) -> Vec<f64> {
699 let n = dims.0 * dims.1 * dims.2;
700 let mut out = vec![0.0_f64; n * ne];
701 for e in 0..ne {
702 let sm = gaussian_filter_3d(&stack[e * n..e * n + n], dims, sigma);
703 out[e * n..e * n + n].copy_from_slice(&sm);
704 }
705 out
706}
707
708fn gaussian_filter_3d(data: &[f64], dims: (usize, usize, usize), sigma: f64) -> Vec<f64> {
711 if sigma <= 0.0 {
712 return data.to_vec();
713 }
714 let (nx, ny, nz) = dims;
715 let radius = (4.0 * sigma + 0.5) as usize;
716 let mut kernel = vec![0.0_f64; 2 * radius + 1];
717 let mut ksum = 0.0;
718 for (k, w) in kernel.iter_mut().enumerate() {
719 let x = k as f64 - radius as f64;
720 *w = (-0.5 * (x / sigma) * (x / sigma)).exp();
721 ksum += *w;
722 }
723 for w in kernel.iter_mut() {
724 *w /= ksum;
725 }
726
727 let idx = |x: usize, y: usize, z: usize| (z * ny + y) * nx + x;
728 let mut a = data.to_vec();
729 let mut b = vec![0.0_f64; a.len()];
730
731 for z in 0..nz {
733 for y in 0..ny {
734 for x in 0..nx {
735 let mut s = 0.0;
736 for (kk, &w) in kernel.iter().enumerate() {
737 let xi = reflect_index(x as isize + kk as isize - radius as isize, nx);
738 s += w * a[idx(xi, y, z)];
739 }
740 b[idx(x, y, z)] = s;
741 }
742 }
743 }
744 std::mem::swap(&mut a, &mut b);
745 for z in 0..nz {
747 for y in 0..ny {
748 for x in 0..nx {
749 let mut s = 0.0;
750 for (kk, &w) in kernel.iter().enumerate() {
751 let yi = reflect_index(y as isize + kk as isize - radius as isize, ny);
752 s += w * a[idx(x, yi, z)];
753 }
754 b[idx(x, y, z)] = s;
755 }
756 }
757 }
758 std::mem::swap(&mut a, &mut b);
759 for z in 0..nz {
761 for y in 0..ny {
762 for x in 0..nx {
763 let mut s = 0.0;
764 for (kk, &w) in kernel.iter().enumerate() {
765 let zi = reflect_index(z as isize + kk as isize - radius as isize, nz);
766 s += w * a[idx(x, y, zi)];
767 }
768 b[idx(x, y, z)] = s;
769 }
770 }
771 }
772 b
773}
774
775fn median_filter_3(data: &[f64], dims: (usize, usize, usize)) -> Vec<f64> {
777 let (nx, ny, nz) = dims;
778 let idx = |x: usize, y: usize, z: usize| (z * ny + y) * nx + x;
779 let mut out = vec![0.0_f64; data.len()];
780 for z in 0..nz {
781 for y in 0..ny {
782 for x in 0..nx {
783 let mut win = [0.0_f64; 27];
784 let mut c = 0;
785 for dz in -1..=1_isize {
786 let zi = reflect_index(z as isize + dz, nz);
787 for dy in -1..=1_isize {
788 let yi = reflect_index(y as isize + dy, ny);
789 for dx in -1..=1_isize {
790 let xi = reflect_index(x as isize + dx, nx);
791 win[c] = data[idx(xi, yi, zi)];
792 c += 1;
793 }
794 }
795 }
796 win.sort_by(|a, b| a.partial_cmp(b).unwrap());
797 out[idx(x, y, z)] = win[13];
798 }
799 }
800 }
801 out
802}
803
804fn percentile(vals: &[f64], q: f64) -> f64 {
810 if vals.is_empty() {
811 return 0.0;
812 }
813 let mut v = vals.to_vec();
814 v.sort_by(|a, b| a.partial_cmp(b).unwrap());
815 let rank = q / 100.0 * (v.len() as f64 - 1.0);
816 let lo = rank.floor() as usize;
817 let hi = rank.ceil() as usize;
818 if lo == hi {
819 v[lo]
820 } else {
821 let frac = rank - lo as f64;
822 v[lo] * (1.0 - frac) + v[hi] * frac
823 }
824}
825
826fn median(vals: &[f64]) -> f64 {
827 if vals.is_empty() {
828 return 0.0;
829 }
830 let mut v = vals.to_vec();
831 v.sort_by(|a, b| a.partial_cmp(b).unwrap());
832 let m = v.len() / 2;
833 if v.len() % 2 == 1 {
834 v[m]
835 } else {
836 0.5 * (v[m - 1] + v[m])
837 }
838}
839
840#[cfg(test)]
841mod tests {
842 use super::*;
843
844 #[test]
845 fn reflect_index_matches_scipy() {
846 let len = 4;
849 let got: Vec<usize> = (-2..6).map(|p| reflect_index(p, len)).collect();
850 assert_eq!(got, vec![1, 0, 0, 1, 2, 3, 3, 2]);
851 }
852
853 #[test]
854 fn gaussian_preserves_constant() {
855 let dims = (8, 8, 8);
856 let n = 8 * 8 * 8;
857 let data = vec![3.0_f64; n];
858 let out = gaussian_filter_3d(&data, dims, 1.0);
859 for &v in &out {
860 assert!((v - 3.0).abs() < 1e-9, "constant not preserved: {v}");
861 }
862 }
863
864 #[test]
865 fn median_removes_impulse() {
866 let dims = (5, 5, 5);
867 let n = 125;
868 let mut data = vec![1.0_f64; n];
869 let idx = |x: usize, y: usize, z: usize| (z * 5 + y) * 5 + x;
870 data[idx(2, 2, 2)] = 100.0; let out = median_filter_3(&data, dims);
872 assert!(
873 (out[idx(2, 2, 2)] - 1.0).abs() < 1e-9,
874 "impulse survived median"
875 );
876 }
877
878 #[test]
879 fn percentile_and_median_basic() {
880 let v = [1.0, 2.0, 3.0, 4.0];
881 assert!((percentile(&v, 0.0) - 1.0).abs() < 1e-12);
882 assert!((percentile(&v, 100.0) - 4.0).abs() < 1e-12);
883 assert!((median(&v) - 2.5).abs() < 1e-12);
884 let v2 = [1.0, 2.0, 3.0];
885 assert!((median(&v2) - 2.0).abs() < 1e-12);
886 }
887
888 #[test]
891 fn fitter_recovers_planted_theta_mwf() {
892 let tes = [0.004, 0.012, 0.020, 0.028];
893 let b0 = 7.0;
894 let fitter = AnchoredFitter::new(&tes, b0, &[]);
895
896 let theta_deg = 30.0_f64;
897 let mwf = 0.10_f64;
898 let theta = theta_deg.to_radians();
899 let meso = 8.0;
901 let h = hc_wm_r2prime(theta, mwf, &tes, b0);
902 let rho = h + meso;
903 let first = hc_wm_signal_mag(tes[0], theta, b0, mwf, meso);
905 let sig: Vec<f64> = tes
906 .iter()
907 .map(|&t| hc_wm_signal_mag(t, theta, b0, mwf, meso) / first)
908 .collect();
909
910 let (t, m, _s) = fitter.fit(&sig, &[rho], None, None, 0.25);
911 assert!((m[0] - mwf).abs() <= 0.02, "mwf {} vs {}", m[0], mwf);
915 assert!(
916 (t[0] - theta_deg).abs() <= 12.0,
917 "theta {} vs {}",
918 t[0],
919 theta_deg
920 );
921 }
922
923 fn synth_phantom(
929 dims: (usize, usize, usize),
930 wm_x_max: usize,
931 b0: f64,
932 tes: &[f64],
933 se_tes: &[f64],
934 ) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<u8>) {
935 let (nx, ny, nz) = dims;
936 let n = nx * ny * nz;
937 let ne = tes.len();
938 let se_e = se_tes.len();
939 let dr = 137.0 * b0 / 3.0;
940 let theta = 40.0_f64.to_radians();
941 let mwf = 0.12;
942 let meso = 6.0;
943 let h = hc_wm_r2prime(theta, mwf, tes, b0);
944 let r2star = 25.0;
945
946 let mut chi = vec![0.0_f64; n];
947 let mut r2p = vec![0.0_f64; n];
948 let mut mag = vec![0.0_f64; n * ne];
949 let mut se = vec![0.0_f64; n * se_e];
950 let mask = vec![1u8; n];
951 for k in 0..nz {
952 for j in 0..ny {
953 for i in 0..nx {
954 let idx = (k * ny + j) * nx + i;
955 if i < wm_x_max {
956 chi[idx] = -0.03;
957 r2p[idx] = h + meso;
958 for (e, &t) in tes.iter().enumerate() {
959 mag[idx * ne + e] = hc_wm_signal_mag(t, theta, b0, mwf, meso);
960 }
961 for (e, &t) in se_tes.iter().enumerate() {
962 se[idx * se_e + e] = hc_wm_se_signal(t, mwf);
963 }
964 } else {
965 chi[idx] = 0.05;
966 r2p[idx] = dr * 0.05;
967 for (e, &t) in tes.iter().enumerate() {
968 mag[idx * ne + e] = (-r2star * t).exp();
969 }
970 for (e, &t) in se_tes.iter().enumerate() {
971 se[idx * se_e + e] = (-r2star * t).exp();
972 }
973 }
974 }
975 }
976 }
977 (chi, r2p, mag, se, mask)
978 }
979
980 #[test]
984 fn hc_chisep_end_to_end_wm_and_gm() {
985 let tes = [0.004, 0.012, 0.020, 0.028];
986 let se_tes = [0.01, 0.03, 0.05, 0.07];
987 let dims = (20, 20, 20);
988 let n = dims.0 * dims.1 * dims.2;
989 let (chi, r2p, mag, se, mask) = synth_phantom(dims, 6, 7.0, &tes, &se_tes);
990 let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
991 let params = HcChisepParams { b0: 7.0, se_echo_times: se_tes.to_vec(), ..Default::default() };
992
993 let mut stages = 0;
994 let (pos, neg, tot) = hc_chisep(
995 &chi, &r2p, &mag, &tes, Some(&se), &mask, &grid, ¶ms, |i, _| stages = i,
996 );
997 assert_eq!(stages, 5, "all pipeline stages ran (beat detected)");
998
999 for i in 0..n {
1001 assert!(pos[i] >= 0.0 && neg[i] <= 0.0, "signs at {i}");
1002 assert!((tot[i] - (pos[i] + neg[i])).abs() < 1e-9, "invariant at {i}");
1003 }
1004 assert!(pos.iter().any(|&v| v > 1e-3), "some χ+ present");
1006 assert!(neg.iter().any(|&v| v < -1e-3), "some χ− present (WM)");
1007 }
1008
1009 #[test]
1012 fn hc_chisep_no_beat_falls_back_to_closed_form() {
1013 let tes = [0.004, 0.012, 0.020, 0.028];
1014 let dims = (10, 10, 10);
1015 let n = dims.0 * dims.1 * dims.2;
1016 let (chi, r2p, mag, _se, mask) = synth_phantom(dims, 0, 7.0, &tes, &[]);
1017 let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
1018 let params = HcChisepParams { b0: 7.0, ..Default::default() };
1019
1020 let (pos, neg, _tot) =
1021 hc_chisep(&chi, &r2p, &mag, &tes, None, &mask, &grid, ¶ms, |_, _| {});
1022 for i in 0..n {
1024 assert!((pos[i] - 0.05).abs() < 1e-6, "closed-form χ+ at {i}: {}", pos[i]);
1025 assert!(neg[i].abs() < 1e-9, "closed-form χ− at {i}: {}", neg[i]);
1026 }
1027 }
1028
1029 #[test]
1031 fn hc_chisep_empty_mask_returns_zero() {
1032 let tes = [0.004, 0.012, 0.020, 0.028];
1033 let dims = (6, 6, 6);
1034 let n = dims.0 * dims.1 * dims.2;
1035 let (chi, r2p, mag, _se, _m) = synth_phantom(dims, 0, 7.0, &tes, &[]);
1036 let mask = vec![0u8; n];
1037 let grid = Grid::new(dims.0, dims.1, dims.2, 1.0, 1.0, 1.0);
1038 let params = HcChisepParams { b0: 7.0, ..Default::default() };
1039
1040 let (pos, neg, tot) =
1041 hc_chisep(&chi, &r2p, &mag, &tes, None, &mask, &grid, ¶ms, |_, _| {});
1042 assert!(pos.iter().all(|&v| v == 0.0), "χ+ all zero");
1043 assert!(neg.iter().all(|&v| v == 0.0), "χ− all zero");
1044 assert!(tot.iter().all(|&v| v == 0.0), "χ_total all zero");
1045 }
1046}