1use std::cell::RefCell;
18use std::collections::HashMap;
19use std::ops::ControlFlow;
20use std::str::FromStr;
21
22use levenberg_marquardt::LevenbergMarquardt;
23use nalgebra::DVector;
24use spectrafit_graph::{evaluate_compiled, CompiledGraph};
25use spectrafit_types::{
26 CoreError, FitGraphSpec, FitOptionsSpec, FitResultSpec, MeasurementSpec, TerminationReason,
27};
28
29use crate::error::SolverError;
30use crate::global::{solve_global, DeConfig};
31use crate::irls::{solve_irls, WeightFn, DEFAULT_MAX_OUTER_ITER, DEFAULT_TOL_WEIGHTS};
32use crate::lm_problem::LmProblem;
33use crate::postfit::{self, LmSolveOutcome};
34
35pub(crate) fn point_major_x(ds: &MeasurementSpec) -> Vec<f64> {
41 let n_dims = ds.x.len();
42 if n_dims == 0 {
43 return Vec::new();
44 }
45 let n_points = ds.x[0].len();
46 let mut out = Vec::with_capacity(n_points * n_dims);
47 for i in 0..n_points {
48 for dim in &ds.x {
49 out.push(dim[i]);
50 }
51 }
52 out
53}
54
55#[derive(Debug, Clone, Copy, PartialEq)]
74enum Solver {
75 Lm,
77 LmLegacy,
79 Trf,
81 Geodesic,
83 Dogleg,
85 NewtonCg,
87 Irls(WeightFn),
89 Global,
91 Varpro,
93 Auto,
95}
96
97impl Solver {
98 fn parse(s: &str) -> Result<Self, SolverError> {
105 match s {
106 "lm" => Ok(Solver::Lm),
107 "lm-legacy" => Ok(Solver::LmLegacy),
108 "trf" => Ok(Solver::Trf),
109 "geodesic" | "lm-geodesic" => Ok(Solver::Geodesic),
110 "dogleg" => Ok(Solver::Dogleg),
111 "newton-cg" | "newton_cg" | "newtoncg" | "steihaug" => Ok(Solver::NewtonCg),
112 "global" => Ok(Solver::Global),
113 "varpro" => Ok(Solver::Varpro),
114 "auto" => Ok(Solver::Auto),
115 _ if s.starts_with("irls") => {
116 let name = s.split_once(':').map(|(_, name)| name).unwrap_or("huber");
117 Ok(Solver::Irls(WeightFn::from_str(name)?))
118 }
119 other => Err(SolverError::UnrecognisedSolver(other.to_string())),
120 }
121 }
122}
123
124fn graph_has_tied_params(graph: &FitGraphSpec) -> bool {
130 !graph.expr_edges.is_empty()
131 || graph
132 .nodes
133 .iter()
134 .any(|n| n.parameters.values().any(|p| p.expr.is_some()))
135}
136
137fn graph_prefers_varpro(graph: &FitGraphSpec) -> bool {
141 !graph_has_tied_params(graph)
142 && graph.nodes.iter().all(|n| n.dataset_index.is_none())
146 && spectrafit_varpro::is_separable(graph)
147 && graph.nodes.iter().all(|n| {
148 n.parameters
149 .iter()
150 .filter(|(name, _)| name.as_str() != "amplitude")
155 .filter(|(_, p)| p.vary)
156 .all(|(_, p)| p.min.is_infinite() && p.max.is_infinite())
157 })
158}
159
160fn solve_varpro_path(
163 graph: &FitGraphSpec,
164 datasets: &[MeasurementSpec],
165 options: &FitOptionsSpec,
166) -> Result<FitResultSpec, CoreError> {
167 let mut param_specs: HashMap<String, spectrafit_types::ParameterSpec> = HashMap::new();
168 for node in &graph.nodes {
169 for (pname, pspec) in &node.parameters {
170 param_specs.insert(format!("{}.{}", node.id, pname), pspec.clone());
171 }
172 }
173 spectrafit_varpro::solve_varpro(graph, datasets, ¶m_specs, options)
174}
175
176pub fn fit(
211 graph: &FitGraphSpec,
212 datasets: Vec<MeasurementSpec>,
213 options: &FitOptionsSpec,
214) -> Result<FitResultSpec, CoreError> {
215 faer::set_global_parallelism(faer::Par::Seq);
233
234 let solver = Solver::parse(&options.solver)?;
240 let datasets = match dispatch_solver(solver, graph, datasets, options) {
241 ControlFlow::Break(result) => return result,
242 ControlFlow::Continue(datasets) => datasets,
243 };
244
245 let LmPreSolve {
248 cg,
249 free_keys,
250 all_params,
251 bounds,
252 scales,
253 sigma,
254 x_all,
255 y_all,
256 init_fit,
257 init_params,
258 node_param_bufs,
259 free_to_node_param,
260 tied_to_node_param,
261 } = prepare_lm_pre_solve(graph, &datasets)?;
262
263 let problem = LmProblem {
265 compiled: &cg,
266 free_keys: free_keys.clone(),
267 bounds: bounds.clone(),
268 all_params: all_params.clone(),
269 node_param_bufs,
270 free_to_node_param,
271 tied_to_node_param,
272 x_concat: x_all.clone(),
273 y_concat: y_all.clone(),
274 params: init_params,
275 scales: scales.clone(),
276 sigma,
277 residual_buf: RefCell::new(vec![0.0; y_all.len()]),
278 jacobian_buf: RefCell::new(vec![0.0; y_all.len() * free_keys.len()]),
281 residual_count: RefCell::new(0),
282 jacobian_count: RefCell::new(0),
283 residual_time_ns: RefCell::new(0),
284 jacobian_time_ns: RefCell::new(0),
285 };
286
287 let outcome = run_lm_solve(problem, solver, options, free_keys.len(), y_all.len())?;
289
290 postfit::assemble_result(
292 outcome,
293 postfit::PostfitInputs {
294 cg: &cg,
295 graph,
296 datasets: &datasets,
297 x_all: &x_all,
298 y_all: &y_all,
299 },
300 init_fit,
301 )
302}
303
304fn dispatch_solver(
320 solver: Solver,
321 graph: &FitGraphSpec,
322 datasets: Vec<MeasurementSpec>,
323 options: &FitOptionsSpec,
324) -> ControlFlow<Result<FitResultSpec, CoreError>, Vec<MeasurementSpec>> {
325 match solver {
326 Solver::Irls(weight) => ControlFlow::Break(solve_irls(
327 graph,
328 datasets,
329 options,
330 weight,
331 DEFAULT_MAX_OUTER_ITER,
332 DEFAULT_TOL_WEIGHTS,
336 )),
337
338 Solver::Global => {
339 let config = DeConfig {
344 max_gen: if options.max_iterations > 0 {
345 options.max_iterations as usize
346 } else {
347 DeConfig::default().max_gen
348 },
349 ..DeConfig::default()
350 };
351 ControlFlow::Break(solve_global(graph, datasets, options, config))
352 }
353 Solver::Varpro => {
354 if !spectrafit_varpro::is_separable(graph) {
359 return ControlFlow::Break(Err(SolverError::VarproNotSeparable.into()));
360 }
361 if graph_has_tied_params(graph) {
362 return ControlFlow::Break(Err(SolverError::VarproExprEdgesUnsupported.into()));
363 }
364 if graph.nodes.iter().any(|n| n.dataset_index.is_some()) {
369 return ControlFlow::Break(
370 Err(SolverError::VarproDatasetScopingUnsupported.into()),
371 );
372 }
373 ControlFlow::Break(run_varpro_guarded(graph, &datasets, options))
377 }
378 Solver::Auto if graph_prefers_varpro(graph) => {
379 ControlFlow::Break(run_varpro_guarded(graph, &datasets, options))
387 }
388 Solver::Lm
391 | Solver::LmLegacy
392 | Solver::Trf
393 | Solver::Geodesic
394 | Solver::Dogleg
395 | Solver::NewtonCg
396 | Solver::Auto => ControlFlow::Continue(datasets),
397 }
398}
399
400struct LmPreSolve {
407 cg: CompiledGraph,
410 free_keys: Vec<String>,
412 all_params: HashMap<String, f64>,
414 bounds: Vec<(f64, f64)>,
416 scales: Vec<f64>,
418 sigma: Vec<f64>,
420 x_all: Vec<f64>,
422 y_all: Vec<f64>,
424 init_fit: Vec<f64>,
426 init_params: DVector<f64>,
428 node_param_bufs: Vec<Vec<f64>>,
430 free_to_node_param: Vec<(usize, usize)>,
432 tied_to_node_param: Vec<(usize, usize)>,
434}
435
436fn prepare_lm_pre_solve(
441 graph: &FitGraphSpec,
442 datasets: &[MeasurementSpec],
443) -> Result<LmPreSolve, CoreError> {
444 let mut cg = CompiledGraph::compile(graph)?;
446 cg.dataset_offsets = {
451 let mut offs = Vec::with_capacity(datasets.len() + 1);
452 let mut acc = 0usize;
453 offs.push(0);
454 for ds in datasets {
455 acc += ds.y.len();
456 offs.push(acc);
457 }
458 offs
459 };
460 let free_keys = cg.free_keys.clone();
461
462 let (all_params, init_vals, bounds, scales) = collect_free_param_specs(graph, &free_keys)?;
464
465 let init_vals_clone = init_vals.clone();
466 let init_params = DVector::from_vec(
471 init_vals
472 .iter()
473 .zip(scales.iter())
474 .map(|(&v, &s)| v / s)
475 .collect::<Vec<f64>>(),
476 );
477
478 let sigma: Vec<f64> = datasets
480 .iter()
481 .flat_map(|ds| {
482 let n = ds.y.len();
483 match &ds.sigma {
484 Some(s) => s.clone(),
485 None => vec![1.0_f64; n],
486 }
487 })
488 .collect();
489
490 let x_all: Vec<f64> = datasets.iter().flat_map(point_major_x).collect();
494
495 let y_all: Vec<f64> = datasets
496 .iter()
497 .flat_map(|ds| ds.y.iter().copied())
498 .collect();
499
500 let init_fit = evaluate_compiled(&cg, &all_params, &x_all)?;
502
503 let (node_param_bufs, free_to_node_param, tied_to_node_param) =
505 build_node_param_buffers(&cg, &all_params, &bounds, &init_vals_clone)?;
506
507 debug_assert_eq!(
528 scales.len(),
529 free_keys.len(),
530 "one scale factor per free parameter"
531 );
532
533 Ok(LmPreSolve {
534 cg,
535 free_keys,
536 all_params,
537 bounds,
538 scales,
539 sigma,
540 x_all,
541 y_all,
542 init_fit,
543 init_params,
544 node_param_bufs,
545 free_to_node_param,
546 tied_to_node_param,
547 })
548}
549
550fn run_lm_solve<'a>(
557 mut problem: LmProblem<'a>,
558 solver: Solver,
559 options: &FitOptionsSpec,
560 free_keys_len: usize,
561 y_all_len: usize,
562) -> Result<LmSolveOutcome<'a>, CoreError> {
563 if problem.has_tied() {
568 let p0: Vec<f64> = problem.params.iter().copied().collect();
569 problem.set_free_and_tied(&p0);
570 }
571
572 let patience = (options.max_iterations as usize).max(1);
574
575 let tol = if options.tolerance > 0.0 {
580 options.tolerance
581 } else {
582 1e-8
583 };
584 let max_nfev = patience.saturating_mul(free_keys_len + 1);
585 let _solve_t0 = std::time::Instant::now();
586 let (
587 result_problem,
588 n_iter_val,
589 success_val,
590 message_val,
591 cost_history,
592 gradient_norm_history,
593 params_history,
594 ) = match solver {
595 Solver::LmLegacy => {
596 let lm = LevenbergMarquardt::new().with_patience(patience);
600 let lm = if options.tolerance > 0.0 {
601 lm.with_tol(options.tolerance)
602 } else {
603 lm
604 };
605 let (rp, report) = lm.minimize(problem);
606 (
607 rp,
608 report.number_of_evaluations as u64,
609 report.termination.was_successful(),
610 map_termination(&report.termination).as_str().to_string(),
611 Vec::new(),
612 Vec::new(),
613 Vec::new(),
614 )
615 }
616 Solver::Dogleg | Solver::NewtonCg => {
617 let mut cfg = spectrafit_dogleg::TrustRegionConfig {
624 ftol: tol,
625 xtol: tol,
626 gtol: tol,
627 max_nfev,
628 ..Default::default()
629 };
630 if let Some(d0) = options.delta0 {
631 cfg.delta0 = d0;
632 }
633 if let Some(md) = options.max_delta {
634 cfg.max_delta = md;
635 }
636 if let Some(e) = options.eta {
637 if !(e.is_finite() && (0.0..0.25).contains(&e)) {
643 return Err(SolverError::InvalidEta(e).into());
644 }
645 cfg.eta = e;
646 }
647 let report = if solver == Solver::Dogleg {
648 spectrafit_dogleg::minimize(&mut problem, &cfg)
649 } else {
650 spectrafit_newton_cg::minimize(&mut problem, &cfg)
651 };
652 let msg = faer_termination_str(report.termination).to_string();
653 (
654 problem,
655 report.n_iter as u64,
656 report.termination.was_successful(),
657 msg,
658 report.cost_history,
659 report.gradient_norm_history,
660 report.params_history,
661 )
662 }
663 _ => {
667 let cfg = spectrafit_levenberg_marquardt::StrategyConfig {
668 kind: spectrafit_levenberg_marquardt::select_regime(y_all_len, free_keys_len),
669 ftol: tol,
670 xtol: tol,
671 gtol: tol,
672 max_nfev,
673 geodesic: solver == Solver::Geodesic,
676 bound_scaling: solver == Solver::Trf || solver == Solver::Auto,
685 ..Default::default()
686 };
687 let report = spectrafit_levenberg_marquardt::minimize(&mut problem, &cfg);
688 let msg = faer_termination_str(report.termination).to_string();
689 (
690 problem,
691 report.n_iter as u64,
692 report.termination.was_successful(),
693 msg,
694 report.cost_history,
695 report.gradient_norm_history,
696 report.params_history,
697 )
698 }
699 };
700 let solve_ns = _solve_t0.elapsed().as_nanos();
701
702 Ok(LmSolveOutcome {
703 result_problem,
704 n_iter_val,
705 success_val,
706 message_val,
707 solve_ns,
708 cost_history,
709 gradient_norm_history,
710 params_history,
711 })
712}
713
714type FreeParamSpecs = (HashMap<String, f64>, Vec<f64>, Vec<(f64, f64)>, Vec<f64>);
716
717type NodeParamBuffers = (Vec<Vec<f64>>, Vec<(usize, usize)>, Vec<(usize, usize)>);
720
721fn collect_free_param_specs(
729 graph: &FitGraphSpec,
730 free_keys: &[String],
731) -> Result<FreeParamSpecs, CoreError> {
732 let mut all_params: HashMap<String, f64> = HashMap::new();
733 for node in &graph.nodes {
734 for (pname, pspec) in &node.parameters {
735 let key = format!("{}.{}", node.id, pname);
736 all_params.insert(key, pspec.value);
737 }
738 }
739
740 let mut init_vals: Vec<f64> = Vec::with_capacity(free_keys.len());
741 let mut bounds: Vec<(f64, f64)> = Vec::with_capacity(free_keys.len());
742 let mut scales: Vec<f64> = Vec::with_capacity(free_keys.len());
743
744 for key in free_keys {
745 let (node_id, param_name) = key
747 .split_once('.')
748 .ok_or_else(|| SolverError::Dispatch(format!("malformed free key: '{}'", key)))?;
749
750 let node = graph
751 .nodes
752 .iter()
753 .find(|n| n.id == node_id)
754 .ok_or_else(|| SolverError::Dispatch(format!("node '{}' not found", node_id)))?;
755
756 let pspec = node.parameters.get(param_name).ok_or_else(|| {
757 SolverError::Dispatch(format!(
758 "param '{}' not found in node '{}'",
759 param_name, node_id
760 ))
761 })?;
762
763 init_vals.push(pspec.value);
764 bounds.push((pspec.min, pspec.max));
765 let s = pspec.scale.unwrap_or(1.0);
768 scales.push(if s.is_finite() && s > 0.0 { s } else { 1.0 });
769 }
770
771 Ok((all_params, init_vals, bounds, scales))
772}
773
774fn build_node_param_buffers(
787 cg: &CompiledGraph,
788 all_params: &HashMap<String, f64>,
789 bounds: &[(f64, f64)],
790 init_vals_clone: &[f64],
791) -> Result<NodeParamBuffers, CoreError> {
792 let mut node_param_bufs: Vec<Vec<f64>> = (0..cg.nodes.len())
793 .map(|i| {
794 cg.node_params(i, all_params)
795 .unwrap_or_else(|_| vec![0.0; cg.nodes[i].param_names.len()])
796 })
797 .collect();
798
799 let free_to_node_param: Vec<(usize, usize)> = {
800 let mut mapping = vec![(0usize, 0usize); cg.free_keys.len()];
803 for (node_idx, pairs) in cg.node_free_cols.iter().enumerate() {
804 for &(local_idx, col) in pairs {
805 mapping[col] = (node_idx, local_idx);
806 }
807 }
808 mapping
809 };
810
811 for (i, &(node_idx, param_pos)) in free_to_node_param.iter().enumerate() {
813 let (lo, hi) = bounds[i];
814 node_param_bufs[node_idx][param_pos] = init_vals_clone[i].clamp(lo, hi);
815 }
816
817 let tied_to_node_param: Vec<(usize, usize)> = cg
818 .tied_plan
819 .order
820 .iter()
821 .map(|tp| {
822 let (nid, pname) = tp
823 .target
824 .split_once('.')
825 .ok_or_else(|| SolverError::MalformedTiedTarget(tp.target.clone()))?;
826 let ni = cg
827 .nodes
828 .iter()
829 .position(|n| n.id == nid)
830 .ok_or_else(|| SolverError::TiedTargetNodeMissing(nid.to_string()))?;
831 let pos = cg.nodes[ni]
832 .param_names
833 .iter()
834 .position(|p| p == pname)
835 .ok_or_else(|| SolverError::TiedTargetParamMissing(pname.to_string()))?;
836 Ok::<(usize, usize), CoreError>((ni, pos))
837 })
838 .collect::<Result<_, _>>()?;
839
840 Ok((node_param_bufs, free_to_node_param, tied_to_node_param))
841}
842
843fn run_varpro_guarded(
854 graph: &FitGraphSpec,
855 datasets: &[MeasurementSpec],
856 options: &FitOptionsSpec,
857) -> Result<FitResultSpec, CoreError> {
858 let result = solve_varpro_path(graph, datasets, options)?;
859 Ok(finalize_varpro_result(graph, datasets, result))
860}
861
862fn finalize_varpro_result(
872 graph: &FitGraphSpec,
873 datasets: &[MeasurementSpec],
874 mut result: FitResultSpec,
875) -> FitResultSpec {
876 let x_all: Vec<f64> = datasets.iter().flat_map(point_major_x).collect();
877 let y_all: Vec<f64> = datasets
878 .iter()
879 .flat_map(|ds| ds.y.iter().copied())
880 .collect();
881 let final_flat: HashMap<String, f64> = result
882 .parameters
883 .iter()
884 .map(|(k, p)| (k.clone(), p.value))
885 .collect();
886 let vp_free_keys: Vec<String> = result
887 .parameters
888 .iter()
889 .filter(|(_, p)| p.vary)
890 .map(|(k, _)| k.clone())
891 .collect();
892 let (success, message) = postfit::apply_postfit_guards(
893 graph,
894 &vp_free_keys,
895 &final_flat,
896 &x_all,
897 &y_all,
898 result.r_squared,
899 vp_free_keys.len(),
900 result.success,
901 result.message.clone(),
902 );
903 result.success = success;
904 result.message = message;
905 result
906}
907
908fn faer_termination_str(t: spectrafit_levenberg_marquardt::Termination) -> &'static str {
910 use spectrafit_levenberg_marquardt::Termination as T;
911 match t {
912 T::Gtol => "converged_gtol",
913 T::Ftol => "converged_ftol",
914 T::Xtol => "converged_xtol",
915 T::ResidualsZero => "residuals_zero",
916 T::MaxEval => "max_iterations",
917 T::NoImprovement => "no_improvement_possible",
918 T::NumericalError => "numerical_error",
919 }
920}
921
922fn map_termination(lm: &levenberg_marquardt::TerminationReason) -> TerminationReason {
926 use levenberg_marquardt::TerminationReason as L;
927 match lm {
928 L::ResidualsZero => TerminationReason::ResidualsZero,
929 L::Orthogonal => TerminationReason::Orthogonal,
930 L::Converged { .. } => TerminationReason::Converged,
931 L::LostPatience => TerminationReason::MaxIterations,
932 L::NoImprovementPossible(_) => TerminationReason::NoImprovementPossible,
933 L::NoParameters => TerminationReason::NoParameters,
934 L::NoResiduals => TerminationReason::NoResiduals,
935 L::WrongDimensions(_) => TerminationReason::WrongDimensions,
936 L::Numerical(_) => TerminationReason::NumericalError,
937 L::User(_) => TerminationReason::UserCancelled,
938 }
939}
940
941#[cfg(test)]
945mod tests {
946 use super::*;
947 use approx::assert_relative_eq;
948 use spectrafit_types::{
949 FitGraphSpec, FitOptionsSpec, MeasurementSpec, ModelNodeSpec, ModelTypeStr, ParameterSpec,
950 };
951 use std::collections::HashMap;
952
953 fn make_param(value: f64, vary: bool) -> ParameterSpec {
956 ParameterSpec {
957 value,
958 min: f64::NEG_INFINITY,
959 max: f64::INFINITY,
960 vary,
961 expr: None,
962 scale: None,
963 }
964 }
965
966 fn default_options() -> FitOptionsSpec {
967 FitOptionsSpec {
968 schema_version: None,
969 solver: "lm".to_string(),
970 max_iterations: 200,
971 tolerance: 1e-8,
972 delta0: None,
973 max_delta: None,
974 eta: None,
975 }
976 }
977
978 fn gaussian(x: f64, amplitude: f64, center: f64, sigma: f64) -> f64 {
980 amplitude * (-(x - center).powi(2) / (2.0 * sigma * sigma)).exp()
981 }
982
983 #[test]
986 fn test_gaussian_recovery() {
987 let (true_a, true_c, true_s) = (5.0_f64, 2.0_f64, 0.5_f64);
989
990 let n = 50usize;
992 let x: Vec<f64> = (0..n)
993 .map(|i| -1.0 + 6.0 * i as f64 / (n - 1) as f64)
994 .collect();
995 let y: Vec<f64> = x
996 .iter()
997 .map(|&xi| gaussian(xi, true_a, true_c, true_s))
998 .collect();
999
1000 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1002 params.insert("amplitude".into(), make_param(4.0, true));
1003 params.insert("center".into(), make_param(1.8, true));
1004 params.insert("sigma".into(), make_param(0.6, true));
1005
1006 let graph = FitGraphSpec {
1007 schema_version: "0.1".into(),
1008 nodes: vec![ModelNodeSpec {
1009 id: "g1".into(),
1010 model_type: ModelTypeStr::Gaussian,
1011 dataset_index: None,
1012 parameters: params,
1013 }],
1014 expr_edges: vec![],
1015 };
1016
1017 let dataset = MeasurementSpec {
1018 schema_version: None,
1019 x: vec![x],
1020 y,
1021 sigma: None,
1022 label: None,
1023 };
1024
1025 let result = fit(&graph, vec![dataset], &default_options()).expect("fit should not error");
1026
1027 assert!(
1028 result.success,
1029 "LM should converge; message: {}",
1030 result.message
1031 );
1032 assert!(
1033 result.n_iter > 0,
1034 "should have done at least one evaluation"
1035 );
1036 assert!(
1037 result.chi2 < 1e-10,
1038 "chi2 = {} should be near zero",
1039 result.chi2
1040 );
1041
1042 let a = result.parameters["g1.amplitude"].value;
1043 let c = result.parameters["g1.center"].value;
1044 let s = result.parameters["g1.sigma"].value;
1045
1046 assert_relative_eq!(a, true_a, max_relative = 0.01);
1047 assert_relative_eq!(c, true_c, max_relative = 0.01);
1048 assert_relative_eq!(s, true_s, max_relative = 0.01);
1049 }
1050
1051 #[test]
1061 fn fit_rejects_unrecognised_solver_string_instead_of_defaulting_to_lm() {
1062 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1063 params.insert("amplitude".into(), make_param(4.0, true));
1064 params.insert("center".into(), make_param(1.8, true));
1065 params.insert("sigma".into(), make_param(0.6, true));
1066 let graph = FitGraphSpec {
1067 schema_version: "0.1".into(),
1068 nodes: vec![ModelNodeSpec {
1069 id: "g1".into(),
1070 model_type: ModelTypeStr::Gaussian,
1071 dataset_index: None,
1072 parameters: params,
1073 }],
1074 expr_edges: vec![],
1075 };
1076 let dataset = MeasurementSpec {
1077 schema_version: None,
1078 x: vec![vec![0.0, 1.0, 2.0]],
1079 y: vec![1.0, 2.0, 1.0],
1080 sigma: None,
1081 label: None,
1082 };
1083 let opts = FitOptionsSpec {
1084 solver: "lmm".to_string(), ..default_options()
1086 };
1087
1088 let err = fit(&graph, vec![dataset], &opts)
1089 .expect_err("a typo'd solver name must error, not silently run LM");
1090 let msg = err.to_string();
1091 assert!(
1092 msg.contains("lmm"),
1093 "error should name the offending string, got: {msg}"
1094 );
1095 assert!(
1096 msg.contains("expected one of"),
1097 "error should list valid solver names, got: {msg}"
1098 );
1099 }
1100
1101 #[test]
1112 fn fit_rejects_eta_at_or_above_a_quarter_on_dogleg() {
1113 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1114 params.insert("amplitude".into(), make_param(4.0, true));
1115 params.insert("center".into(), make_param(1.8, true));
1116 params.insert("sigma".into(), make_param(0.6, true));
1117 let graph = FitGraphSpec {
1118 schema_version: "0.1".into(),
1119 nodes: vec![ModelNodeSpec {
1120 id: "g1".into(),
1121 model_type: ModelTypeStr::Gaussian,
1122 dataset_index: None,
1123 parameters: params,
1124 }],
1125 expr_edges: vec![],
1126 };
1127 let dataset = MeasurementSpec {
1128 schema_version: None,
1129 x: vec![vec![0.0, 1.0, 2.0]],
1130 y: vec![1.0, 2.0, 1.0],
1131 sigma: None,
1132 label: None,
1133 };
1134 let opts = FitOptionsSpec {
1135 solver: "dogleg".to_string(),
1136 eta: Some(0.5),
1137 ..default_options()
1138 };
1139
1140 let err = fit(&graph, vec![dataset], &opts)
1141 .expect_err("eta >= 0.25 must error, not silently run with an invalid ratio");
1142 let msg = err.to_string();
1143 assert!(
1144 msg.contains("eta"),
1145 "error should name the offending option, got: {msg}"
1146 );
1147 }
1148
1149 #[test]
1150 fn fit_rejects_unrecognised_irls_weight_fn_instead_of_defaulting_to_huber() {
1151 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1152 params.insert("amplitude".into(), make_param(4.0, true));
1153 params.insert("center".into(), make_param(1.8, true));
1154 params.insert("sigma".into(), make_param(0.6, true));
1155 let graph = FitGraphSpec {
1156 schema_version: "0.1".into(),
1157 nodes: vec![ModelNodeSpec {
1158 id: "g1".into(),
1159 model_type: ModelTypeStr::Gaussian,
1160 dataset_index: None,
1161 parameters: params,
1162 }],
1163 expr_edges: vec![],
1164 };
1165 let dataset = MeasurementSpec {
1166 schema_version: None,
1167 x: vec![vec![0.0, 1.0, 2.0]],
1168 y: vec![1.0, 2.0, 1.0],
1169 sigma: None,
1170 label: None,
1171 };
1172 let opts = FitOptionsSpec {
1173 solver: "irls:buisquare".to_string(), ..default_options()
1175 };
1176
1177 let err = fit(&graph, vec![dataset], &opts)
1178 .expect_err("a typo'd irls weight-fn name must error, not silently run Huber");
1179 let msg = err.to_string();
1180 assert!(
1181 msg.contains("buisquare"),
1182 "error should name the offending string, got: {msg}"
1183 );
1184 }
1185
1186 #[test]
1189 fn test_constant_recovery() {
1190 let n = 20usize;
1191 let x: Vec<f64> = (0..n).map(|i| i as f64 / (n - 1) as f64).collect();
1192 let y: Vec<f64> = x
1194 .iter()
1195 .enumerate()
1196 .map(|(i, _)| 3.0 + 1e-6 * (i as f64 * 0.1).sin())
1197 .collect();
1198
1199 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1200 params.insert("c".into(), make_param(0.0, true)); let graph = FitGraphSpec {
1203 schema_version: "0.1".into(),
1204 nodes: vec![ModelNodeSpec {
1205 id: "const1".into(),
1206 model_type: ModelTypeStr::Constant,
1207 dataset_index: None,
1208 parameters: params,
1209 }],
1210 expr_edges: vec![],
1211 };
1212
1213 let dataset = MeasurementSpec {
1214 schema_version: None,
1215 x: vec![x],
1216 y,
1217 sigma: None,
1218 label: None,
1219 };
1220
1221 let result = fit(&graph, vec![dataset], &default_options()).expect("fit should not error");
1222
1223 assert!(
1224 result.success,
1225 "LM should converge; message: {}",
1226 result.message
1227 );
1228
1229 let c_val = result.parameters["const1.c"].value;
1230 assert!(
1231 (c_val - 3.0).abs() < 1e-4,
1232 "recovered constant = {}, expected ≈ 3.0",
1233 c_val
1234 );
1235 }
1236
1237 fn run_gaussian_nd_recovery(d: usize) {
1244 let amp = 3.0_f64;
1245 let centers: Vec<f64> = (0..d).map(|i| -1.0 + 0.5 * i as f64).collect();
1246 let sigmas: Vec<f64> = (0..d).map(|i| 0.8 + 0.1 * i as f64).collect();
1247
1248 let axis_n = if d <= 3 { 7 } else { 5 };
1250 let axis: Vec<f64> = (0..axis_n)
1251 .map(|i| -3.0 + 6.0 * i as f64 / (axis_n - 1) as f64)
1252 .collect();
1253 let mut coords: Vec<Vec<f64>> = vec![vec![]];
1255 for _dim in 0..d {
1256 let mut next = Vec::with_capacity(coords.len() * axis.len());
1257 for prefix in &coords {
1258 for &a in &axis {
1259 let mut p = prefix.clone();
1260 p.push(a);
1261 next.push(p);
1262 }
1263 }
1264 coords = next;
1265 }
1266 let g = |pt: &[f64]| -> f64 {
1267 let mut z = 0.0;
1268 for i in 0..d {
1269 let dx = pt[i] - centers[i];
1270 z -= dx * dx / (2.0 * sigmas[i] * sigmas[i]);
1271 }
1272 amp * z.exp()
1273 };
1274 let y: Vec<f64> = coords.iter().map(|pt| g(pt)).collect();
1275 let n_points = coords.len();
1276 let x: Vec<Vec<f64>> = (0..d)
1278 .map(|dim| (0..n_points).map(|p| coords[p][dim]).collect())
1279 .collect();
1280
1281 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1282 params.insert("amplitude".into(), make_param(2.0, true));
1283 for (i, ¢er) in centers.iter().enumerate() {
1284 params.insert(format!("center_{i}"), make_param(center + 0.3, true));
1285 params.insert(format!("sigma_{i}"), make_param(1.0, true));
1286 }
1287 let graph = FitGraphSpec {
1288 schema_version: "0.1".into(),
1289 nodes: vec![ModelNodeSpec {
1290 id: "gnd".into(),
1291 model_type: ModelTypeStr::GaussianNd,
1292 dataset_index: None,
1293 parameters: params,
1294 }],
1295 expr_edges: vec![],
1296 };
1297 let dataset = MeasurementSpec {
1298 schema_version: None,
1299 x,
1300 y,
1301 sigma: None,
1302 label: None,
1303 };
1304 let result =
1305 fit(&graph, vec![dataset], &default_options()).expect("N-D fit should not error");
1306 assert!(
1307 result.success,
1308 "LM should converge on {d}-D; msg: {}",
1309 result.message
1310 );
1311 let p = &result.parameters;
1312 assert_relative_eq!(p["gnd.amplitude"].value, amp, max_relative = 0.02);
1313 for i in 0..d {
1314 assert_relative_eq!(
1315 p[&format!("gnd.center_{i}")].value,
1316 centers[i],
1317 epsilon = 0.05
1318 );
1319 assert_relative_eq!(
1320 p[&format!("gnd.sigma_{i}")].value,
1321 sigmas[i],
1322 epsilon = 0.05
1323 );
1324 }
1325 }
1326
1327 #[test]
1328 fn gaussian_nd_fit_recovers_3d() {
1329 run_gaussian_nd_recovery(3);
1330 }
1331
1332 #[test]
1333 fn gaussian_nd_fit_recovers_5d_arbitrary_n() {
1334 run_gaussian_nd_recovery(5);
1335 }
1336
1337 fn tied_two_gaussian_graph(k: f64) -> FitGraphSpec {
1352 let mut g1: HashMap<String, ParameterSpec> = HashMap::new();
1353 g1.insert("amplitude".into(), make_param(4.0, true));
1354 g1.insert("center".into(), make_param(-1.0, true));
1355 g1.insert("sigma".into(), make_param(0.5, true));
1356
1357 let mut g2: HashMap<String, ParameterSpec> = HashMap::new();
1358 g2.insert("amplitude".into(), make_param(0.0, false));
1360 g2.insert("center".into(), make_param(1.0, true));
1361 g2.insert("sigma".into(), make_param(0.5, true));
1362
1363 FitGraphSpec {
1364 schema_version: "0.1".into(),
1365 nodes: vec![
1366 ModelNodeSpec {
1367 id: "g1".into(),
1368 model_type: ModelTypeStr::Gaussian,
1369 dataset_index: None,
1370 parameters: g1,
1371 },
1372 ModelNodeSpec {
1373 id: "g2".into(),
1374 model_type: ModelTypeStr::Gaussian,
1375 dataset_index: None,
1376 parameters: g2,
1377 },
1378 ],
1379 expr_edges: vec![spectrafit_types::ExprEdge {
1380 target_node: "g2".into(),
1381 target_param: "amplitude".into(),
1382 expression: format!("{} * g1.amplitude", k),
1383 }],
1384 }
1385 }
1386
1387 #[test]
1389 fn test_tied_amplitude_fit_recovers_ratio() {
1390 let k = 0.5_f64;
1391 let (true_a, true_s) = (5.0_f64, 0.5_f64);
1392 let n = 80usize;
1393 let x: Vec<f64> = (0..n)
1394 .map(|i| -3.0 + 6.0 * i as f64 / (n - 1) as f64)
1395 .collect();
1396 let y: Vec<f64> = x
1397 .iter()
1398 .map(|&xi| gaussian(xi, true_a, -1.0, true_s) + gaussian(xi, k * true_a, 1.0, true_s))
1399 .collect();
1400
1401 let graph = tied_two_gaussian_graph(k);
1402 let dataset = MeasurementSpec {
1403 schema_version: None,
1404 x: vec![x],
1405 y,
1406 sigma: None,
1407 label: None,
1408 };
1409 let result = fit(&graph, vec![dataset], &default_options()).unwrap();
1410
1411 assert!(
1412 result.success,
1413 "tied fit should converge: {}",
1414 result.message
1415 );
1416 let a1 = result.parameters["g1.amplitude"].value;
1417 let a2 = result.parameters["g2.amplitude"].value;
1418 assert_relative_eq!(a2, k * a1, epsilon = 1e-9);
1420 assert_relative_eq!(a1, true_a, max_relative = 1e-3);
1422 }
1423
1424 #[test]
1434 fn test_param_expr_fit_recovers_derived_value() {
1435 let k = 0.5_f64;
1436 let (true_a, true_s) = (5.0_f64, 0.5_f64);
1437 let n = 80usize;
1438 let x: Vec<f64> = (0..n)
1439 .map(|i| -3.0 + 6.0 * i as f64 / (n - 1) as f64)
1440 .collect();
1441 let y: Vec<f64> = x
1443 .iter()
1444 .map(|&xi| gaussian(xi, true_a, -1.0, true_s) + gaussian(xi, k * true_a, 1.0, true_s))
1445 .collect();
1446
1447 let mut g1: HashMap<String, ParameterSpec> = HashMap::new();
1450 g1.insert("amplitude".into(), make_param(4.0, true));
1451 g1.insert("center".into(), make_param(-1.0, true));
1452 g1.insert("sigma".into(), make_param(0.5, true));
1453
1454 let mut g2: HashMap<String, ParameterSpec> = HashMap::new();
1455 let mut tied_amp = make_param(0.0, false);
1456 tied_amp.expr = Some(format!("{} * g1.amplitude", k));
1457 g2.insert("amplitude".into(), tied_amp);
1458 g2.insert("center".into(), make_param(1.0, true));
1459 g2.insert("sigma".into(), make_param(0.5, true));
1460
1461 let graph = FitGraphSpec {
1462 schema_version: "0.1".into(),
1463 nodes: vec![
1464 ModelNodeSpec {
1465 id: "g1".into(),
1466 model_type: ModelTypeStr::Gaussian,
1467 dataset_index: None,
1468 parameters: g1,
1469 },
1470 ModelNodeSpec {
1471 id: "g2".into(),
1472 model_type: ModelTypeStr::Gaussian,
1473 dataset_index: None,
1474 parameters: g2,
1475 },
1476 ],
1477 expr_edges: vec![],
1479 };
1480
1481 let dataset = MeasurementSpec {
1482 schema_version: None,
1483 x: vec![x],
1484 y,
1485 sigma: None,
1486 label: None,
1487 };
1488 let result = fit(&graph, vec![dataset], &default_options()).unwrap();
1489
1490 assert!(
1491 result.success,
1492 "param-expr tied fit should converge: {}",
1493 result.message
1494 );
1495 let a1 = result.parameters["g1.amplitude"].value;
1496 let a2 = result.parameters["g2.amplitude"].value;
1497 assert!(
1500 a2.abs() > 1e-6,
1501 "tied param must not be frozen at placeholder 0.0"
1502 );
1503 assert_relative_eq!(a2, k * a1, epsilon = 1e-9);
1504 assert_relative_eq!(a1, true_a, max_relative = 1e-3);
1506 }
1507
1508 #[test]
1511 fn test_tied_fit_reduces_free_param_count() {
1512 let graph = tied_two_gaussian_graph(0.5);
1513 let cg = CompiledGraph::compile(&graph).unwrap();
1514 assert_eq!(cg.free_keys.len(), 5);
1516 assert_eq!(cg.tied_plan.len(), 1);
1517 }
1518
1519 fn gaussian_fit_inputs(scale: Option<f64>) -> (FitGraphSpec, MeasurementSpec) {
1532 let (true_a, true_c, true_s) = (5.0_f64, 2.0_f64, 0.5_f64);
1533 let n = 50usize;
1534 let x: Vec<f64> = (0..n)
1535 .map(|i| -1.0 + 6.0 * i as f64 / (n - 1) as f64)
1536 .collect();
1537 let y: Vec<f64> = x
1538 .iter()
1539 .map(|&xi| gaussian(xi, true_a, true_c, true_s))
1540 .collect();
1541
1542 let mk = |value: f64, scale: Option<f64>| ParameterSpec {
1543 value,
1544 min: f64::NEG_INFINITY,
1545 max: f64::INFINITY,
1546 vary: true,
1547 expr: None,
1548 scale,
1549 };
1550 let (sa, sc, ss) = match scale {
1552 None => (None, None, None),
1553 Some(base) => (Some(base), Some(1.0), Some(1.0 / base)),
1554 };
1555 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1556 params.insert("amplitude".into(), mk(4.0, sa));
1557 params.insert("center".into(), mk(1.8, sc));
1558 params.insert("sigma".into(), mk(0.6, ss));
1559
1560 let graph = FitGraphSpec {
1561 schema_version: "0.1".into(),
1562 nodes: vec![ModelNodeSpec {
1563 id: "g1".into(),
1564 model_type: ModelTypeStr::Gaussian,
1565 dataset_index: None,
1566 parameters: params,
1567 }],
1568 expr_edges: vec![],
1569 };
1570 let dataset = MeasurementSpec {
1571 schema_version: None,
1572 x: vec![x],
1573 y,
1574 sigma: None,
1575 label: None,
1576 };
1577 (graph, dataset)
1578 }
1579
1580 #[test]
1583 fn condition_number_is_some_and_finite_after_gaussian_fit() {
1584 let (graph, dataset) = gaussian_fit_inputs(None);
1585 let result = fit(&graph, vec![dataset], &default_options()).expect("fit should not error");
1586 let cond = result
1587 .condition_number
1588 .expect("well-conditioned Gaussian fit should report a condition number");
1589 assert!(cond.is_finite(), "condition number must be finite: {cond}");
1590 assert!(cond >= 1.0, "condition number must be ≥ 1.0: {cond}");
1591 }
1592
1593 #[test]
1596 fn parameter_scale_changes_effective_conditioning() {
1597 let (g_unscaled, d_unscaled) = gaussian_fit_inputs(None);
1604 let (g_scaled, d_scaled) = gaussian_fit_inputs(Some(1000.0));
1605
1606 let r_unscaled =
1607 fit(&g_unscaled, vec![d_unscaled], &default_options()).expect("unscaled fit");
1608 let r_scaled = fit(&g_scaled, vec![d_scaled], &default_options()).expect("scaled fit");
1609
1610 let c0 = r_unscaled
1611 .condition_number
1612 .expect("unscaled fit should report κ");
1613 let c1 = r_scaled
1614 .condition_number
1615 .expect("scaled fit should report κ");
1616 assert!(
1617 (c0 - c1).abs() > 1e-6,
1618 "Parameter.scale should change effective conditioning: κ_unscaled={c0}, κ_scaled={c1}"
1619 );
1620 }
1621
1622 #[test]
1623 fn parameter_scale_of_one_is_a_bitwise_no_op() {
1624 let (g_none, d_none) = gaussian_fit_inputs(None);
1630 let g_one = {
1631 let mut g = g_none.clone();
1632 for p in g.nodes[0].parameters.values_mut() {
1633 p.scale = Some(1.0);
1634 }
1635 g
1636 };
1637
1638 let r_none = fit(&g_none, vec![d_none.clone()], &default_options()).expect("none fit");
1639 let r_one = fit(&g_one, vec![d_none], &default_options()).expect("scale=1 fit");
1640
1641 assert_eq!(
1642 r_none.condition_number, r_one.condition_number,
1643 "scale=Some(1.0) must reproduce the un-scaled κ bit-for-bit"
1644 );
1645 for (key, p_none) in &r_none.parameters {
1646 let p_one = &r_one.parameters[key];
1647 assert_eq!(
1648 p_none.value, p_one.value,
1649 "scale=1 changed converged value of {key}: {} vs {}",
1650 p_none.value, p_one.value
1651 );
1652 }
1653 }
1654
1655 #[test]
1658 fn auto_routes_to_trf_for_bounded_graph() {
1659 let (true_a, true_c, true_s) = (5.0_f64, 2.0_f64, 0.5_f64);
1664 let n = 60usize;
1665 let x: Vec<f64> = (0..n)
1666 .map(|i| -1.0 + 6.0 * i as f64 / (n - 1) as f64)
1667 .collect();
1668 let y: Vec<f64> = x
1669 .iter()
1670 .map(|&xi| gaussian(xi, true_a, true_c, true_s))
1671 .collect();
1672
1673 let bounded = |value: f64, min: f64, max: f64| ParameterSpec {
1674 value,
1675 min,
1676 max,
1677 vary: true,
1678 expr: None,
1679 scale: None,
1680 };
1681 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1682 params.insert("amplitude".into(), bounded(4.0, 0.0, f64::INFINITY));
1683 params.insert("center".into(), make_param(1.8, true));
1684 params.insert("sigma".into(), bounded(0.6, 1e-6, 10.0)); let graph = FitGraphSpec {
1686 schema_version: "0.1".into(),
1687 nodes: vec![ModelNodeSpec {
1688 id: "g1".into(),
1689 model_type: ModelTypeStr::Gaussian,
1690 dataset_index: None,
1691 parameters: params,
1692 }],
1693 expr_edges: vec![],
1694 };
1695 assert!(
1696 !graph_prefers_varpro(&graph),
1697 "bounded sigma must disqualify the graph from VarPro"
1698 );
1699 let data = MeasurementSpec {
1700 schema_version: None,
1701 x: vec![x],
1702 y,
1703 sigma: None,
1704 label: None,
1705 };
1706 let opts = |s: &str| FitOptionsSpec {
1707 solver: s.to_string(),
1708 ..default_options()
1709 };
1710 let r_auto = fit(&graph, vec![data.clone()], &opts("auto")).expect("auto fit");
1711 let r_trf = fit(&graph, vec![data], &opts("trf")).expect("trf fit");
1712 assert_eq!(
1714 r_auto.n_iter, r_trf.n_iter,
1715 "auto should match trf iterations"
1716 );
1717 assert_relative_eq!(r_auto.chi2, r_trf.chi2, max_relative = 1e-12);
1718 for key in ["g1.amplitude", "g1.center", "g1.sigma"] {
1719 assert_relative_eq!(
1720 r_auto.parameters[key].value,
1721 r_trf.parameters[key].value,
1722 max_relative = 1e-12
1723 );
1724 }
1725 }
1726
1727 #[test]
1730 fn degenerate_peak_collapse_is_flagged_unsuccessful() {
1731 let n = 200usize;
1735 let x: Vec<f64> = (0..n).map(|i| 10.0 * i as f64 / (n - 1) as f64).collect();
1736 let y: Vec<f64> = x.iter().map(|&xi| gaussian(xi, 3.0, 7.0, 0.3)).collect();
1737
1738 let bounded = |value: f64, min: f64, max: f64| ParameterSpec {
1739 value,
1740 min,
1741 max,
1742 vary: true,
1743 expr: None,
1744 scale: None,
1745 };
1746 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1747 params.insert("amplitude".into(), bounded(1.0, 0.0, f64::INFINITY));
1748 params.insert("center".into(), make_param(0.0, true)); params.insert("sigma".into(), bounded(0.3, 1e-6, f64::INFINITY));
1750 let graph = FitGraphSpec {
1751 schema_version: "0.1".into(),
1752 nodes: vec![ModelNodeSpec {
1753 id: "g1".into(),
1754 model_type: ModelTypeStr::Gaussian,
1755 dataset_index: None,
1756 parameters: params,
1757 }],
1758 expr_edges: vec![],
1759 };
1760 let data = MeasurementSpec {
1761 schema_version: None,
1762 x: vec![x],
1763 y,
1764 sigma: None,
1765 label: None,
1766 };
1767 let r = fit(&graph, vec![data], &default_options()).expect("fit runs");
1768 assert!(
1769 !r.success,
1770 "a collapsed-peak fit (R²<0, amplitude≈0) must be flagged unsuccessful, got success=true r2={}",
1771 r.r_squared
1772 );
1773 assert!(
1774 r.message.contains("degenerate_fit"),
1775 "expected degenerate_fit message, got {:?}",
1776 r.message
1777 );
1778 }
1779
1780 #[test]
1783 fn varpro_rejects_dataset_index_scoped_graph() {
1784 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1785 params.insert("amplitude".into(), make_param(1.0, true));
1786 params.insert("center".into(), make_param(0.0, true));
1787 params.insert("sigma".into(), make_param(1.0, true));
1788 let graph = FitGraphSpec {
1789 schema_version: "0.1".into(),
1790 nodes: vec![ModelNodeSpec {
1791 id: "g1".into(),
1792 model_type: ModelTypeStr::Gaussian,
1793 dataset_index: Some(0), parameters: params,
1795 }],
1796 expr_edges: vec![],
1797 };
1798 let x: Vec<f64> = (0..10).map(|i| i as f64).collect();
1799 let data = MeasurementSpec {
1800 schema_version: None,
1801 x: vec![x],
1802 y: vec![0.0; 10],
1803 sigma: None,
1804 label: None,
1805 };
1806 let opts = FitOptionsSpec {
1807 solver: "varpro".into(),
1808 ..default_options()
1809 };
1810 let err = fit(&graph, vec![data], &opts);
1811 assert!(
1812 err.is_err(),
1813 "varpro must reject dataset_index-scoped graphs rather than mis-scope them"
1814 );
1815 let msg = format!("{:?}", err.unwrap_err());
1816 assert!(
1817 msg.contains("dataset_index"),
1818 "error should mention dataset_index, got {msg}"
1819 );
1820 }
1821
1822 fn two_gaussian_param_expr_graph(tie: bool) -> FitGraphSpec {
1828 let mut g1: HashMap<String, ParameterSpec> = HashMap::new();
1829 g1.insert("amplitude".into(), make_param(5.0, true));
1830 g1.insert("center".into(), make_param(-1.0, true));
1831 g1.insert("sigma".into(), make_param(0.5, true));
1832
1833 let mut g2: HashMap<String, ParameterSpec> = HashMap::new();
1834 g2.insert("amplitude".into(), make_param(3.0, true));
1835 g2.insert("center".into(), make_param(1.5, true));
1836 let mut sig2 = make_param(0.5, !tie);
1838 if tie {
1839 sig2.expr = Some("g1.sigma".into());
1840 }
1841 g2.insert("sigma".into(), sig2);
1842
1843 FitGraphSpec {
1844 schema_version: "0.1".into(),
1845 nodes: vec![
1846 ModelNodeSpec {
1847 id: "g1".into(),
1848 model_type: ModelTypeStr::Gaussian,
1849 dataset_index: None,
1850 parameters: g1,
1851 },
1852 ModelNodeSpec {
1853 id: "g2".into(),
1854 model_type: ModelTypeStr::Gaussian,
1855 dataset_index: None,
1856 parameters: g2,
1857 },
1858 ],
1859 expr_edges: vec![],
1860 }
1861 }
1862
1863 #[test]
1864 fn graph_prefers_varpro_false_for_param_expr_tie() {
1865 let graph = two_gaussian_param_expr_graph(true);
1870 assert!(
1871 !graph_prefers_varpro(&graph),
1872 "a Parameter.expr tie must disqualify the graph from VarPro auto-routing"
1873 );
1874 }
1875
1876 #[test]
1877 fn graph_prefers_varpro_true_for_untied_unbounded_separable() {
1878 let graph = two_gaussian_param_expr_graph(false);
1881 assert!(
1882 graph_prefers_varpro(&graph),
1883 "an untied, unbounded, separable graph must remain VarPro-eligible"
1884 );
1885 }
1886
1887 #[test]
1888 fn auto_and_explicit_varpro_agree_on_success() {
1889 let graph = two_gaussian_param_expr_graph(false);
1901 assert!(
1902 graph_prefers_varpro(&graph),
1903 "fixture must actually route auto -> varpro, or this proves nothing"
1904 );
1905
1906 let n = 64usize;
1907 let x: Vec<f64> = (0..n)
1908 .map(|i| -3.0 + 6.0 * i as f64 / (n - 1) as f64)
1909 .collect();
1910 let y: Vec<f64> = x.iter().map(|xi| (20.0 * xi).sin()).collect();
1916 let data = MeasurementSpec {
1917 schema_version: None,
1918 x: vec![x],
1919 y,
1920 sigma: None,
1921 label: None,
1922 };
1923 let opts = |s: &str| FitOptionsSpec {
1924 solver: s.to_string(),
1925 ..default_options()
1926 };
1927
1928 let r_auto = fit(&graph, vec![data.clone()], &opts("auto")).expect("auto fit");
1929 let r_varpro = fit(&graph, vec![data], &opts("varpro")).expect("varpro fit");
1930
1931 assert_eq!(
1932 r_auto.success, r_varpro.success,
1933 "auto and explicit varpro must reach the same success verdict on the \
1934 same graph and data; a difference means one path skipped the \
1935 post-fit guards in `finalize_varpro_result`"
1936 );
1937 assert_relative_eq!(r_auto.chi2, r_varpro.chi2, max_relative = 1e-12);
1938 }
1939
1940 #[test]
1941 fn varpro_explicit_rejects_param_expr_tie() {
1942 let graph = two_gaussian_param_expr_graph(true);
1945 let n = 64usize;
1946 let x: Vec<f64> = (0..n)
1947 .map(|i| -3.0 + 6.0 * i as f64 / (n - 1) as f64)
1948 .collect();
1949 let y: Vec<f64> = x
1950 .iter()
1951 .map(|&xi| gaussian(xi, 5.0, -1.0, 0.5) + gaussian(xi, 3.0, 1.5, 0.5))
1952 .collect();
1953 let data = MeasurementSpec {
1954 schema_version: None,
1955 x: vec![x],
1956 y,
1957 sigma: None,
1958 label: None,
1959 };
1960 let opts = FitOptionsSpec {
1961 solver: "varpro".into(),
1962 ..default_options()
1963 };
1964 let msg = format!(
1965 "{}",
1966 fit(&graph, vec![data], &opts).expect_err(
1967 "solver='varpro' with a Parameter.expr tie must error, not drop the tie"
1968 )
1969 );
1970 assert!(
1971 msg.contains("tied parameters") || msg.contains("Parameter.expr"),
1972 "error should name the tied-params limitation, got: {msg}"
1973 );
1974 }
1975
1976 #[test]
1981 fn graph_with_unbounded_amplitude_but_bounded_sigma_is_not_varpro() {
1982 let bounded_sigma = ParameterSpec {
1983 value: 0.6,
1984 min: 0.1,
1985 max: 2.0,
1986 vary: true,
1987 expr: None,
1988 scale: None,
1989 };
1990 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
1991 params.insert("amplitude".into(), make_param(4.0, true)); params.insert("center".into(), make_param(0.0, true)); params.insert("sigma".into(), bounded_sigma); let graph = FitGraphSpec {
1995 schema_version: "0.1".into(),
1996 nodes: vec![ModelNodeSpec {
1997 id: "g1".into(),
1998 model_type: ModelTypeStr::Gaussian,
1999 dataset_index: None,
2000 parameters: params,
2001 }],
2002 expr_edges: vec![],
2003 };
2004 assert!(
2005 !graph_prefers_varpro(&graph),
2006 "a bounded nonlinear param must block VarPro auto-routing (bounds would be ignored)"
2007 );
2008 }
2009
2010 fn make_varpro_inputs() -> (FitGraphSpec, MeasurementSpec) {
2016 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
2017 params.insert("amplitude".into(), make_param(1.0, true));
2018 params.insert("center".into(), make_param(0.0, true));
2019 params.insert("sigma".into(), make_param(1.0, true));
2020 let graph = FitGraphSpec {
2021 schema_version: "0.1".into(),
2022 nodes: vec![ModelNodeSpec {
2023 id: "g1".into(),
2024 model_type: ModelTypeStr::Gaussian,
2025 dataset_index: None,
2026 parameters: params,
2027 }],
2028 expr_edges: vec![],
2029 };
2030 let x: Vec<f64> = (0..10).map(|i| i as f64 * 0.1).collect();
2031 let y: Vec<f64> = x.iter().map(|&xi| gaussian(xi, 1.0, 0.0, 1.0)).collect();
2032 let dataset = MeasurementSpec {
2033 schema_version: None,
2034 x: vec![x],
2035 y,
2036 sigma: None,
2037 label: None,
2038 };
2039 (graph, dataset)
2040 }
2041
2042 #[test]
2046 fn varpro_with_expr_edges_emits_solver_error_variant() {
2047 use spectrafit_types::ExprEdge;
2048
2049 let (mut graph, dataset) = make_varpro_inputs();
2050 graph.nodes[0]
2052 .parameters
2053 .insert("amplitude".to_string(), make_param(1.0, false));
2054 graph.expr_edges.push(ExprEdge {
2055 target_node: "g1".to_string(),
2056 target_param: "amplitude".to_string(),
2057 expression: "2.0".to_string(),
2058 });
2059
2060 let mut options = default_options();
2061 options.solver = "varpro".to_string();
2062 let err = fit(&graph, vec![dataset], &options).unwrap_err();
2063 let expected: CoreError = SolverError::VarproExprEdgesUnsupported.into();
2064 assert_eq!(format!("{err}"), format!("{expected}"));
2065 }
2066
2067 #[test]
2070 fn varpro_with_dataset_index_emits_solver_error_variant() {
2071 let (mut graph, dataset) = make_varpro_inputs();
2072 graph.nodes[0].dataset_index = Some(0);
2073
2074 let mut options = default_options();
2075 options.solver = "varpro".to_string();
2076 let err = fit(&graph, vec![dataset], &options).unwrap_err();
2077 let expected: CoreError = SolverError::VarproDatasetScopingUnsupported.into();
2078 assert_eq!(format!("{err}"), format!("{expected}"));
2079 }
2080
2081 #[test]
2090 fn fit_rejects_varpro_on_non_separable_graph() {
2091 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
2092 params.insert("amplitude".into(), make_param(1.0, true));
2093 params.insert("center".into(), make_param(0.0, true));
2094 params.insert("sigma".into(), make_param(1.0, true));
2095 params.insert("gamma".into(), make_param(1.0, true));
2096 let graph = FitGraphSpec {
2097 schema_version: "0.1".into(),
2098 nodes: vec![ModelNodeSpec {
2099 id: "v1".into(),
2100 model_type: ModelTypeStr::TrueVoigt,
2101 dataset_index: None,
2102 parameters: params,
2103 }],
2104 expr_edges: vec![],
2105 };
2106 let x: Vec<f64> = (0..10).map(|i| i as f64 * 0.1).collect();
2107 let dataset = MeasurementSpec {
2108 schema_version: None,
2109 x: vec![x],
2110 y: vec![0.0; 10],
2111 sigma: None,
2112 label: None,
2113 };
2114 let opts = FitOptionsSpec {
2115 solver: "varpro".to_string(),
2116 ..default_options()
2117 };
2118 let err = fit(&graph, vec![dataset], &opts).expect_err(
2119 "solver='varpro' on a non-separable graph (true_voigt) must error, not fall back",
2120 );
2121 let expected: CoreError = SolverError::VarproNotSeparable.into();
2122 assert_eq!(format!("{err}"), format!("{expected}"));
2123 }
2124
2125 fn constant_global_fit_inputs() -> (FitGraphSpec, MeasurementSpec) {
2134 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
2135 params.insert(
2136 "c".into(),
2137 ParameterSpec {
2138 value: 0.0,
2139 min: -10.0,
2140 max: 10.0,
2141 vary: true,
2142 expr: None,
2143 scale: None,
2144 },
2145 );
2146 let graph = FitGraphSpec {
2147 schema_version: "0.1".into(),
2148 nodes: vec![ModelNodeSpec {
2149 id: "const1".into(),
2150 model_type: ModelTypeStr::Constant,
2151 dataset_index: None,
2152 parameters: params,
2153 }],
2154 expr_edges: vec![],
2155 };
2156 let n = 10usize;
2157 let x: Vec<f64> = (0..n).map(|i| i as f64).collect();
2158 let y = vec![3.0; n];
2159 let dataset = MeasurementSpec {
2160 schema_version: None,
2161 x: vec![x],
2162 y,
2163 sigma: None,
2164 label: None,
2165 };
2166 (graph, dataset)
2167 }
2168
2169 #[test]
2173 fn dispatch_global_solver_uses_max_iterations_as_de_generation_budget() {
2174 let (graph, dataset) = constant_global_fit_inputs();
2175 let opts = FitOptionsSpec {
2176 solver: "global".to_string(),
2177 max_iterations: 7,
2178 ..default_options()
2179 };
2180 let result = fit(&graph, vec![dataset], &opts).expect("global fit should not error");
2181 let gens = result
2182 .n_de_generations
2183 .expect("global solver must report n_de_generations");
2184 assert!(
2185 gens <= 7,
2186 "max_iterations=7 must cap the DE generation budget at 7, got {gens}"
2187 );
2188 }
2189
2190 #[test]
2193 fn dispatch_global_solver_falls_back_to_default_generation_budget_when_max_iterations_is_zero()
2194 {
2195 let (graph, dataset) = constant_global_fit_inputs();
2196 let opts = FitOptionsSpec {
2197 solver: "global".to_string(),
2198 max_iterations: 0,
2199 ..default_options()
2200 };
2201 let result = fit(&graph, vec![dataset], &opts)
2202 .expect("global fit with max_iterations=0 must still run, using the DE default");
2203 let gens = result
2204 .n_de_generations
2205 .expect("global solver must report n_de_generations");
2206 assert!(
2207 gens <= DeConfig::default().max_gen as u64,
2208 "max_iterations=0 must fall back to DeConfig::default().max_gen (100), got {gens}"
2209 );
2210 }
2211
2212 #[test]
2223 fn run_lm_solve_falls_back_to_default_tolerance_when_non_positive() {
2224 let (graph, dataset) = gaussian_fit_inputs(None);
2225 let opts = FitOptionsSpec {
2226 solver: "lm-legacy".to_string(),
2227 tolerance: 0.0,
2228 ..default_options()
2229 };
2230 let result = fit(&graph, vec![dataset], &opts)
2231 .expect("lm-legacy fit with tolerance<=0.0 should still run on the 1e-8 default");
2232 assert!(
2233 result.success,
2234 "lm-legacy should converge on a clean Gaussian even via the default tolerance; \
2235 message: {}",
2236 result.message
2237 );
2238 }
2239
2240 #[test]
2244 fn run_lm_solve_applies_explicit_delta0_and_max_delta_overrides() {
2245 let (graph, dataset) = gaussian_fit_inputs(None);
2246 let opts = FitOptionsSpec {
2247 solver: "dogleg".to_string(),
2248 delta0: Some(0.5),
2249 max_delta: Some(50.0),
2250 ..default_options()
2251 };
2252 let result = fit(&graph, vec![dataset], &opts)
2253 .expect("dogleg fit with explicit delta0/max_delta should not error");
2254 assert!(
2255 result.success,
2256 "dogleg should converge on a clean Gaussian with an explicit \
2257 delta0/max_delta override; message: {}",
2258 result.message
2259 );
2260 }
2261
2262 #[test]
2271 fn collect_free_param_specs_errors_when_free_key_param_missing_from_node() {
2272 let mut params: HashMap<String, ParameterSpec> = HashMap::new();
2273 params.insert("amplitude".into(), make_param(1.0, true));
2274 let graph = FitGraphSpec {
2275 schema_version: "0.1".into(),
2276 nodes: vec![ModelNodeSpec {
2277 id: "g1".into(),
2278 model_type: ModelTypeStr::Gaussian,
2279 dataset_index: None,
2280 parameters: params,
2281 }],
2282 expr_edges: vec![],
2283 };
2284 let free_keys = vec!["g1.sigma".to_string()];
2288 let err = collect_free_param_specs(&graph, &free_keys)
2289 .expect_err("a free key naming a param absent from its node must error");
2290 let msg = format!("{err}");
2291 assert!(
2292 msg.contains("sigma") && msg.contains("g1"),
2293 "error should name the missing param and its node, got: {msg}"
2294 );
2295 }
2296
2297 #[test]
2305 fn faer_termination_str_maps_numerical_error() {
2306 assert_eq!(
2307 faer_termination_str(spectrafit_levenberg_marquardt::Termination::NumericalError),
2308 "numerical_error"
2309 );
2310 }
2311
2312 #[test]
2317 fn map_termination_covers_every_lm_legacy_variant() {
2318 use levenberg_marquardt::TerminationReason as L;
2319
2320 let cases: Vec<(L, TerminationReason)> = vec![
2321 (L::ResidualsZero, TerminationReason::ResidualsZero),
2322 (L::Orthogonal, TerminationReason::Orthogonal),
2323 (L::LostPatience, TerminationReason::MaxIterations),
2324 (
2325 L::NoImprovementPossible("test"),
2326 TerminationReason::NoImprovementPossible,
2327 ),
2328 (L::NoParameters, TerminationReason::NoParameters),
2329 (L::NoResiduals, TerminationReason::NoResiduals),
2330 (
2331 L::WrongDimensions("test"),
2332 TerminationReason::WrongDimensions,
2333 ),
2334 (L::Numerical("test"), TerminationReason::NumericalError),
2335 (L::User("test"), TerminationReason::UserCancelled),
2336 ];
2337 for (input, expected) in cases {
2338 let mapped = map_termination(&input);
2339 assert_eq!(
2340 mapped, expected,
2341 "map_termination({input:?}) should map to {expected:?}, got {mapped:?}"
2342 );
2343 }
2344 }
2345}