Skip to main content

spectrafit_graph/
lib.rs

1//! spectrafit-graph — DAG compilation and evaluation engine.
2//!
3//! Exposes three top-level functions consumed by the solver and pyo3 bindings:
4//!   - [`evaluate`]            — sum of all node contributions at given x-points
5//!   - [`evaluate_components`] — per-node contributions
6//!   - [`jacobian`]            — full analytical Jacobian matrix \[n_points × n_free_params\]
7#![warn(missing_docs)]
8
9// Modules are private: the crate's entire public surface is the curated
10// `pub use` list below (house rule 26). Previously these were `pub mod`, so a
11// consumer could reach any item in `compiler`/`executor`/`expr` directly,
12// bypassing this list.
13mod compiler;
14mod error;
15mod executor;
16mod expr;
17
18pub use compiler::CompiledGraph;
19pub use error::GraphError;
20
21use std::collections::HashMap;
22
23use nalgebra::DMatrix;
24use spectrafit_types::{CoreError, FitGraphSpec};
25
26pub use executor::{
27    evaluate_compiled, evaluate_compiled_indexed, evaluate_components_compiled, jacobian_compiled,
28    jacobian_compiled_indexed, jacobian_compiled_indexed_into,
29    jacobian_compiled_indexed_weighted_into, residuals_compiled_indexed_into,
30};
31pub use expr::{parse as parse_expr, BinOp, Expr, TiedParam, TiedPlan};
32
33/// Evaluate the model sum across all nodes at the given x-points.
34///
35/// `params_flat` keys follow the `"node_id.param_name"` convention.
36///
37/// # Errors
38/// Returns [`CoreError::Eval`] if: the graph fails to compile (unknown model
39/// type, missing parameter, duplicate node id, or a bad `expr_edges` cycle —
40/// see [`CompiledGraph::compile`]); `params_flat` is missing a
41/// `"node_id.param_name"` entry a node needs; `x.len()` is not an exact
42/// multiple of the graph's coordinate dimensionality; or a node's
43/// `dataset_index` is out of range for the compiled `dataset_offsets`.
44pub fn evaluate(
45    graph: &FitGraphSpec,
46    params_flat: &HashMap<String, f64>,
47    x: &[f64],
48) -> Result<Vec<f64>, CoreError> {
49    let cg = compiler::CompiledGraph::compile(graph)?;
50    executor::evaluate_compiled(&cg, params_flat, x)
51}
52
53/// Evaluate each node independently.
54///
55/// Returns a map `{ "node_id" => Vec<f64> }`.
56///
57/// # Errors
58/// Returns [`CoreError::Eval`] if: the graph fails to compile (see
59/// [`CompiledGraph::compile`]); `params_flat` is missing a required
60/// `"node_id.param_name"` entry; `x.len()` is not an exact multiple of the
61/// graph's coordinate dimensionality; or a node's `dataset_index` is out of
62/// range for the compiled `dataset_offsets`.
63pub fn evaluate_components(
64    graph: &FitGraphSpec,
65    params_flat: &HashMap<String, f64>,
66    x: &[f64],
67) -> Result<HashMap<String, Vec<f64>>, CoreError> {
68    let cg = compiler::CompiledGraph::compile(graph)?;
69    executor::evaluate_components_compiled(&cg, params_flat, x)
70}
71
72/// Compute the full analytical Jacobian matrix for use by the solver.
73///
74/// Returns a matrix of shape `[n_points × n_free_params]`.
75///
76/// Free parameters are ordered by (node_id alphabetically, then param order
77/// as declared by the model's `param_names()`).
78///
79/// # Errors
80/// Returns [`CoreError::Eval`] if: the graph fails to compile (see
81/// [`CompiledGraph::compile`]); `params_flat` is missing a required
82/// `"node_id.param_name"` entry; `x.len()` is not an exact multiple of the
83/// graph's coordinate dimensionality; or a node's `dataset_index` is out of
84/// range for the compiled `dataset_offsets`.
85pub fn jacobian(
86    graph: &FitGraphSpec,
87    params_flat: &HashMap<String, f64>,
88    x: &[f64],
89) -> Result<DMatrix<f64>, CoreError> {
90    let cg = compiler::CompiledGraph::compile(graph)?;
91    executor::jacobian_compiled(&cg, params_flat, x)
92}
93
94// ---------------------------------------------------------------------------
95// Tests
96// ---------------------------------------------------------------------------
97#[cfg(test)]
98mod tests {
99    use super::*;
100    use approx::assert_relative_eq;
101    use spectrafit_types::{FitGraphSpec, ModelNodeSpec, ModelTypeStr, ParameterSpec};
102
103    /// Build a minimal `ParameterSpec` with `vary = true`.
104    fn make_param(value: f64, vary: bool) -> ParameterSpec {
105        ParameterSpec {
106            value,
107            min: f64::NEG_INFINITY,
108            max: f64::INFINITY,
109            vary,
110            expr: None,
111            scale: None,
112        }
113    }
114
115    /// Convenience: build `params_flat` from a list of `("node.param", value)` pairs.
116    fn flat(pairs: &[(&str, f64)]) -> HashMap<String, f64> {
117        pairs.iter().map(|(k, v)| (k.to_string(), *v)).collect()
118    }
119
120    /// Build a single-Gaussian `FitGraphSpec`.
121    fn single_gaussian_graph() -> (FitGraphSpec, HashMap<String, f64>) {
122        let mut params = HashMap::new();
123        params.insert("amplitude".to_string(), make_param(3.0, true));
124        params.insert("center".to_string(), make_param(0.0, true));
125        params.insert("sigma".to_string(), make_param(1.0, true));
126
127        let graph = FitGraphSpec {
128            schema_version: "0.1".to_string(),
129            nodes: vec![ModelNodeSpec {
130                id: "g1".to_string(),
131                model_type: ModelTypeStr::Gaussian,
132                dataset_index: None,
133                parameters: params,
134            }],
135            expr_edges: vec![],
136        };
137        let pf = flat(&[("g1.amplitude", 3.0), ("g1.center", 0.0), ("g1.sigma", 1.0)]);
138        (graph, pf)
139    }
140
141    // -----------------------------------------------------------------------
142    // Test 1: single Gaussian — evaluate at center returns amplitude
143    // -----------------------------------------------------------------------
144    #[test]
145    fn test_single_gaussian_at_center() {
146        let (graph, pf) = single_gaussian_graph();
147        let result = evaluate(&graph, &pf, &[0.0]).unwrap();
148        assert_eq!(result.len(), 1);
149        // At x=center the Gaussian equals amplitude (exp(0)=1).
150        assert_relative_eq!(result[0], 3.0, epsilon = 1e-12);
151    }
152
153    // -----------------------------------------------------------------------
154    // Test 2: two-node sum — Gaussian + Constant
155    // -----------------------------------------------------------------------
156    #[test]
157    fn test_two_node_sum_gaussian_constant() {
158        let mut g_params = HashMap::new();
159        g_params.insert("amplitude".to_string(), make_param(2.0, true));
160        g_params.insert("center".to_string(), make_param(0.0, true));
161        g_params.insert("sigma".to_string(), make_param(1.0, true));
162
163        let mut c_params = HashMap::new();
164        c_params.insert("c".to_string(), make_param(5.0, true));
165
166        let graph = FitGraphSpec {
167            schema_version: "0.1".to_string(),
168            nodes: vec![
169                ModelNodeSpec {
170                    id: "g1".to_string(),
171                    model_type: ModelTypeStr::Gaussian,
172                    dataset_index: None,
173                    parameters: g_params,
174                },
175                ModelNodeSpec {
176                    id: "bg".to_string(),
177                    model_type: ModelTypeStr::Constant,
178                    dataset_index: None,
179                    parameters: c_params,
180                },
181            ],
182            expr_edges: vec![],
183        };
184        let pf = flat(&[
185            ("g1.amplitude", 2.0),
186            ("g1.center", 0.0),
187            ("g1.sigma", 1.0),
188            ("bg.c", 5.0),
189        ]);
190
191        // At x=0: Gaussian(0) = 2.0; Constant = 5.0  → sum = 7.0
192        let result = evaluate(&graph, &pf, &[0.0]).unwrap();
193        assert_relative_eq!(result[0], 7.0, epsilon = 1e-12);
194
195        // At x=1 (one sigma from center): Gaussian(1) = 2*exp(-0.5); sum += 5.0
196        let result2 = evaluate(&graph, &pf, &[1.0]).unwrap();
197        let expected = 2.0 * (-0.5f64).exp() + 5.0;
198        assert_relative_eq!(result2[0], expected, epsilon = 1e-12);
199    }
200
201    // -----------------------------------------------------------------------
202    // Test 3: evaluate_components returns correct node keys
203    // -----------------------------------------------------------------------
204    #[test]
205    fn test_evaluate_components_keys() {
206        let mut g_params = HashMap::new();
207        g_params.insert("amplitude".to_string(), make_param(1.0, true));
208        g_params.insert("center".to_string(), make_param(0.0, true));
209        g_params.insert("sigma".to_string(), make_param(1.0, true));
210
211        let mut c_params = HashMap::new();
212        c_params.insert("c".to_string(), make_param(2.0, false));
213
214        let graph = FitGraphSpec {
215            schema_version: "0.1".to_string(),
216            nodes: vec![
217                ModelNodeSpec {
218                    id: "peak".to_string(),
219                    model_type: ModelTypeStr::Gaussian,
220                    dataset_index: None,
221                    parameters: g_params,
222                },
223                ModelNodeSpec {
224                    id: "baseline".to_string(),
225                    model_type: ModelTypeStr::Constant,
226                    dataset_index: None,
227                    parameters: c_params,
228                },
229            ],
230            expr_edges: vec![],
231        };
232        let pf = flat(&[
233            ("peak.amplitude", 1.0),
234            ("peak.center", 0.0),
235            ("peak.sigma", 1.0),
236            ("baseline.c", 2.0),
237        ]);
238
239        let comps = evaluate_components(&graph, &pf, &[0.0, 1.0]).unwrap();
240        assert!(comps.contains_key("peak"), "must have 'peak' component");
241        assert!(
242            comps.contains_key("baseline"),
243            "must have 'baseline' component"
244        );
245        assert_eq!(comps.len(), 2);
246
247        // peak at x=0 → 1.0; baseline → 2.0
248        assert_relative_eq!(comps["peak"][0], 1.0, epsilon = 1e-12);
249        assert_relative_eq!(comps["baseline"][0], 2.0, epsilon = 1e-12);
250    }
251
252    // -----------------------------------------------------------------------
253    // Test 4: jacobian shape is [n_points × n_free_params]
254    // -----------------------------------------------------------------------
255    #[test]
256    fn test_jacobian_shape() {
257        let (graph, pf) = single_gaussian_graph();
258        // 5 x-points, 3 free params (amplitude, center, sigma)
259        let x: Vec<f64> = vec![0.0, 0.5, 1.0, 1.5, 2.0];
260        let jac = jacobian(&graph, &pf, &x).unwrap();
261        assert_eq!(jac.nrows(), 5, "rows = n_points");
262        assert_eq!(jac.ncols(), 3, "cols = n_free_params");
263    }
264
265    // -----------------------------------------------------------------------
266    // Test 5: jacobian values match finite-difference for Gaussian
267    // -----------------------------------------------------------------------
268    #[test]
269    fn test_jacobian_vs_finite_diff_gaussian() {
270        let (graph, pf) = single_gaussian_graph();
271        let x = vec![0.5f64];
272        let h = 1e-6;
273
274        let jac = jacobian(&graph, &pf, &x).unwrap();
275
276        // free_keys are sorted by node_id ("g1") then model param order:
277        //   col 0 = g1.amplitude, col 1 = g1.center, col 2 = g1.sigma
278        let param_keys = ["g1.amplitude", "g1.center", "g1.sigma"];
279        for (col, key) in param_keys.iter().enumerate() {
280            let mut pf_plus = pf.clone();
281            *pf_plus.get_mut(*key).unwrap() += h;
282            let f_plus = evaluate(&graph, &pf_plus, &x).unwrap()[0];
283            let f_base = evaluate(&graph, &pf, &x).unwrap()[0];
284            let fd = (f_plus - f_base) / h;
285            assert!(
286                (jac[(0, col)] - fd).abs() < 1e-5,
287                "Jacobian col {} ({}) mismatch: got {}, expected {}",
288                col,
289                key,
290                jac[(0, col)],
291                fd
292            );
293        }
294    }
295
296    // -----------------------------------------------------------------------
297    // Test 6: compile() rejects a graph that would have cyclic param targets
298    // -----------------------------------------------------------------------
299    #[test]
300    fn test_compile_rejects_duplicate_expr_target() {
301        use spectrafit_types::ExprEdge;
302
303        let mut g_params = HashMap::new();
304        g_params.insert("amplitude".to_string(), make_param(1.0, false));
305        g_params.insert("center".to_string(), make_param(0.0, false));
306        g_params.insert("sigma".to_string(), make_param(1.0, false));
307
308        let graph = FitGraphSpec {
309            schema_version: "0.1".to_string(),
310            nodes: vec![ModelNodeSpec {
311                id: "g1".to_string(),
312                model_type: ModelTypeStr::Gaussian,
313                dataset_index: None,
314                parameters: g_params,
315            }],
316            // Two edges pointing to the same target param → cycle/conflict
317            expr_edges: vec![
318                ExprEdge {
319                    target_node: "g1".to_string(),
320                    target_param: "amplitude".to_string(),
321                    expression: "2.0".to_string(),
322                },
323                ExprEdge {
324                    target_node: "g1".to_string(),
325                    target_param: "amplitude".to_string(),
326                    expression: "3.0".to_string(),
327                },
328            ],
329        };
330
331        let result = compiler::CompiledGraph::compile(&graph);
332        assert!(
333            result.is_err(),
334            "compile should reject duplicate expr targets"
335        );
336    }
337}