1use nereids_endf::resonance::ResonanceData;
16use nereids_fitting::joint_resolution::{Densities, JointResolutionModel, SpectrumSpec};
17use nereids_fitting::lm::{self, LmConfig};
18use nereids_fitting::parameters::{FitParameter, ParameterSet};
19use nereids_fitting::resolution_calib::{GAUSSIAN_DELTA_L_BOUNDS_M, GAUSSIAN_DELTA_T_BOUNDS_US};
20
21use crate::error::PipelineError;
22use crate::pipeline::TEMPERATURE_BOUNDS_K;
23
24pub struct SampleSpectrum {
26 pub energies: Vec<f64>,
28 pub transmission: Vec<f64>,
30 pub uncertainty: Vec<f64>,
32 pub isotopes: Vec<(ResonanceData, f64)>,
34 pub temperature_k: f64,
36 pub fit_temperature: bool,
38}
39
40pub struct CalibrantSpectrum {
46 pub energies: Vec<f64>,
48 pub transmission: Vec<f64>,
50 pub uncertainty: Vec<f64>,
52 pub isotopes: Vec<(ResonanceData, f64)>,
54 pub temperature_k: f64,
56}
57
58#[derive(Debug, Clone)]
63pub struct JointFitResult {
64 pub densities: Vec<f64>,
66 pub density_uncertainties: Option<Vec<f64>>,
69 pub temperature_k: Option<f64>,
71 pub temperature_k_unc: Option<f64>,
77 pub delta_t_us: f64,
79 pub delta_l_m: f64,
81 pub reduced_chi_squared: f64,
83 pub converged: bool,
85 pub iterations: usize,
87}
88
89const fn resolution_indices(n_density: usize, fit_temperature: bool) -> (usize, usize) {
92 let after = n_density + if fit_temperature { 1 } else { 0 };
93 (after, after + 1)
94}
95
96pub fn fit_with_calibrant(
111 sample: &SampleSpectrum,
112 calibrant: &CalibrantSpectrum,
113 flight_path_m: f64,
114 delta_t_init: f64,
115 delta_l_init: f64,
116) -> Result<JointFitResult, PipelineError> {
117 for (label, energies, transmission, uncertainty, isotopes) in [
118 (
119 "sample",
120 &sample.energies,
121 &sample.transmission,
122 &sample.uncertainty,
123 &sample.isotopes,
124 ),
125 (
126 "calibrant",
127 &calibrant.energies,
128 &calibrant.transmission,
129 &calibrant.uncertainty,
130 &calibrant.isotopes,
131 ),
132 ] {
133 check_spectrum(label, energies, transmission, uncertainty)?;
134 if isotopes.is_empty() {
135 return Err(PipelineError::InvalidParameter(format!(
136 "the {label} needs at least one isotope"
137 )));
138 }
139 }
140 let positive = (f64::MIN_POSITIVE, f64::INFINITY);
148 for (label, value, range) in [
149 ("flight_path_m", flight_path_m, positive),
150 ("temperature_k", sample.temperature_k, TEMPERATURE_BOUNDS_K),
151 (
152 "the calibrant temperature",
153 calibrant.temperature_k,
154 TEMPERATURE_BOUNDS_K,
155 ),
156 ("delta_t_init", delta_t_init, GAUSSIAN_DELTA_T_BOUNDS_US),
157 ("delta_l_init", delta_l_init, GAUSSIAN_DELTA_L_BOUNDS_M),
158 ]
159 .into_iter()
160 .chain(
161 sample
162 .isotopes
163 .iter()
164 .map(|(_, n)| ("a sample density", *n, positive)),
165 )
166 .chain(
167 calibrant
168 .isotopes
169 .iter()
170 .map(|(_, n)| ("a calibrant density", *n, positive)),
171 ) {
172 check_range(label, value, range)?;
173 }
174
175 let n_isotopes = sample.isotopes.len();
176 let (delta_t_index, delta_l_index) = resolution_indices(n_isotopes, sample.fit_temperature);
177 let temperature_index = sample.fit_temperature.then_some(n_isotopes);
178 let (sample_data, initial_densities): (Vec<_>, Vec<_>) =
179 sample.isotopes.iter().cloned().unzip();
180 let (calibrant_data, calibrant_densities): (Vec<_>, Vec<_>) =
181 calibrant.isotopes.iter().cloned().unzip();
182
183 let model = JointResolutionModel::new(
184 SpectrumSpec {
185 energies: sample.energies.clone(),
186 resonance_data: sample_data,
187 densities: Densities::Fitted((0..n_isotopes).collect()),
188 temperature_index,
189 temperature_k: sample.temperature_k,
190 },
191 SpectrumSpec {
192 energies: calibrant.energies.clone(),
193 resonance_data: calibrant_data,
194 densities: Densities::Known(calibrant_densities),
195 temperature_index: None,
196 temperature_k: calibrant.temperature_k,
197 },
198 flight_path_m,
199 delta_t_index,
200 delta_l_index,
201 )
202 .map_err(PipelineError::Fitting)?;
203
204 let mut values: Vec<FitParameter> = initial_densities
205 .iter()
206 .enumerate()
207 .map(|(i, &n)| FitParameter::non_negative(format!("density_{i}"), n))
208 .collect();
209 if sample.fit_temperature {
210 values.push(FitParameter {
211 name: "temperature_k".into(),
212 value: sample.temperature_k,
213 lower: TEMPERATURE_BOUNDS_K.0,
214 upper: TEMPERATURE_BOUNDS_K.1,
215 fixed: false,
216 });
217 }
218 for (name, seed, (lo, hi)) in [
224 ("delta_t_us_sq", delta_t_init, GAUSSIAN_DELTA_T_BOUNDS_US),
225 ("delta_l_m_sq", delta_l_init, GAUSSIAN_DELTA_L_BOUNDS_M),
226 ] {
227 values.push(FitParameter {
228 name: name.into(),
229 value: seed * seed,
230 lower: lo * lo,
231 upper: hi * hi,
232 fixed: false,
233 });
234 }
235 let mut params = ParameterSet::new(values);
236
237 let mut data = sample.transmission.clone();
239 data.extend_from_slice(&calibrant.transmission);
240 let mut sigma = sample.uncertainty.clone();
241 sigma.extend_from_slice(&calibrant.uncertainty);
242
243 let result = lm::levenberg_marquardt(
244 &model,
245 &data,
246 &sigma,
247 &mut params,
248 &LmConfig {
249 compute_covariance: true,
250 ..Default::default()
251 },
252 )
253 .map_err(PipelineError::Fitting)?;
254
255 let sigma_of = |i: usize| {
258 result
259 .uncertainties
260 .as_ref()
261 .and_then(|u| u.get(i).copied())
262 };
263 Ok(JointFitResult {
264 densities: result.params[..n_isotopes].to_vec(),
265 density_uncertainties: result
266 .uncertainties
267 .as_ref()
268 .map(|u| u[..n_isotopes].to_vec()),
269 temperature_k: temperature_index.map(|i| result.params[i]),
270 temperature_k_unc: temperature_index.and_then(sigma_of),
271 delta_t_us: result.params[delta_t_index].max(0.0).sqrt(),
272 delta_l_m: result.params[delta_l_index].max(0.0).sqrt(),
273 reduced_chi_squared: result.reduced_chi_squared,
274 converged: result.converged,
275 iterations: result.iterations,
276 })
277}
278
279fn check_range(label: &str, value: f64, (lo, hi): (f64, f64)) -> Result<(), PipelineError> {
281 if value.is_finite() && (lo..=hi).contains(&value) {
282 return Ok(());
283 }
284 Err(PipelineError::InvalidParameter(format!(
285 "{label} must be finite and lie in [{lo}, {hi}], got {value}"
286 )))
287}
288
289fn check_spectrum(
290 label: &str,
291 energies: &[f64],
292 transmission: &[f64],
293 uncertainty: &[f64],
294) -> Result<(), PipelineError> {
295 if energies.is_empty() {
296 return Err(PipelineError::InvalidParameter(format!(
297 "the {label} energy grid is empty"
298 )));
299 }
300 if transmission.len() != energies.len() || uncertainty.len() != energies.len() {
301 return Err(PipelineError::ShapeMismatch(format!(
302 "{label}: {} energies, {} transmission points, {} uncertainties",
303 energies.len(),
304 transmission.len(),
305 uncertainty.len(),
306 )));
307 }
308 if uncertainty.iter().any(|s| !s.is_finite() || *s <= 0.0) {
309 return Err(PipelineError::InvalidParameter(format!(
310 "the {label} uncertainties must all be finite and positive"
311 )));
312 }
313 if energies.iter().any(|e| !e.is_finite() || *e <= 0.0) {
319 return Err(PipelineError::InvalidParameter(format!(
320 "the {label} energies must all be finite and positive"
321 )));
322 }
323 if let Some(i) = energies.windows(2).position(|w| w[1] <= w[0]) {
324 return Err(PipelineError::InvalidParameter(format!(
325 "the {label} energies must increase: [{i}] = {} is not below [{}] = {}",
326 energies[i],
327 i + 1,
328 energies[i + 1],
329 )));
330 }
331 if transmission.iter().any(|t| !t.is_finite()) {
332 return Err(PipelineError::InvalidParameter(format!(
333 "the {label} transmission values must all be finite"
334 )));
335 }
336 Ok(())
337}
338
339#[cfg(test)]
340mod tests {
341 use super::*;
342 use nereids_endf::resonance::test_support::synthetic_isotope;
343 use nereids_fitting::resolution_calib::{
344 GAUSSIAN_DELTA_L_BOUNDS_M, GAUSSIAN_DELTA_T_BOUNDS_US,
345 };
346 use nereids_physics::resolution::{ResolutionFunction, ResolutionParams};
347 use nereids_physics::transmission::{InstrumentParams, SampleParams, forward_model};
348
349 const L: f64 = 25.0;
350 const W: f64 = 0.30;
351 const DL: f64 = 0.05;
352
353 fn spectrum(
354 iso: &nereids_endf::resonance::ResonanceData,
355 density: f64,
356 temperature_k: f64,
357 energies: &[f64],
358 ) -> Vec<f64> {
359 let sample = SampleParams::new(temperature_k, vec![(iso.clone(), density)]).unwrap();
360 let inst = InstrumentParams {
361 resolution: ResolutionFunction::Gaussian(ResolutionParams::new(L, W, DL, 0.0).unwrap()),
362 };
363 forward_model(energies, &sample, Some(&inst)).unwrap()
364 }
365
366 fn sample_arm(
368 iso: &nereids_endf::resonance::ResonanceData,
369 energies: &[f64],
370 t: &[f64],
371 unc: &[f64],
372 density: f64,
373 temperature_k: f64,
374 ) -> SampleSpectrum {
375 SampleSpectrum {
376 energies: energies.to_vec(),
377 transmission: t.to_vec(),
378 uncertainty: unc.to_vec(),
379 isotopes: vec![(iso.clone(), density)],
380 temperature_k,
381 fit_temperature: true,
382 }
383 }
384
385 fn calibrant_arm(
387 iso: &nereids_endf::resonance::ResonanceData,
388 energies: &[f64],
389 t: &[f64],
390 unc: &[f64],
391 density: f64,
392 temperature_k: f64,
393 ) -> CalibrantSpectrum {
394 CalibrantSpectrum {
395 energies: energies.to_vec(),
396 transmission: t.to_vec(),
397 uncertainty: unc.to_vec(),
398 isotopes: vec![(iso.clone(), density)],
399 temperature_k,
400 }
401 }
402
403 #[test]
414 fn a_joint_fit_recovers_the_sample_and_the_shared_resolution() {
415 let sample_iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
416 let calibrant_iso = synthetic_isotope(72, 177, 31.0, 0.04, 0.07);
417 let sample_e: Vec<f64> = (0..160).map(|i| 18.0 + i as f64 * 0.03).collect();
418 let calibrant_e: Vec<f64> = (0..140).map(|i| 29.0 + i as f64 * 0.032).collect();
419 let sample_t = spectrum(&sample_iso, 2.0e-3, 320.0, &sample_e);
420 let calibrant_t = spectrum(&calibrant_iso, 3.0e-3, 300.0, &calibrant_e);
421 let unc = vec![1.0e-3; sample_e.len()];
422
423 let calibrant = CalibrantSpectrum {
424 energies: calibrant_e,
425 transmission: calibrant_t,
426 uncertainty: vec![1.0e-3; 140],
427 isotopes: vec![(calibrant_iso, 3.0e-3)],
428 temperature_k: 300.0,
429 };
430
431 let r = fit_with_calibrant(
432 &sample_arm(&sample_iso, &sample_e, &sample_t, &unc, 1.6e-3, 300.0),
433 &calibrant,
434 L,
435 0.45,
437 0.02,
438 )
439 .expect("the joint fit runs");
440
441 assert!(r.converged, "joint fit did not converge");
442 assert!(
443 (r.densities[0] - 2.0e-3).abs() < 0.05 * 2.0e-3,
444 "density {} is not within 5 % of 2.0e-3",
445 r.densities[0]
446 );
447 let temperature = r.temperature_k.expect("temperature was fitted");
448 assert!(
449 (temperature - 320.0).abs() < 10.0,
450 "temperature {temperature} K is not within 10 K of 320 K"
451 );
452 assert!(
453 (r.delta_t_us - W).abs() < 0.2 * W,
454 "the shared width {} is not within 20 % of {W}, so the fit did not \
455 move off its seed of 0.45",
456 r.delta_t_us
457 );
458 assert!(
459 r.temperature_k_unc
460 .is_some_and(|s| s.is_finite() && s > 0.0),
461 "a joint fit must report a temperature uncertainty"
462 );
463 }
464
465 #[test]
475 fn a_width_the_grid_cannot_carry_is_rejected_and_unreachable() {
476 let iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
477 let energies: Vec<f64> = (0..40).map(|i| 18.0 + i as f64 * 0.05).collect();
478 let t = spectrum(&iso, 2.0e-3, 300.0, &energies);
479 let unc = vec![1.0e-3; energies.len()];
480 let calibrant = calibrant_arm(&iso, &energies, &t, &unc, 2.0e-3, 300.0);
481 let run = |dt: f64, dl: f64| {
482 fit_with_calibrant(
483 &sample_arm(&iso, &energies, &t, &unc, 2.0e-3, 300.0),
484 &calibrant,
485 L,
486 dt,
487 dl,
488 )
489 };
490
491 let (t_lo, t_hi) = GAUSSIAN_DELTA_T_BOUNDS_US;
492 let (_, l_hi) = GAUSSIAN_DELTA_L_BOUNDS_M;
493 assert!(
494 run(t_hi * 2.0, DL).is_err(),
495 "a timing width above {t_hi} µs must be rejected before any \
496 residual is evaluated"
497 );
498 assert!(
499 run(t_lo / 2.0, DL).is_err(),
500 "a timing width below {t_lo} µs must be rejected"
501 );
502 assert!(
503 run(W, l_hi * 2.0).is_err(),
504 "a flight-path width above {l_hi} m must be rejected"
505 );
506 let r = run(W, DL).expect("a width inside the box is accepted");
509 assert!(
510 (t_lo..=t_hi).contains(&r.delta_t_us),
511 "the fitted timing width {} left [{t_lo}, {t_hi}]",
512 r.delta_t_us
513 );
514 assert!(
515 r.delta_l_m <= l_hi,
516 "the fitted flight-path width {} exceeded {l_hi}",
517 r.delta_l_m
518 );
519 }
520
521 #[test]
529 fn unreadable_grids_and_spectra_are_rejected_on_either_arm() {
530 let iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
531 let energies: Vec<f64> = (0..40).map(|i| 18.0 + i as f64 * 0.05).collect();
532 let good = spectrum(&iso, 2.0e-3, 300.0, &energies);
533 let unc = vec![1.0e-3; energies.len()];
534
535 let mut descending = energies.clone();
536 descending.reverse();
537 let mut repeated = energies.clone();
538 repeated[7] = repeated[6];
539 let mut negative = energies.clone();
540 negative[0] = -1.0;
541 let mut nan_t = good.clone();
542 nan_t[3] = f64::NAN;
543
544 let cases: [(&str, Vec<f64>, Vec<f64>); 4] = [
545 ("descending energies", descending, good.clone()),
546 ("a repeated energy", repeated, good.clone()),
547 ("a negative energy", negative, good.clone()),
548 ("a non-finite transmission", energies.clone(), nan_t),
549 ];
550 for (what, grid, values) in cases {
551 let sound = calibrant_arm(&iso, &energies, &good, &unc, 2.0e-3, 300.0);
554 assert!(
555 fit_with_calibrant(
556 &sample_arm(&iso, &grid, &values, &unc, 2.0e-3, 300.0),
557 &sound,
558 L,
559 W,
560 DL,
561 )
562 .is_err(),
563 "{what} must be rejected on the sample arm"
564 );
565 let broken = CalibrantSpectrum {
566 energies: grid,
567 transmission: values,
568 uncertainty: unc.clone(),
569 isotopes: vec![(iso.clone(), 2.0e-3)],
570 temperature_k: 300.0,
571 };
572 assert!(
573 fit_with_calibrant(
574 &sample_arm(&iso, &energies, &good, &unc, 2.0e-3, 300.0),
575 &broken,
576 L,
577 W,
578 DL,
579 )
580 .is_err(),
581 "{what} must be rejected on the calibrant arm"
582 );
583 }
584 }
585
586 #[test]
594 fn a_calibrant_that_constrains_nothing_is_rejected() {
595 let iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
596 let energies: Vec<f64> = (0..40).map(|i| 18.0 + i as f64 * 0.05).collect();
597 let t = spectrum(&iso, 2.0e-3, 300.0, &energies);
598 let unc = vec![1.0e-3; energies.len()];
599
600 for bad in [0.0, -1.0e-3, f64::NAN] {
601 let calibrant = CalibrantSpectrum {
602 energies: energies.clone(),
603 transmission: t.clone(),
604 uncertainty: unc.clone(),
605 isotopes: vec![(iso.clone(), bad)],
606 temperature_k: 300.0,
607 };
608 let Err(err) = fit_with_calibrant(
609 &sample_arm(&iso, &energies, &t, &unc, 2.0e-3, 300.0),
610 &calibrant,
611 L,
612 W,
613 DL,
614 ) else {
615 panic!("a calibrant density of {bad} must be rejected");
616 };
617 assert!(
618 matches!(err, PipelineError::InvalidParameter(ref m) if m.contains("calibrant")),
619 "expected a calibrant density rejection for {bad}, got {err}"
620 );
621 }
622 }
623
624 #[test]
626 fn a_temperature_outside_the_shared_box_is_rejected() {
627 let iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
628 let energies: Vec<f64> = (0..40).map(|i| 18.0 + i as f64 * 0.05).collect();
629 let t = spectrum(&iso, 2.0e-3, 300.0, &energies);
630 let unc = vec![1.0e-3; energies.len()];
631 let calibrant = calibrant_arm(&iso, &energies, &t, &unc, 2.0e-3, 300.0);
632 for seed in [0.0, 6000.0] {
635 assert!(
636 fit_with_calibrant(
637 &sample_arm(&iso, &energies, &t, &unc, 2.0e-3, seed),
638 &calibrant,
639 L,
640 W,
641 DL,
642 )
643 .is_err(),
644 "a sample temperature of {seed} K must be rejected"
645 );
646 }
647 }
648
649 #[test]
651 fn mismatched_spectra_are_rejected_before_fitting() {
652 let iso = synthetic_isotope(72, 178, 20.0, 0.05, 0.06);
653 let energies: Vec<f64> = (0..40).map(|i| 18.0 + i as f64 * 0.05).collect();
654 let t = spectrum(&iso, 2.0e-3, 300.0, &energies);
655 let unc = vec![1.0e-3; energies.len()];
656 let calibrant = CalibrantSpectrum {
657 energies: energies.clone(),
658 transmission: t.clone(),
659 uncertainty: unc[1..].to_vec(),
661 isotopes: vec![(iso.clone(), 2.0e-3)],
662 temperature_k: 300.0,
663 };
664 let err = fit_with_calibrant(
665 &sample_arm(&iso, &energies, &t, &unc, 2.0e-3, 300.0),
666 &calibrant,
667 L,
668 W,
669 DL,
670 )
671 .expect_err("a calibrant whose arrays disagree must be rejected");
672 assert!(
673 matches!(err, PipelineError::ShapeMismatch(ref m) if m.contains("calibrant")),
674 "expected a calibrant shape mismatch, got {err}"
675 );
676 }
677}