spectrafit_varpro/
model.rs1use 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#[derive(Debug)]
34struct NodeSpec {
35 param_values: Vec<f64>,
38 alpha_indices: Vec<Option<usize>>,
41 is_invariant: bool,
43}
44
45pub struct GraphSeparableModel {
52 models: Vec<Box<dyn Model>>,
54 node_specs: Vec<NodeSpec>,
56 x: Vec<f64>,
58 n_data: usize,
60 params: OVector<f64, Dyn>,
62 phi: OMatrix<f64, Dyn, Dyn>,
64 pub alpha_keys: Vec<String>,
66}
67
68impl GraphSeparableModel {
69 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 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 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 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 let phi = OMatrix::<f64, Dyn, Dyn>::zeros(n_data, n_basis);
145
146 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 model_obj.set_params(params)?;
160 Ok(model_obj)
161 }
162
163 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 for (i, ai) in spec.alpha_indices.iter().enumerate() {
171 if let Some(idx) = ai {
172 pv[i] = self.params[*idx];
173 }
174 }
175 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 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 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 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
217impl 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
231impl 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}