1use std::collections::HashMap;
8
9use nalgebra::DMatrix;
10use rayon::prelude::*;
11use spectrafit_types::CoreError;
12
13use crate::compiler::CompiledGraph;
14use crate::error::GraphError;
15
16const POINTS_PER_THREAD_CUTOFF: usize = 512;
29const MIN_TOTAL_WORK_CUTOFF: usize = 1_500_000;
30
31#[inline]
37fn should_parallel(n_points: usize, work_per_point: usize) -> bool {
38 let n_threads = rayon::current_num_threads();
39 if n_threads <= 1 {
40 return false;
41 }
42
43 let point_cutoff = POINTS_PER_THREAD_CUTOFF.saturating_mul(n_threads);
44 let work_cutoff = MIN_TOTAL_WORK_CUTOFF;
45 let total_work = n_points.saturating_mul(work_per_point.max(1));
46
47 n_points >= point_cutoff && total_work >= work_cutoff
48}
49
50#[inline]
62fn coord_layout(cg: &CompiledGraph, x_len: usize) -> Result<(usize, usize), CoreError> {
63 let stride = cg.n_dims()?;
64 debug_assert!(stride >= 1);
65 if !x_len.is_multiple_of(stride) {
66 return Err(GraphError::XBufferStrideMismatch {
67 x_len,
68 n_dims: stride,
69 }
70 .into());
71 }
72 Ok((stride, x_len / stride))
73}
74
75pub fn evaluate_compiled(
91 cg: &CompiledGraph,
92 flat: &HashMap<String, f64>,
93 x: &[f64],
94) -> Result<Vec<f64>, CoreError> {
95 let node_params: Vec<Vec<f64>> = (0..cg.nodes.len())
97 .map(|i| cg.node_params(i, flat))
98 .collect::<Result<_, _>>()?;
99
100 let (stride, n_points) = coord_layout(cg, x.len())?;
101
102 if scoping_active(cg) {
106 validate_dataset_scope(cg)?;
107 let offsets = cg.dataset_offsets.as_slice();
108 let values: Vec<f64> = x
109 .chunks_exact(stride)
110 .enumerate()
111 .map(|(p, coord)| {
112 let ds = dataset_of_point(offsets, p);
113 let mut sum = 0.0_f64;
114 for (node, params) in cg.nodes.iter().zip(node_params.iter()) {
115 if let Some(i) = node.dataset_index {
116 if i != ds {
117 continue;
118 }
119 }
120 sum += node.model.eval(coord, params);
121 }
122 sum
123 })
124 .collect();
125 return Ok(values);
126 }
127
128 let eval_point = |coord: &[f64]| -> f64 {
132 let mut sum = 0.0_f64;
133 for (node, params) in cg.nodes.iter().zip(node_params.iter()) {
134 sum += node.model.eval(coord, params);
135 }
136 sum
137 };
138
139 let values: Vec<f64> = if should_parallel(n_points, cg.nodes.len().saturating_mul(stride)) {
140 x.par_chunks_exact(stride).map(eval_point).collect()
141 } else {
142 x.chunks_exact(stride).map(eval_point).collect()
143 };
144 Ok(values)
145}
146
147pub fn evaluate_components_compiled(
164 cg: &CompiledGraph,
165 flat: &HashMap<String, f64>,
166 x: &[f64],
167) -> Result<HashMap<String, Vec<f64>>, CoreError> {
168 let mut result: HashMap<String, Vec<f64>> = HashMap::with_capacity(cg.nodes.len());
169
170 let (stride, n_points) = coord_layout(cg, x.len())?;
171 validate_dataset_scope(cg)?;
174
175 for (i, node) in cg.nodes.iter().enumerate() {
176 let params = cg.node_params(i, flat)?;
177
178 let eval_pt = |coord: &[f64]| node.model.eval(coord, ¶ms);
179
180 let mut values: Vec<f64> = if should_parallel(n_points, stride) {
181 x.par_chunks_exact(stride).map(eval_pt).collect()
182 } else {
183 x.chunks_exact(stride).map(eval_pt).collect()
184 };
185
186 if scoping_active(cg) {
189 if let Some(di) = node.dataset_index {
190 let offs = &cg.dataset_offsets;
191 let (a, b) = (offs[di], offs[di + 1]);
192 for (p, v) in values.iter_mut().enumerate() {
193 if p < a || p >= b {
194 *v = 0.0;
195 }
196 }
197 }
198 }
199
200 result.insert(node.id.clone(), values);
201 }
202
203 Ok(result)
204}
205
206pub fn evaluate_compiled_indexed(
226 cg: &CompiledGraph,
227 node_params: &[Vec<f64>],
228 x: &[f64],
229) -> Result<Vec<f64>, CoreError> {
230 let (_stride, n_points) = coord_layout(cg, x.len())?;
233 let mut values = vec![0.0_f64; n_points];
234 evaluate_compiled_indexed_into(cg, node_params, x, &mut values)?;
235 Ok(values)
236}
237
238pub fn evaluate_compiled_indexed_into(
251 cg: &CompiledGraph,
252 node_params: &[Vec<f64>],
253 x: &[f64],
254 out: &mut [f64],
255) -> Result<(), CoreError> {
256 let (stride, n_points) = coord_layout(cg, x.len())?;
257
258 if out.len() != n_points {
259 return Err(GraphError::OutputBufferLength {
260 actual: out.len(),
261 expected: n_points,
262 }
263 .into());
264 }
265
266 let eval_point = |coord: &[f64]| {
267 let mut sum = 0.0_f64;
268 for (node, params) in cg.nodes.iter().zip(node_params.iter()) {
269 sum += node.model.eval(coord, params);
270 }
271 sum
272 };
273
274 if should_parallel(n_points, cg.nodes.len().saturating_mul(stride)) {
275 out.par_iter_mut()
276 .zip(x.par_chunks_exact(stride))
277 .for_each(|(slot, coord)| *slot = eval_point(coord));
278 } else {
279 for (slot, coord) in out.iter_mut().zip(x.chunks_exact(stride)) {
280 *slot = eval_point(coord);
281 }
282 }
283
284 Ok(())
285}
286
287#[inline]
296fn scoping_active(cg: &CompiledGraph) -> bool {
297 cg.dataset_offsets.len() > 2 && cg.nodes.iter().any(|n| n.dataset_index.is_some())
298}
299
300#[inline]
312fn validate_dataset_scope(cg: &CompiledGraph) -> Result<(), CoreError> {
313 if !scoping_active(cg) {
314 return Ok(());
315 }
316 let n_datasets = cg.dataset_offsets.len() - 1;
317 for node in &cg.nodes {
318 if let Some(di) = node.dataset_index {
319 if di >= n_datasets {
320 return Err(GraphError::DatasetIndexOutOfRange {
321 node: node.id.clone(),
322 dataset_index: di,
323 n_datasets,
324 }
325 .into());
326 }
327 }
328 }
329 Ok(())
330}
331
332#[inline]
335fn dataset_of_point(offsets: &[usize], p: usize) -> usize {
336 offsets.partition_point(|&o| o <= p).saturating_sub(1)
337}
338
339fn apply_jacobian_dataset_scope(
345 cg: &CompiledGraph,
346 n_points: usize,
347 n_free: usize,
348 data: &mut [f64],
349) {
350 if !scoping_active(cg) {
351 return;
352 }
353 let offsets = &cg.dataset_offsets;
354 for (node_idx, node) in cg.nodes.iter().enumerate() {
355 let Some(di) = node.dataset_index else {
356 continue;
357 };
358 let (a, b) = (offsets[di], offsets[di + 1]);
359 for &(_local, col) in &cg.node_free_cols[node_idx] {
360 for row in 0..n_points {
361 if row < a || row >= b {
362 data[row * n_free + col] = 0.0;
363 }
364 }
365 }
366 }
367}
368
369pub fn residuals_compiled_indexed_into(
383 cg: &CompiledGraph,
384 node_params: &[Vec<f64>],
385 x: &[f64],
386 y: &[f64],
387 sigma: &[f64],
388 out: &mut [f64],
389) -> Result<(), CoreError> {
390 let (stride, n_points) = coord_layout(cg, x.len())?;
391
392 if out.len() != n_points || y.len() != n_points || sigma.len() != n_points {
393 return Err(GraphError::OutputBufferShape.into());
394 }
395
396 if scoping_active(cg) {
401 validate_dataset_scope(cg)?;
402 let offsets = cg.dataset_offsets.as_slice();
403 for (p, (slot, (coord, (&obs, &s)))) in out
404 .iter_mut()
405 .zip(x.chunks_exact(stride).zip(y.iter().zip(sigma.iter())))
406 .enumerate()
407 {
408 let ds = dataset_of_point(offsets, p);
409 let mut sum = 0.0_f64;
410 for (node, params) in cg.nodes.iter().zip(node_params.iter()) {
411 if let Some(i) = node.dataset_index {
412 if i != ds {
413 continue;
414 }
415 }
416 sum += node.model.eval(coord, params);
417 }
418 *slot = (sum - obs) / s;
419 }
420 return Ok(());
421 }
422
423 if stride == 1 && !cg.nodes.is_empty() {
430 debug_assert_eq!(out.len(), x.len());
431 cg.nodes[0].model.eval_slice_into(x, &node_params[0], out);
433 if cg.nodes.len() > 1 {
434 let mut scratch = vec![0.0_f64; n_points];
435 for (node, params) in cg.nodes.iter().zip(node_params.iter()).skip(1) {
436 node.model.eval_slice_into(x, params, &mut scratch);
437 for (slot, &c) in out.iter_mut().zip(scratch.iter()) {
438 *slot += c;
439 }
440 }
441 }
442 for ((slot, &obs), &s) in out.iter_mut().zip(y.iter()).zip(sigma.iter()) {
443 *slot = (*slot - obs) / s;
444 }
445 return Ok(());
446 }
447
448 let eval_point = |coord: &[f64]| {
449 let mut sum = 0.0_f64;
450 for (node, params) in cg.nodes.iter().zip(node_params.iter()) {
451 sum += node.model.eval(coord, params);
452 }
453 sum
454 };
455
456 if should_parallel(n_points, cg.nodes.len().saturating_mul(stride)) {
457 out.par_iter_mut()
458 .zip(
459 x.par_chunks_exact(stride)
460 .zip(y.par_iter().zip(sigma.par_iter())),
461 )
462 .for_each(|(slot, (coord, (&obs, &s)))| *slot = (eval_point(coord) - obs) / s);
463 } else {
464 for (slot, (coord, (&obs, &s))) in out
465 .iter_mut()
466 .zip(x.chunks_exact(stride).zip(y.iter().zip(sigma.iter())))
467 {
468 *slot = (eval_point(coord) - obs) / s;
469 }
470 }
471
472 Ok(())
473}
474
475pub fn jacobian_compiled(
499 cg: &CompiledGraph,
500 flat: &HashMap<String, f64>,
501 x: &[f64],
502) -> Result<DMatrix<f64>, CoreError> {
503 let node_params: Vec<Vec<f64>> = (0..cg.nodes.len())
505 .map(|i| cg.node_params(i, flat))
506 .collect::<Result<_, _>>()?;
507
508 jacobian_compiled_indexed(cg, &node_params, x)
509}
510
511pub fn jacobian_compiled_indexed(
524 cg: &CompiledGraph,
525 node_params: &[Vec<f64>],
526 x: &[f64],
527) -> Result<DMatrix<f64>, CoreError> {
528 let mut data = Vec::new();
529 jacobian_compiled_indexed_into(cg, node_params, x, &mut data)?;
530 let (_stride, n_points) = coord_layout(cg, x.len())?;
531 let n_free = cg.free_keys.len();
532 if n_free == 0 {
533 Ok(DMatrix::zeros(n_points, 0))
534 } else {
535 Ok(DMatrix::from_row_slice(n_points, n_free, &data))
536 }
537}
538
539pub fn jacobian_compiled_indexed_into(
558 cg: &CompiledGraph,
559 node_params: &[Vec<f64>],
560 x: &[f64],
561 data: &mut Vec<f64>,
562) -> Result<(), CoreError> {
563 let (stride, n_points) = coord_layout(cg, x.len())?;
564 let n_free = cg.free_keys.len();
565
566 validate_dataset_scope(cg)?;
569
570 if n_free == 0 {
571 data.clear();
572 return Ok(());
573 }
574
575 let required = n_points.saturating_mul(n_free);
578 if data.len() != required {
579 data.resize(required, 0.0);
580 } else {
581 data.fill(0.0);
582 }
583
584 let max_params = cg
586 .nodes
587 .iter()
588 .map(|n| n.param_names.len())
589 .max()
590 .unwrap_or(0);
591
592 if should_parallel(n_points, cg.nodes.len().saturating_mul(n_free)) {
593 data.par_chunks_mut(n_free)
597 .zip(x.par_chunks_exact(stride))
598 .for_each_with(vec![0.0_f64; max_params], |scratch, (row_buf, coord)| {
599 for (node_idx, node) in cg.nodes.iter().enumerate() {
600 let free_cols = &cg.node_free_cols[node_idx];
601 if free_cols.is_empty() {
602 continue;
603 }
604 let params = &node_params[node_idx];
605 node.model
606 .jacobian_into(coord, params, &mut scratch[..params.len()]);
607 for &(local_idx, col) in free_cols {
608 row_buf[col] = scratch[local_idx];
609 }
610 }
611 });
612 } else {
613 let mut scratch = vec![0.0_f64; max_params];
616 for (i, coord) in x.chunks_exact(stride).enumerate() {
617 let row_buf = &mut data[i * n_free..(i + 1) * n_free];
618 for (node_idx, node) in cg.nodes.iter().enumerate() {
619 let free_cols = &cg.node_free_cols[node_idx];
620 if free_cols.is_empty() {
621 continue;
622 }
623 let params = &node_params[node_idx];
624 node.model
625 .jacobian_into(coord, params, &mut scratch[..params.len()]);
626 for &(local_idx, col) in free_cols {
627 row_buf[col] = scratch[local_idx];
628 }
629 }
630 }
631 }
632
633 apply_jacobian_dataset_scope(cg, n_points, n_free, data);
634 Ok(())
635}
636
637pub fn jacobian_compiled_indexed_weighted_into(
651 cg: &CompiledGraph,
652 node_params: &[Vec<f64>],
653 x: &[f64],
654 sigma: &[f64],
655 data: &mut Vec<f64>,
656) -> Result<(), CoreError> {
657 let (stride, n_points) = coord_layout(cg, x.len())?;
658 let n_free = cg.free_keys.len();
659
660 if sigma.len() != n_points {
661 return Err(GraphError::OutputBufferLength {
662 actual: sigma.len(),
663 expected: n_points,
664 }
665 .into());
666 }
667
668 validate_dataset_scope(cg)?;
671
672 if n_free == 0 {
673 data.clear();
674 return Ok(());
675 }
676
677 let required = n_points.saturating_mul(n_free);
678 if data.len() != required {
679 data.resize(required, 0.0);
680 } else {
681 data.fill(0.0);
682 }
683
684 if !scoping_active(cg)
689 && stride == 1
690 && cg.nodes.len() == 1
691 && cg.nodes[0].model.n_dims() == 1
692 && n_free == cg.nodes[0].param_names.len()
693 && cg.node_free_cols[0]
694 .iter()
695 .enumerate()
696 .all(|(i, &(local, col))| local == i && col == i)
697 {
698 let node = &cg.nodes[0];
699 node.model.jac_slice_into(x, &node_params[0], data);
700 for i in 0..n_points {
702 let inv_s = 1.0 / sigma[i];
703 for v in &mut data[i * n_free..(i + 1) * n_free] {
704 *v *= inv_s;
705 }
706 }
707 return Ok(());
708 }
709
710 if stride == 1 && !scoping_active(cg) {
719 let mut node_scratch: Vec<f64> = Vec::new();
721 for (node_idx, node) in cg.nodes.iter().enumerate() {
722 let free_cols = &cg.node_free_cols[node_idx];
723 if free_cols.is_empty() {
724 continue;
725 }
726 let n_local = node.param_names.len();
727 node_scratch.clear();
728 node_scratch.resize(n_points * n_local, 0.0);
729 node.model
730 .jac_slice_into(x, &node_params[node_idx], &mut node_scratch);
731 for i in 0..n_points {
732 let inv_s = 1.0 / sigma[i];
733 let base = i * n_local;
734 let row = &mut data[i * n_free..(i + 1) * n_free];
735 for &(local_idx, col) in free_cols {
736 row[col] = node_scratch[base + local_idx] * inv_s;
737 }
738 }
739 }
740 return Ok(());
741 }
742
743 let max_params = cg
744 .nodes
745 .iter()
746 .map(|n| n.param_names.len())
747 .max()
748 .unwrap_or(0);
749
750 if should_parallel(n_points, cg.nodes.len().saturating_mul(n_free)) {
751 data.par_chunks_mut(n_free)
752 .zip(x.par_chunks_exact(stride).zip(sigma.par_iter()))
753 .for_each_with(
754 vec![0.0_f64; max_params],
755 |scratch, (row_buf, (coord, &s))| {
756 for (node_idx, node) in cg.nodes.iter().enumerate() {
757 let free_cols = &cg.node_free_cols[node_idx];
758 if free_cols.is_empty() {
759 continue;
760 }
761 let params = &node_params[node_idx];
762 node.model
763 .jacobian_into(coord, params, &mut scratch[..params.len()]);
764 for &(local_idx, col) in free_cols {
765 row_buf[col] = scratch[local_idx] / s;
766 }
767 }
768 },
769 );
770 } else {
771 let mut scratch = vec![0.0_f64; max_params];
772 for (i, (coord, &s)) in x.chunks_exact(stride).zip(sigma.iter()).enumerate() {
773 let row_buf = &mut data[i * n_free..(i + 1) * n_free];
774 for (node_idx, node) in cg.nodes.iter().enumerate() {
775 let free_cols = &cg.node_free_cols[node_idx];
776 if free_cols.is_empty() {
777 continue;
778 }
779 let params = &node_params[node_idx];
780 node.model
781 .jacobian_into(coord, params, &mut scratch[..params.len()]);
782 for &(local_idx, col) in free_cols {
783 row_buf[col] = scratch[local_idx] / s;
784 }
785 }
786 }
787 }
788
789 apply_jacobian_dataset_scope(cg, n_points, n_free, data);
790 Ok(())
791}
792
793#[cfg(test)]
797mod tests {
798 use super::*;
799 use crate::compiler::CompiledGraph;
800 use approx::assert_relative_eq;
801 use spectrafit_types::{FitGraphSpec, ModelNodeSpec, ModelTypeStr, ParameterSpec};
802
803 fn make_param(value: f64, vary: bool) -> ParameterSpec {
804 ParameterSpec {
805 value,
806 min: f64::NEG_INFINITY,
807 max: f64::INFINITY,
808 vary,
809 expr: None,
810 scale: None,
811 }
812 }
813
814 fn flat(pairs: &[(&str, f64)]) -> HashMap<String, f64> {
815 pairs.iter().map(|(k, v)| (k.to_string(), *v)).collect()
816 }
817
818 fn single_gaussian() -> (FitGraphSpec, HashMap<String, f64>) {
819 let mut params = HashMap::new();
820 params.insert("amplitude".to_string(), make_param(2.0, true));
821 params.insert("center".to_string(), make_param(1.0, true));
822 params.insert("sigma".to_string(), make_param(0.5, true));
823 let graph = FitGraphSpec {
824 schema_version: "0.1".to_string(),
825 nodes: vec![ModelNodeSpec {
826 id: "g".to_string(),
827 model_type: ModelTypeStr::Gaussian,
828 dataset_index: None,
829 parameters: params,
830 }],
831 expr_edges: vec![],
832 };
833 let pf = flat(&[("g.amplitude", 2.0), ("g.center", 1.0), ("g.sigma", 0.5)]);
834 (graph, pf)
835 }
836
837 #[test]
838 fn evaluate_single_gaussian_at_center() {
839 let (graph, pf) = single_gaussian();
840 let cg = CompiledGraph::compile(&graph).unwrap();
841 let result = evaluate_compiled(&cg, &pf, &[1.0]).unwrap();
842 assert_relative_eq!(result[0], 2.0, epsilon = 1e-12);
844 }
845
846 #[test]
847 fn evaluate_multi_point() {
848 let (graph, pf) = single_gaussian();
849 let cg = CompiledGraph::compile(&graph).unwrap();
850 let x = vec![0.0, 1.0, 2.0];
851 let result = evaluate_compiled(&cg, &pf, &x).unwrap();
852 assert_eq!(result.len(), 3);
853 for v in &result {
855 assert!(*v >= 0.0);
856 }
857 }
858
859 #[test]
860 fn evaluate_components_returns_all_nodes() {
861 let mut g_params = HashMap::new();
862 g_params.insert("amplitude".to_string(), make_param(1.0, true));
863 g_params.insert("center".to_string(), make_param(0.0, true));
864 g_params.insert("sigma".to_string(), make_param(1.0, true));
865 let mut c_params = HashMap::new();
866 c_params.insert("c".to_string(), make_param(3.0, true));
867
868 let graph = FitGraphSpec {
869 schema_version: "0.1".to_string(),
870 nodes: vec![
871 ModelNodeSpec {
872 id: "peak".to_string(),
873 model_type: ModelTypeStr::Gaussian,
874 dataset_index: None,
875 parameters: g_params,
876 },
877 ModelNodeSpec {
878 id: "bg".to_string(),
879 model_type: ModelTypeStr::Constant,
880 dataset_index: None,
881 parameters: c_params,
882 },
883 ],
884 expr_edges: vec![],
885 };
886 let pf = flat(&[
887 ("peak.amplitude", 1.0),
888 ("peak.center", 0.0),
889 ("peak.sigma", 1.0),
890 ("bg.c", 3.0),
891 ]);
892 let cg = CompiledGraph::compile(&graph).unwrap();
893 let comps = evaluate_components_compiled(&cg, &pf, &[0.0]).unwrap();
894 assert_eq!(comps.len(), 2);
895 assert_relative_eq!(comps["peak"][0], 1.0, epsilon = 1e-12);
896 assert_relative_eq!(comps["bg"][0], 3.0, epsilon = 1e-12);
897 }
898
899 #[test]
900 fn jacobian_shape_single_gaussian() {
901 let (graph, pf) = single_gaussian();
902 let cg = CompiledGraph::compile(&graph).unwrap();
903 let x = vec![0.0, 0.5, 1.0, 1.5, 2.0];
904 let jac = jacobian_compiled(&cg, &pf, &x).unwrap();
905 assert_eq!(jac.nrows(), 5);
906 assert_eq!(jac.ncols(), 3); }
908
909 #[test]
910 fn jacobian_matches_finite_difference() {
911 let (graph, pf) = single_gaussian();
912 let cg = CompiledGraph::compile(&graph).unwrap();
913 let x = vec![0.8f64]; let h = 1e-6;
915 let jac = jacobian_compiled(&cg, &pf, &x).unwrap();
916
917 let keys = ["g.amplitude", "g.center", "g.sigma"];
918 for (col, key) in keys.iter().enumerate() {
919 let mut pf_plus = pf.clone();
920 *pf_plus.get_mut(*key).unwrap() += h;
921 let f_plus = evaluate_compiled(&cg, &pf_plus, &x).unwrap()[0];
922 let f_base = evaluate_compiled(&cg, &pf, &x).unwrap()[0];
923 let fd = (f_plus - f_base) / h;
924 assert!(
925 (jac[(0, col)] - fd).abs() < 1e-5,
926 "col {} ({}) analytical vs FD mismatch: got {}, expected {}",
927 col,
928 key,
929 jac[(0, col)],
930 fd
931 );
932 }
933 }
934
935 #[test]
936 fn jacobian_fixed_param_excluded() {
937 let mut params = HashMap::new();
939 params.insert("amplitude".to_string(), make_param(2.0, false)); params.insert("center".to_string(), make_param(0.0, true));
941 params.insert("sigma".to_string(), make_param(1.0, true));
942 let graph = FitGraphSpec {
943 schema_version: "0.1".to_string(),
944 nodes: vec![ModelNodeSpec {
945 id: "g".to_string(),
946 model_type: ModelTypeStr::Gaussian,
947 dataset_index: None,
948 parameters: params,
949 }],
950 expr_edges: vec![],
951 };
952 let pf = flat(&[("g.amplitude", 2.0), ("g.center", 0.0), ("g.sigma", 1.0)]);
953 let cg = CompiledGraph::compile(&graph).unwrap();
954 let jac = jacobian_compiled(&cg, &pf, &[0.0]).unwrap();
955 assert_eq!(jac.ncols(), 2, "only 2 free params (center, sigma)");
956 }
957
958 fn two_gauss_plus_const() -> (FitGraphSpec, HashMap<String, f64>) {
965 let mut g1 = HashMap::new();
966 g1.insert("amplitude".to_string(), make_param(2.0, true));
967 g1.insert("center".to_string(), make_param(-1.0, true));
968 g1.insert("sigma".to_string(), make_param(0.7, false)); let mut g2 = HashMap::new();
970 g2.insert("amplitude".to_string(), make_param(1.3, true));
971 g2.insert("center".to_string(), make_param(1.2, true));
972 g2.insert("sigma".to_string(), make_param(0.5, true));
973 let mut bg = HashMap::new();
974 bg.insert("c".to_string(), make_param(0.4, true));
975 let graph = FitGraphSpec {
976 schema_version: "0.1".to_string(),
977 nodes: vec![
978 ModelNodeSpec {
979 id: "a".to_string(),
980 model_type: ModelTypeStr::Gaussian,
981 dataset_index: None,
982 parameters: g1,
983 },
984 ModelNodeSpec {
985 id: "b".to_string(),
986 model_type: ModelTypeStr::Gaussian,
987 dataset_index: None,
988 parameters: g2,
989 },
990 ModelNodeSpec {
991 id: "k".to_string(),
992 model_type: ModelTypeStr::Constant,
993 dataset_index: None,
994 parameters: bg,
995 },
996 ],
997 expr_edges: vec![],
998 };
999 let pf = flat(&[
1000 ("a.amplitude", 2.0),
1001 ("a.center", -1.0),
1002 ("a.sigma", 0.7),
1003 ("b.amplitude", 1.3),
1004 ("b.center", 1.2),
1005 ("b.sigma", 0.5),
1006 ("k.c", 0.4),
1007 ]);
1008 (graph, pf)
1009 }
1010
1011 #[test]
1012 fn multi_node_residuals_and_jacobian_match_scalar_reference() {
1013 let (graph, pf) = two_gauss_plus_const();
1014 let cg = CompiledGraph::compile(&graph).unwrap();
1015 let n = 9usize;
1016 let x: Vec<f64> = (0..n)
1017 .map(|i| -2.0 + 4.0 * i as f64 / (n - 1) as f64)
1018 .collect();
1019 let y: Vec<f64> = x.iter().map(|xi| 0.3 * xi + 0.1).collect();
1020 let sigma: Vec<f64> = (0..n).map(|i| 0.5 + 0.1 * i as f64).collect();
1021
1022 let node_params: Vec<Vec<f64>> = (0..cg.nodes.len())
1023 .map(|i| cg.node_params(i, &pf).unwrap())
1024 .collect();
1025 let n_free = cg.free_keys.len();
1026
1027 let mut res = vec![0.0; n];
1029 residuals_compiled_indexed_into(&cg, &node_params, &x, &y, &sigma, &mut res).unwrap();
1030 let mut jac = Vec::new();
1031 jacobian_compiled_indexed_weighted_into(&cg, &node_params, &x, &sigma, &mut jac).unwrap();
1032
1033 let mut res_ref = vec![0.0; n];
1035 let mut jac_ref = vec![0.0; n * n_free];
1036 for i in 0..n {
1037 let mut sum = 0.0;
1038 for (ni, node) in cg.nodes.iter().enumerate() {
1039 sum += node.model.eval(&[x[i]], &node_params[ni]);
1040 let jn = node.model.jacobian(&[x[i]], &node_params[ni]);
1041 for &(local, col) in &cg.node_free_cols[ni] {
1042 jac_ref[i * n_free + col] = jn[local] / sigma[i];
1043 }
1044 }
1045 res_ref[i] = (sum - y[i]) / sigma[i];
1046 }
1047
1048 for i in 0..n {
1049 assert_relative_eq!(res[i], res_ref[i], epsilon = 1e-12);
1050 }
1051 for k in 0..n * n_free {
1052 assert_relative_eq!(jac[k], jac_ref[k], epsilon = 1e-12);
1053 }
1054 assert_eq!(n_free, 6);
1057 }
1058
1059 fn single_gaussian2d() -> (FitGraphSpec, HashMap<String, f64>) {
1063 let mut params = HashMap::new();
1064 params.insert("amplitude".to_string(), make_param(3.0, true));
1065 params.insert("center_x".to_string(), make_param(0.5, true));
1066 params.insert("center_y".to_string(), make_param(-1.0, true));
1067 params.insert("sigma_x".to_string(), make_param(1.0, true));
1068 params.insert("sigma_y".to_string(), make_param(1.5, true));
1069 let graph = FitGraphSpec {
1070 schema_version: "0.1".to_string(),
1071 nodes: vec![ModelNodeSpec {
1072 id: "g2".to_string(),
1073 model_type: ModelTypeStr::Gaussian2D,
1074 dataset_index: None,
1075 parameters: params,
1076 }],
1077 expr_edges: vec![],
1078 };
1079 let pf = flat(&[
1080 ("g2.amplitude", 3.0),
1081 ("g2.center_x", 0.5),
1082 ("g2.center_y", -1.0),
1083 ("g2.sigma_x", 1.0),
1084 ("g2.sigma_y", 1.5),
1085 ]);
1086 (graph, pf)
1087 }
1088
1089 #[test]
1094 fn executor_strides_two_column_x_for_constant_model() {
1095 let (graph, pf) = single_gaussian2d();
1101 let cg = CompiledGraph::compile(&graph).unwrap();
1102 assert_eq!(cg.n_dims().unwrap(), 2);
1103
1104 let x_flat = vec![0.5, -1.0, 10.0, 10.0, 0.5, -1.0];
1106 let vals = evaluate_compiled(&cg, &pf, &x_flat).unwrap();
1107 assert_eq!(vals.len(), 3, "len/stride = 6/2 = 3 points");
1108
1109 assert_relative_eq!(vals[0], 3.0, epsilon = 1e-12);
1111 assert_relative_eq!(vals[2], 3.0, epsilon = 1e-12);
1112 assert!(
1114 vals[1] < 1e-6,
1115 "far-field point should be ~0, got {}",
1116 vals[1]
1117 );
1118
1119 assert!(evaluate_compiled(&cg, &pf, &[0.0, 0.0, 0.0]).is_err());
1121 }
1122
1123 #[test]
1126 fn executor_one_d_constant_path_unchanged() {
1127 let mut params = HashMap::new();
1128 params.insert("c".to_string(), make_param(7.0, true));
1129 let graph = FitGraphSpec {
1130 schema_version: "0.1".to_string(),
1131 nodes: vec![ModelNodeSpec {
1132 id: "bg".to_string(),
1133 model_type: ModelTypeStr::Constant,
1134 dataset_index: None,
1135 parameters: params,
1136 }],
1137 expr_edges: vec![],
1138 };
1139 let pf = flat(&[("bg.c", 7.0)]);
1140 let cg = CompiledGraph::compile(&graph).unwrap();
1141 assert_eq!(cg.n_dims().unwrap(), 1);
1142 let vals = evaluate_compiled(&cg, &pf, &[0.0, 1.0, 2.0, 3.0]).unwrap();
1143 assert_eq!(vals.len(), 4);
1144 for v in &vals {
1145 assert_relative_eq!(*v, 7.0, epsilon = 1e-12);
1146 }
1147 }
1148
1149 #[test]
1150 fn gaussian2d_grid_evaluation_is_peaked_at_center() {
1151 let (graph, pf_true) = single_gaussian2d();
1160 let cg = CompiledGraph::compile(&graph).unwrap();
1161
1162 let mut x_flat = Vec::new();
1164 let mut coords = Vec::new();
1165 for i in 0..5 {
1166 for j in 0..5 {
1167 let xi = -2.0 + i as f64;
1168 let yj = -3.0 + j as f64;
1169 x_flat.push(xi);
1170 x_flat.push(yj);
1171 coords.push((xi, yj));
1172 }
1173 }
1174 let y = evaluate_compiled(&cg, &pf_true, &x_flat).unwrap();
1175 assert_eq!(y.len(), 25, "len/stride = 50/2 = 25 points");
1176
1177 for &v in &y {
1179 assert!(v > 0.0 && v <= 3.0 + 1e-12, "value out of range: {v}");
1180 }
1181
1182 let (peak_idx, &peak) = y
1185 .iter()
1186 .enumerate()
1187 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
1188 .unwrap();
1189 assert!(peak > 0.0);
1190 let (px, py) = coords[peak_idx];
1191 assert!(
1192 (px - 0.0).abs() <= 1.0 && (py - (-1.0)).abs() <= 1e-12,
1193 "peak at ({px},{py}) not nearest true center (0.5,-1.0)"
1194 );
1195 }
1196
1197 #[test]
1198 fn gaussian2d_jacobian_matches_finite_difference() {
1199 let (graph, pf) = single_gaussian2d();
1203 let cg = CompiledGraph::compile(&graph).unwrap();
1204 let x = vec![0.9f64, -0.2f64];
1206 let h = 1e-6;
1207 let jac = jacobian_compiled(&cg, &pf, &x).unwrap();
1208 assert_eq!(jac.nrows(), 1);
1209 assert_eq!(jac.ncols(), 5);
1210
1211 let keys = [
1212 "g2.amplitude",
1213 "g2.center_x",
1214 "g2.center_y",
1215 "g2.sigma_x",
1216 "g2.sigma_y",
1217 ];
1218 for (col, key) in keys.iter().enumerate() {
1219 let mut pf_plus = pf.clone();
1220 *pf_plus.get_mut(*key).unwrap() += h;
1221 let f_plus = evaluate_compiled(&cg, &pf_plus, &x).unwrap()[0];
1222 let f_base = evaluate_compiled(&cg, &pf, &x).unwrap()[0];
1223 let fd = (f_plus - f_base) / h;
1224 assert!(
1225 (jac[(0, col)] - fd).abs() < 1e-5,
1226 "col {} ({}) analytical vs FD mismatch: got {}, expected {}",
1227 col,
1228 key,
1229 jac[(0, col)],
1230 fd
1231 );
1232 }
1233 }
1234
1235 #[test]
1236 fn dataset_index_scopes_local_node_to_its_dataset() {
1237 let mk_node = |id: &str, ds: Option<usize>, c: f64| ModelNodeSpec {
1241 id: id.to_string(),
1242 model_type: ModelTypeStr::Gaussian,
1243 dataset_index: ds,
1244 parameters: {
1245 let mut p = HashMap::new();
1246 p.insert("amplitude".to_string(), make_param(1.0, true));
1247 p.insert("center".to_string(), make_param(c, true));
1248 p.insert("sigma".to_string(), make_param(0.5, true));
1249 p
1250 },
1251 };
1252 let graph = FitGraphSpec {
1253 schema_version: "0.1".to_string(),
1254 nodes: vec![mk_node("gg", None, 0.0), mk_node("gl", Some(1), 5.0)],
1255 expr_edges: vec![],
1256 };
1257 let mut cg = CompiledGraph::compile(&graph).unwrap();
1258 cg.dataset_offsets = vec![0, 3, 6];
1259
1260 let x = vec![-0.5, 0.0, 0.5, 4.5, 5.0, 5.5];
1261 let pf = flat(&[
1262 ("gg.amplitude", 1.0),
1263 ("gg.center", 0.0),
1264 ("gg.sigma", 0.5),
1265 ("gl.amplitude", 1.0),
1266 ("gl.center", 5.0),
1267 ("gl.sigma", 0.5),
1268 ]);
1269
1270 let comps = evaluate_components_compiled(&cg, &pf, &x).unwrap();
1272 assert!(
1273 comps["gl"][0..3].iter().all(|&v| v == 0.0),
1274 "local node must contribute zero outside its dataset"
1275 );
1276 assert!(
1277 comps["gl"][3..6].iter().any(|&v| v.abs() > 1e-6),
1278 "local node must contribute inside its dataset"
1279 );
1280 assert!(comps["gg"][0..3].iter().any(|&v| v.abs() > 1e-6));
1281
1282 let best = evaluate_compiled(&cg, &pf, &x).unwrap();
1284 assert!(
1285 (best[1] - comps["gg"][1]).abs() < 1e-12,
1286 "best_fit on dataset 0 must exclude the local node"
1287 );
1288
1289 let node_params: Vec<Vec<f64>> = (0..cg.nodes.len())
1291 .map(|i| cg.node_params(i, &pf).unwrap())
1292 .collect();
1293 let mut jdata = Vec::new();
1294 jacobian_compiled_indexed_into(&cg, &node_params, &x, &mut jdata).unwrap();
1295 let n_free = cg.free_keys.len();
1296 let gl_cols: Vec<usize> = cg
1297 .free_keys
1298 .iter()
1299 .enumerate()
1300 .filter(|(_, k)| k.starts_with("gl."))
1301 .map(|(i, _)| i)
1302 .collect();
1303 assert!(!gl_cols.is_empty());
1304 for row in 0..3 {
1305 for &col in &gl_cols {
1306 assert_eq!(
1307 jdata[row * n_free + col],
1308 0.0,
1309 "local node Jacobian must be zero outside its dataset"
1310 );
1311 }
1312 }
1313 }
1314
1315 #[test]
1319 fn out_of_range_dataset_index_errors_not_panics() {
1320 let mk_node = |id: &str, ds: Option<usize>, c: f64| ModelNodeSpec {
1321 id: id.to_string(),
1322 model_type: ModelTypeStr::Gaussian,
1323 dataset_index: ds,
1324 parameters: {
1325 let mut p = HashMap::new();
1326 p.insert("amplitude".to_string(), make_param(1.0, true));
1327 p.insert("center".to_string(), make_param(c, true));
1328 p.insert("sigma".to_string(), make_param(0.5, true));
1329 p
1330 },
1331 };
1332 let graph = FitGraphSpec {
1335 schema_version: "0.1".to_string(),
1336 nodes: vec![mk_node("gg", None, 0.0), mk_node("gl", Some(2), 5.0)],
1337 expr_edges: vec![],
1338 };
1339 let mut cg = CompiledGraph::compile(&graph).unwrap();
1340 cg.dataset_offsets = vec![0, 3, 6]; let x = vec![-0.5, 0.0, 0.5, 4.5, 5.0, 5.5];
1343 let y = vec![0.0; 6];
1344 let sigma = vec![1.0; 6];
1345 let pf = flat(&[
1346 ("gg.amplitude", 1.0),
1347 ("gg.center", 0.0),
1348 ("gg.sigma", 0.5),
1349 ("gl.amplitude", 1.0),
1350 ("gl.center", 5.0),
1351 ("gl.sigma", 0.5),
1352 ]);
1353
1354 let err = evaluate_compiled(&cg, &pf, &x).unwrap_err();
1356 assert!(
1357 format!("{err}").contains("dataset_index"),
1358 "expected a dataset_index range error, got: {err}"
1359 );
1360
1361 assert!(evaluate_components_compiled(&cg, &pf, &x).is_err());
1363
1364 let node_params: Vec<Vec<f64>> = (0..cg.nodes.len())
1366 .map(|i| cg.node_params(i, &pf).unwrap())
1367 .collect();
1368 let mut res = vec![0.0; 6];
1369 assert!(
1370 residuals_compiled_indexed_into(&cg, &node_params, &x, &y, &sigma, &mut res).is_err()
1371 );
1372
1373 let mut jdata = Vec::new();
1375 assert!(jacobian_compiled_indexed_into(&cg, &node_params, &x, &mut jdata).is_err());
1376 let mut jw = Vec::new();
1377 assert!(
1378 jacobian_compiled_indexed_weighted_into(&cg, &node_params, &x, &sigma, &mut jw)
1379 .is_err()
1380 );
1381 }
1382
1383 #[test]
1388 fn x_buffer_stride_mismatch_emits_graph_error_variant() {
1389 let mut params = HashMap::new();
1392 params.insert("amplitude".to_string(), make_param(1.0, true));
1393 params.insert("center_x".to_string(), make_param(0.0, true));
1394 params.insert("center_y".to_string(), make_param(0.0, true));
1395 params.insert("sigma_x".to_string(), make_param(1.0, true));
1396 params.insert("sigma_y".to_string(), make_param(1.0, true));
1397 let graph = FitGraphSpec {
1398 schema_version: "0.1".to_string(),
1399 nodes: vec![ModelNodeSpec {
1400 id: "g".to_string(),
1401 model_type: ModelTypeStr::Gaussian2D,
1402 dataset_index: None,
1403 parameters: params,
1404 }],
1405 expr_edges: vec![],
1406 };
1407 let cg = CompiledGraph::compile(&graph).unwrap();
1408 let pf = flat(&[
1409 ("g.amplitude", 1.0),
1410 ("g.center_x", 0.0),
1411 ("g.center_y", 0.0),
1412 ("g.sigma_x", 1.0),
1413 ("g.sigma_y", 1.0),
1414 ]);
1415
1416 let bad_x: Vec<f64> = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let err = evaluate_compiled(&cg, &pf, &bad_x).unwrap_err();
1418 let expected: CoreError = GraphError::XBufferStrideMismatch {
1419 x_len: 7,
1420 n_dims: 2,
1421 }
1422 .into();
1423 assert_eq!(format!("{err}"), format!("{expected}"));
1424 }
1425}