Skip to main content

spectrafit_varpro/
model.rs

1//! Manual implementation of [`SeparableNonlinearModel`] backed by a compiled
2//! spectrafit graph.
3//!
4//! # Model layout
5//!
6//! For a graph with `M` model nodes, each node contributes **one column** to the
7//! basis-function matrix Φ (n_data × M).  For separable nodes the column is
8//! evaluated at `amplitude = 1.0` (amplitude is the linear coefficient for the
9//! varpro solve).
10//!
11//! The global nonlinear parameter vector α is the concatenation of each node's
12//! non-amplitude, free parameters in graph-node order, then model-param order.
13//! Fixed params are injected as constants; `amplitude` is excluded (it is a
14//! linear coefficient solved by varpro, not part of α).
15//!
16//! # Invariant nodes
17//! `constant` and `linear` nodes have no nonlinear parameters; their Φ column
18//! is evaluated once at construction time and never changes.
19
20use nalgebra::{DVector, Dyn, OMatrix, OVector};
21use spectrafit_graph::CompiledGraph;
22use spectrafit_models::Model;
23use spectrafit_types::{FitGraphSpec, MeasurementSpec, ParameterSpec};
24use std::collections::HashMap;
25use varpro::model::errors::ModelError;
26use varpro::model::SeparableNonlinearModel;
27
28// ---------------------------------------------------------------------------
29// Internal helpers
30// ---------------------------------------------------------------------------
31
32/// Per-node parameter bookkeeping.
33#[derive(Debug)]
34struct NodeSpec {
35    /// Full resolved param values (`amplitude` = initial value, nonlinear = initial).
36    /// Mutated in `set_params` for the nonlinear params.
37    param_values: Vec<f64>,
38    /// For each param slot: `Some(alpha_idx)` if it is a free nonlinear param,
39    /// `None` if it is the amplitude (linear) or a fixed param.
40    alpha_indices: Vec<Option<usize>>,
41    /// Whether the column is invariant (no nonlinear params — constant/linear).
42    is_invariant: bool,
43}
44
45// ---------------------------------------------------------------------------
46// GraphSeparableModel
47// ---------------------------------------------------------------------------
48
49/// Implements varpro's `SeparableNonlinearModel` trait for an arbitrary
50/// separable spectrafit graph.
51pub struct GraphSeparableModel {
52    /// Boxed model kernels (one per node, in graph order).
53    models: Vec<Box<dyn Model>>,
54    /// Per-node parameter bookkeeping.
55    node_specs: Vec<NodeSpec>,
56    /// Flat x values (single 1-D dataset).
57    x: Vec<f64>,
58    /// Number of data points.
59    n_data: usize,
60    /// Current nonlinear parameters α.
61    params: OVector<f64, Dyn>,
62    /// Cached basis matrix Φ (n_data × n_basis).  Recomputed in `set_params`.
63    phi: OMatrix<f64, Dyn, Dyn>,
64    /// Names of alpha keys, for mapping results back.
65    pub alpha_keys: Vec<String>,
66}
67
68impl GraphSeparableModel {
69    /// Build from a compiled graph, one dataset, and the full parameter map.
70    ///
71    /// `all_params` maps `"node_id.param_name"` → [`ParameterSpec`].
72    ///
73    /// # Errors
74    /// Returns [`ModelError::ParameterNotInModel`] when `graph` fails to
75    /// compile, or [`ModelError::IncorrectParameterCount`] if the initial
76    /// `set_params` call (used to prime the Φ cache) receives a parameter
77    /// vector whose length does not match the model's own count.
78    pub fn new(
79        graph: &FitGraphSpec,
80        dataset: &MeasurementSpec,
81        all_params: &HashMap<String, ParameterSpec>,
82    ) -> Result<Self, ModelError> {
83        let compiled =
84            CompiledGraph::compile(graph).map_err(|_| ModelError::ParameterNotInModel {
85                parameter: "graph compilation failed".into(),
86            })?;
87
88        let x: Vec<f64> = dataset.x.first().cloned().unwrap_or_default();
89        let n_data = x.len();
90        let n_basis = compiled.nodes.len();
91
92        // ── 1. Assign alpha indices ────────────────────────────────────────
93        let mut alpha_keys: Vec<String> = Vec::new();
94        let mut node_specs: Vec<NodeSpec> = Vec::new();
95
96        for node_entry in &compiled.nodes {
97            let model_type_str = graph
98                .nodes
99                .iter()
100                .find(|n| n.id == node_entry.id)
101                .map(|n| n.model_type.as_str())
102                .unwrap_or("constant");
103
104            let is_invariant = crate::INVARIANT_MODEL_TYPES.contains(&model_type_str);
105
106            let mut param_values: Vec<f64> = Vec::new();
107            let mut alpha_indices: Vec<Option<usize>> = Vec::new();
108
109            for (i, pname) in node_entry.param_names.iter().enumerate() {
110                let key = format!("{}.{}", node_entry.id, pname);
111                let spec = all_params.get(&key);
112                let value = spec.map(|s| s.value).unwrap_or(0.0);
113                param_values.push(value);
114
115                // i == 0 is amplitude for separable models (linear coeff, not in α).
116                // For invariant models all params are linear — skip all.
117                let is_amplitude = i == 0 && !is_invariant;
118                let is_fixed = spec.map(|s| !s.vary).unwrap_or(true);
119
120                if !is_amplitude && !is_fixed {
121                    let idx = alpha_keys.len();
122                    alpha_keys.push(key);
123                    alpha_indices.push(Some(idx));
124                } else {
125                    alpha_indices.push(None);
126                }
127            }
128
129            node_specs.push(NodeSpec {
130                param_values,
131                alpha_indices,
132                is_invariant,
133            });
134        }
135
136        // ── 2. Build initial α ─────────────────────────────────────────────
137        let init_alpha: Vec<f64> = alpha_keys
138            .iter()
139            .map(|k| all_params.get(k).map(|s| s.value).unwrap_or(0.0))
140            .collect();
141        let params = OVector::<f64, Dyn>::from_vec(init_alpha);
142
143        // ── 3. Allocate Φ ──────────────────────────────────────────────────
144        let phi = OMatrix::<f64, Dyn, Dyn>::zeros(n_data, n_basis);
145
146        // Extract models from compiled (moves out of compiled.nodes)
147        let models: Vec<Box<dyn Model>> = compiled.nodes.into_iter().map(|n| n.model).collect();
148
149        let mut model_obj = Self {
150            models,
151            node_specs,
152            x,
153            n_data,
154            params: params.clone(),
155            phi,
156            alpha_keys,
157        };
158        // Fill cache with initial params (errors bubble up)
159        model_obj.set_params(params)?;
160        Ok(model_obj)
161    }
162
163    /// Compute basis column `j` from current α.
164    fn eval_column(&self, j: usize) -> DVector<f64> {
165        let spec = &self.node_specs[j];
166        let model = &self.models[j];
167
168        let mut pv = spec.param_values.clone();
169        // Inject current α
170        for (i, ai) in spec.alpha_indices.iter().enumerate() {
171            if let Some(idx) = ai {
172                pv[i] = self.params[*idx];
173            }
174        }
175        // Amplitude (index 0) acts as the linear coefficient; set to 1 so the
176        // column is the normalised shape (varpro solves for the amplitude).
177        if !spec.is_invariant && !pv.is_empty() {
178            pv[0] = 1.0;
179        }
180        DVector::from_iterator(self.n_data, self.x.iter().map(|&xi| model.eval(&[xi], &pv)))
181    }
182
183    /// Compute ∂(column j)/∂α[alpha_idx].
184    fn deriv_column(&self, j: usize, alpha_idx: usize) -> DVector<f64> {
185        let spec = &self.node_specs[j];
186        let model = &self.models[j];
187
188        // Which local param position corresponds to alpha_idx?
189        let param_pos = spec
190            .alpha_indices
191            .iter()
192            .position(|ai| *ai == Some(alpha_idx));
193        let Some(param_pos) = param_pos else {
194            return DVector::zeros(self.n_data);
195        };
196
197        let mut pv = spec.param_values.clone();
198        for (i, ai) in spec.alpha_indices.iter().enumerate() {
199            if let Some(idx) = ai {
200                pv[i] = self.params[*idx];
201            }
202        }
203        if !spec.is_invariant && !pv.is_empty() {
204            pv[0] = 1.0;
205        }
206        // model.jacobian returns one derivative per param in the order of param_names
207        DVector::from_iterator(
208            self.n_data,
209            self.x.iter().map(|&xi| {
210                let jac = model.jacobian(&[xi], &pv);
211                jac[param_pos]
212            }),
213        )
214    }
215}
216
217// ---------------------------------------------------------------------------
218// Manual Debug impl (models are Box<dyn Model> which doesn't derive Debug)
219// ---------------------------------------------------------------------------
220
221impl std::fmt::Debug for GraphSeparableModel {
222    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
223        f.debug_struct("GraphSeparableModel")
224            .field("n_data", &self.n_data)
225            .field("n_basis", &self.models.len())
226            .field("n_alpha", &self.params.len())
227            .finish()
228    }
229}
230
231// ---------------------------------------------------------------------------
232// SeparableNonlinearModel impl
233// ---------------------------------------------------------------------------
234
235impl SeparableNonlinearModel for GraphSeparableModel {
236    type ScalarType = f64;
237    type Error = ModelError;
238
239    fn parameter_count(&self) -> usize {
240        self.params.len()
241    }
242
243    fn base_function_count(&self) -> usize {
244        self.models.len()
245    }
246
247    fn output_len(&self) -> usize {
248        self.n_data
249    }
250
251    fn set_params(&mut self, parameters: OVector<f64, Dyn>) -> Result<(), Self::Error> {
252        if parameters.len() != self.params.len() {
253            return Err(ModelError::IncorrectParameterCount {
254                expected: self.params.len(),
255                actual: parameters.len(),
256            });
257        }
258        self.params = parameters;
259        for j in 0..self.models.len() {
260            let col = self.eval_column(j);
261            self.phi.set_column(j, &col);
262        }
263        Ok(())
264    }
265
266    fn params(&self) -> OVector<f64, Dyn> {
267        self.params.clone()
268    }
269
270    fn eval(&self) -> Result<OMatrix<f64, Dyn, Dyn>, Self::Error> {
271        Ok(self.phi.clone())
272    }
273
274    fn eval_partial_deriv(
275        &self,
276        derivative_index: usize,
277    ) -> Result<OMatrix<f64, Dyn, Dyn>, Self::Error> {
278        if derivative_index >= self.params.len() {
279            return Err(ModelError::DerivativeIndexOutOfBounds {
280                index: derivative_index,
281            });
282        }
283        let n_basis = self.models.len();
284        let mut dphi = OMatrix::<f64, Dyn, Dyn>::zeros(self.n_data, n_basis);
285        for j in 0..n_basis {
286            let col = self.deriv_column(j, derivative_index);
287            dphi.set_column(j, &col);
288        }
289        Ok(dphi)
290    }
291}