Skip to main content

fidget_solver/
lib.rs

1//! Solver for systems of equations expressed as sets of [Function] objects
2#![warn(missing_docs)]
3use fidget_core::{
4    eval::{BulkEvaluator, Function, Tape, TracingEvaluator},
5    types::Grad,
6    var::Var,
7};
8use std::collections::HashMap;
9
10/// Input parameter to the solver
11#[derive(Copy, Clone, Debug)]
12pub enum Parameter {
13    /// Free variable with the given starting position
14    Free(f32),
15    /// Fixed variable at the given value
16    Fixed(f32),
17}
18
19/// Error indicating that we could not solve for matrix pseudo-inverse
20#[derive(thiserror::Error, Debug)]
21#[error("could not solve for matrix pseudo-inverse: {0}")]
22pub struct SingularMatrix(&'static str);
23
24/// Workspace for solvers
25struct Solver<'a, F: Function> {
26    /// Input parameters
27    vars: &'a HashMap<Var, Parameter>,
28
29    /// Tapes for bulk gradient evaluation of each constraint
30    grad_tapes: Vec<<F::GradSliceEval as BulkEvaluator>::Tape>,
31
32    /// Tapes for single-point evaluation of each constraint
33    point_tapes: Vec<<F::PointEval as TracingEvaluator>::Tape>,
34
35    /// Bulk gradient evaluator, for use in computing the Jacobian
36    grad_eval: F::GradSliceEval,
37
38    /// Single-point evaluator, for use in checking our current error
39    point_eval: F::PointEval,
40
41    /// Input data for use when calling the gradient bulk evaluator
42    input_grad: Vec<Vec<Grad>>,
43
44    /// Input data for use when calling the single-point evaluator
45    input_point: Vec<f32>,
46
47    /// Map from (free) variables to the index of their gradient
48    ///
49    /// We evaluate 3x gradients per sample, so for `grad_index = gi`, the
50    /// relevant derivative will be `out[gi / 3].d(gi % 3)`
51    grad_index: HashMap<Var, usize>,
52}
53
54impl<'a, F: Function> Solver<'a, F> {
55    fn new(eqs: &'a [F], vars: &'a HashMap<Var, Parameter>) -> Self {
56        // Build our per-constraint
57        let grad_tapes = eqs
58            .iter()
59            .map(|f| f.grad_slice_tape(Default::default()))
60            .collect::<Vec<_>>();
61        let point_tapes = eqs
62            .iter()
63            .map(|f| f.point_tape(Default::default()))
64            .collect::<Vec<_>>();
65
66        // Build a map from *free* variable to index of its gradient, since
67        // we'll be using tightly-packed Vec everywhere here
68        //
69        // (We ignore the gradient of fixed variables)
70        let grad_index: HashMap<Var, usize> = vars
71            .iter()
72            .filter(|(_v, p)| matches!(p, Parameter::Free(..)))
73            .enumerate()
74            .map(|(i, (v, _p))| (*v, i))
75            .collect();
76
77        let var_count = vars
78            .len()
79            .max(grad_tapes.iter().map(|t| t.vars().len()).max().unwrap_or(0));
80
81        // Build a scratch array with rows for each variable, and enough columns
82        // to simultaneously compute all of the gradients that we need
83        let input_grad =
84            vec![
85                vec![Grad::from(0f32); grad_index.len().div_ceil(3)];
86                var_count
87            ];
88        let input_point = vec![0f32; var_count];
89
90        Self {
91            vars,
92            grad_tapes,
93            point_tapes,
94            grad_eval: Default::default(),
95            point_eval: Default::default(),
96            grad_index,
97
98            input_grad,
99            input_point,
100        }
101    }
102
103    /// Computes the Jacobian into `cur`
104    ///
105    /// # Panics
106    /// If `jacobian` or `result` are an invalid size
107    fn get_jacobian(
108        &mut self,
109        cur: &[f32],
110        jacobian: &mut nalgebra::DMatrix<f32>,
111        result: &mut nalgebra::DVector<f32>,
112    ) {
113        for (ti, tape) in self.grad_tapes.iter().enumerate() {
114            // Update the values in the gradient evaluation array
115            for (v, p) in self.vars {
116                let Some(i) = tape.vars().get(v) else {
117                    continue;
118                };
119                let slice = &mut self.input_grad[i];
120                match p {
121                    Parameter::Free(..) => {
122                        let gi = self.grad_index[v];
123                        for (j, v) in slice.iter_mut().enumerate() {
124                            *v = Grad::new(
125                                cur[gi],
126                                if j * 3 == gi { 1.0 } else { 0.0 },
127                                if j * 3 + 1 == gi { 1.0 } else { 0.0 },
128                                if j * 3 + 2 == gi { 1.0 } else { 0.0 },
129                            );
130                        }
131                    }
132                    Parameter::Fixed(f) => {
133                        slice.fill(Grad::new(*f, 0.0, 0.0, 0.0));
134                    }
135                };
136            }
137            // Do the actual gradient evaluation
138            let out = self.grad_eval.eval(tape, &self.input_grad).unwrap();
139
140            // Populate this row of the Jacobian
141            for gi in 0..self.grad_index.len() {
142                *jacobian.get_mut((ti, gi)).unwrap() = out[0][gi / 3].d(gi % 3);
143            }
144            result[ti] = out[0][0].v;
145        }
146    }
147
148    fn get_err(&mut self, cur: &[f32], delta: &[f32]) -> f32 {
149        let mut err = 0f32;
150        for tape in self.point_tapes.iter() {
151            // Update the free values in the gradient evaluation array
152            //
153            // (we preloaded unit gradients and fixed values in the appropriate
154            // locations, which don't change from evaluation to evaluation)
155            for (v, p) in self.vars {
156                let Some(i) = tape.vars().get(v) else {
157                    continue;
158                };
159                let f = &mut self.input_point[i];
160                match p {
161                    Parameter::Free(..) => {
162                        let gi = self.grad_index[v];
163                        *f = cur[gi] - delta[gi];
164                    }
165                    Parameter::Fixed(p) => {
166                        *f = *p;
167                    }
168                };
169            }
170            // Do the actual gradient evaluation
171            let (out, _t) =
172                self.point_eval.eval(tape, &self.input_point).unwrap();
173            err += out[0].powi(2); // TODO: consolidate into a single tape
174        }
175        err
176    }
177}
178
179/// Least-squares minimization on a set of functions
180///
181/// Returns a map from free variable to its final value
182///
183/// Minimization is accomplished using a relatively basic implementation of
184/// the [Levenberg-Marquardt algorithm](https://en.wikipedia.org/wiki/Levenberg%E2%80%93Marquardt_algorithm).
185///
186/// ## References
187/// - [The Levenberg-Marquardt Algorithm (Ranganathan 2004)](http://ananth.in/docs/lmtut.pdf)
188/// - [Basics on Continuous Optimization ยง Levenberg-Marquardt](https://www.brnt.eu/phd/node10.html#SECTION00622700000000000000)
189/// - [Improvements to the Levenberg-Marquardt algorithm for nonlinear
190///   least-squares minimization (Transtrum 2012)](https://arxiv.org/pdf/1201.5885)
191pub fn solve<F: Function>(
192    eqs: &[F],
193    vars: &HashMap<Var, Parameter>,
194) -> Result<HashMap<Var, f32>, SingularMatrix> {
195    let tapes = eqs
196        .iter()
197        .map(|f| f.grad_slice_tape(Default::default()))
198        .collect::<Vec<_>>();
199
200    // Current values for free variables
201    let mut cur = HashMap::new();
202    for (v, p) in vars {
203        if let Parameter::Free(f) = *p {
204            cur.insert(*v, f);
205        }
206    }
207
208    let mut solver = Solver::new(eqs, vars);
209
210    // Build an array of current values for each free variable
211    let mut cur = vec![0f32; solver.grad_index.len()];
212    for (v, i) in &solver.grad_index {
213        let Parameter::Free(f) = vars[v] else {
214            unreachable!();
215        };
216        cur[*i] = f;
217    }
218
219    // Working arrays for the current Jacobian and result
220    let mut jacobian = nalgebra::DMatrix::repeat(tapes.len(), cur.len(), 0f32);
221    let mut result = nalgebra::DVector::repeat(tapes.len(), 0f32);
222
223    let mut damping = 1.0;
224    let mut prev_err = f32::INFINITY;
225    let mut err_buf = [0f32; 4];
226    for i in 0.. {
227        solver.get_jacobian(&cur, &mut jacobian, &mut result);
228
229        // Early exit if we're done
230        if result.iter().all(|v| *v == 0.0) {
231            break;
232        }
233
234        let jt = jacobian.transpose();
235        let jt_j = &jt * &jacobian;
236
237        let jt_r = jt * &result;
238
239        // TODO: be optimistic and evaluate the full gradient on the first
240        // attempt, since it should usually succeed?
241        let (err, step) = loop {
242            let adjusted = &jt_j
243                + damping * nalgebra::DMatrix::from_diagonal(&jt_j.diagonal());
244
245            let delta = adjusted
246                .svd(true, true)
247                .solve(&jt_r, f32::EPSILON)
248                .map_err(SingularMatrix)?;
249
250            let err = solver.get_err(&cur, delta.as_slice());
251            if err > prev_err {
252                // Keep going in this inner loop, taking smaller steps
253                damping *= 1.5;
254            } else {
255                // We found a good step size, so reduce damping
256                damping /= 3.0;
257                break (err, delta);
258            }
259        };
260
261        // Update our current position, checking whether it actually changed
262        // (i.e. whether our steps are below the floating-point epsilon)
263        //
264        // TODO: improve exit criteria?
265        let mut changed = false;
266        for gi in 0..solver.grad_index.len() {
267            let prev = cur[gi];
268            cur[gi] -= step[gi];
269            changed |= prev != cur[gi];
270        }
271        err_buf[i % err_buf.len()] = err;
272        if !changed
273            || err == 0.0
274            || damping == 0.0
275            || err_buf.iter().all(|e| *e == err_buf[0])
276        {
277            break;
278        }
279        prev_err = err;
280    }
281
282    // Return the new "current" values, which are our optimized position
283    let out = solver
284        .grad_index
285        .into_iter()
286        .map(|(v, i)| (v, cur[i]))
287        .collect();
288    Ok(out)
289}
290
291#[cfg(test)]
292mod test {
293    use super::*;
294    use approx::{assert_relative_eq, relative_eq};
295    use fidget_core::{
296        context::{Context, Tree},
297        eval::MathFunction,
298        vm::VmFunction,
299    };
300
301    #[test]
302    fn basic_solver() {
303        let eqn = Tree::x() + Tree::y();
304        let mut ctx = Context::new();
305        let root = ctx.import(&eqn);
306
307        let f = VmFunction::new(&ctx, &[root]).unwrap();
308        let mut values = HashMap::new();
309        values.insert(Var::X, Parameter::Free(0.0));
310        values.insert(Var::Y, Parameter::Fixed(-1.0));
311        let sol = solve(&[f], &values).unwrap();
312        assert_eq!(sol.len(), 1);
313        assert_relative_eq!(sol[&Var::X], 1.0);
314    }
315
316    #[test]
317    fn four_vars_at_once() {
318        let vs = (0..4).map(|_| Var::new()).collect::<Vec<Var>>();
319        let mut root = Tree::from(vs[0]);
320        for v in &vs[1..] {
321            root += Tree::from(*v);
322        }
323        let mut ctx = Context::new();
324        let root = ctx.import(&root);
325
326        let f = VmFunction::new(&ctx, &[root]).unwrap();
327        let mut values = HashMap::new();
328        for (i, &v) in vs.iter().enumerate() {
329            values.insert(v, Parameter::Free(i as f32));
330        }
331        let sol = solve(&[f], &values).unwrap();
332        assert_eq!(sol.len(), 4);
333        let mut out = 0.0;
334        for v in &vs {
335            out += sol[v];
336        }
337        assert_relative_eq!(out, 0.0);
338    }
339
340    #[test]
341    fn four_vars_independent() {
342        let vs = (0..4).map(|_| Var::new()).collect::<Vec<Var>>();
343        let mut eqns = vec![];
344        let mut ctx = Context::new();
345        for (i, &v) in vs.iter().enumerate() {
346            let eqn = Tree::from(v) - Tree::from(i as f32);
347            let root = ctx.import(&eqn);
348            let f = VmFunction::new(&ctx, &[root]).unwrap();
349            eqns.push(f);
350        }
351
352        let mut values = HashMap::new();
353        for (i, &v) in vs.iter().enumerate() {
354            values.insert(v, Parameter::Free(i as f32 * 2.0));
355        }
356        let sol = solve(&eqns, &values).unwrap();
357        assert_eq!(sol.len(), 4);
358        for (i, v) in vs.iter().enumerate() {
359            assert_relative_eq!(i as f32, sol[v]);
360        }
361    }
362
363    #[test]
364    fn xy_nonlinear() {
365        let constraints = vec![
366            (Tree::x() * 2 + Tree::y() * 3) * (Tree::x() - Tree::y()) - 2,
367            Tree::x() * 3 + Tree::y() - 5,
368        ];
369        let mut ctx = Context::new();
370        let eqns = constraints
371            .into_iter()
372            .map(|c| {
373                let root = ctx.import(&c);
374                VmFunction::new(&ctx, &[root]).unwrap()
375            })
376            .collect::<Vec<_>>();
377
378        let mut values = HashMap::new();
379        values.insert(Var::X, Parameter::Free(0.0));
380        values.insert(Var::Y, Parameter::Free(0.0));
381        let sol = solve(&eqns, &values).unwrap();
382
383        let x = sol[&Var::X];
384        let y = sol[&Var::Y];
385
386        assert_relative_eq!((x * 2.0 + y * 3.0) * (x - y), 2.0);
387        assert_relative_eq!(x * 3.0 + y, 5.0);
388    }
389
390    #[test]
391    fn one_var_no_solution() {
392        // Solve for X == 1 and X == 2 simultaneously
393        let constraints = vec![Tree::x() - 1.0, Tree::x() - 2.0];
394
395        let mut ctx = Context::new();
396        let eqns = constraints
397            .into_iter()
398            .map(|c| {
399                let root = ctx.import(&c);
400                VmFunction::new(&ctx, &[root]).unwrap()
401            })
402            .collect::<Vec<_>>();
403
404        let mut values = HashMap::new();
405        values.insert(Var::X, Parameter::Free(0.0));
406
407        let sol = solve(&eqns, &values).unwrap();
408
409        let x = sol[&Var::X];
410        assert_relative_eq!(x, 1.5);
411    }
412
413    #[test]
414    fn solve_banana() {
415        // See https://en.wikipedia.org/wiki/Rosenbrock_function
416        let a = 1f32;
417        let b = 100f32;
418        let constraints = [a - Tree::x(), b * (Tree::y() - Tree::x().square())];
419
420        let mut ctx = Context::new();
421        let eqns = constraints
422            .into_iter()
423            .map(|c| {
424                let root = ctx.import(&c);
425                VmFunction::new(&ctx, &[root]).unwrap()
426            })
427            .collect::<Vec<_>>();
428
429        let mut values = HashMap::new();
430        values.insert(Var::X, Parameter::Free(0.0));
431        values.insert(Var::Y, Parameter::Free(0.0));
432        let sol = solve(&eqns, &values).unwrap();
433        assert_relative_eq!(sol[&Var::X], 1.0);
434        assert_relative_eq!(sol[&Var::Y], 1.0);
435
436        let mut values = HashMap::new();
437        values.insert(Var::X, Parameter::Free(1.0));
438        values.insert(Var::Y, Parameter::Free(1.0));
439        let sol = solve(&eqns, &values).unwrap();
440        assert_relative_eq!(sol[&Var::X], 1.0);
441        assert_relative_eq!(sol[&Var::Y], 1.0);
442    }
443
444    #[test]
445    fn solve_circle() {
446        let t = (Tree::x().square() + Tree::y().square()).sqrt();
447        let mut ctx = Context::new();
448        let root = ctx.import(&t);
449        let eqn = VmFunction::new(&ctx, &[root]).unwrap();
450        let eqns = [eqn];
451
452        let mut values = HashMap::new();
453        values.insert(Var::X, Parameter::Free(0.0));
454        values.insert(Var::Y, Parameter::Free(0.0));
455        let sol = solve(&eqns, &values).unwrap();
456        assert_relative_eq!(sol[&Var::X], 0.0);
457        assert_relative_eq!(sol[&Var::Y], 0.0);
458
459        let mut values = HashMap::new();
460        values.insert(Var::X, Parameter::Free(1.0));
461        values.insert(Var::Y, Parameter::Free(1.5));
462        let sol = solve(&eqns, &values).unwrap();
463        assert_relative_eq!(sol[&Var::X], 0.0);
464        assert_relative_eq!(sol[&Var::Y], 0.0);
465    }
466
467    fn one_linear(n: usize) {
468        // Build a random matrix of our solutions
469        let mut values = nalgebra::DVector::<f32>::zeros(n);
470        for v in values.iter_mut() {
471            *v = rand::random();
472        }
473
474        let vars = (0..n).map(|_| Var::new()).collect::<Vec<_>>();
475        let trees = vars.iter().map(|v| Tree::from(*v)).collect::<Vec<_>>();
476
477        let mut mat = nalgebra::DMatrix::<f32>::zeros(n, n);
478        for v in mat.iter_mut() {
479            *v = rand::random();
480        }
481
482        let sol = &mat * &values;
483
484        let mut ctx = Context::new();
485        let mut eqns = vec![];
486        for row in 0..n {
487            let mut out = Tree::from(-sol[row]);
488            for (col, t) in trees.iter().enumerate() {
489                out += *mat.get((row, col)).unwrap() * t.clone();
490            }
491            let root = ctx.import(&out);
492            let f = VmFunction::new(&ctx, &[root]).unwrap();
493            eqns.push(f);
494        }
495
496        let params = vars.iter().map(|v| (*v, Parameter::Free(0.0))).collect();
497        let out = solve(&eqns, &params).unwrap();
498
499        // It's possible for there to be multiple solutions here, so we'll check
500        // the actual equations.
501        for i in 0..n {
502            values[i] = out[&vars[i]];
503        }
504        let sol2 = &mat * &values;
505        let err = (&sol - &sol2).norm_squared();
506        assert!(err < 1e-3, "error {err} is too large");
507        for (a, b) in sol.iter().zip(sol2.iter()) {
508            assert_relative_eq!(a, b, epsilon = 1e-2);
509        }
510    }
511
512    #[test]
513    fn small_linear() {
514        for _ in 0..1000 {
515            one_linear(2);
516        }
517    }
518
519    #[test]
520    fn medium_linear() {
521        for _ in 0..1000 {
522            one_linear(10);
523        }
524    }
525
526    #[test]
527    fn big_linear() {
528        for _ in 0..50 {
529            one_linear(50);
530        }
531    }
532
533    fn one_quadratic(n: usize) -> bool {
534        let m: usize = n * n + n;
535
536        // Build a random matrix of our solutions
537        let mut values = nalgebra::DVector::<f32>::zeros(n);
538        for v in values.iter_mut() {
539            *v = rand::random();
540        }
541
542        // Build a column vector of [a b c ... aa ab ac ... ba bb ...]^T
543        let mut col = nalgebra::DVector::<f32>::zeros(m);
544        col.rows_range_mut(..n).copy_from(&values);
545        for i in 0..n {
546            for j in 0..n {
547                let index = i * n + j + n;
548                col[index] = values[i] * values[j];
549            }
550        }
551
552        let vars = (0..n).map(|_| Var::new()).collect::<Vec<_>>();
553        let trees = vars.iter().map(|v| Tree::from(*v)).collect::<Vec<_>>();
554
555        let mut mat = nalgebra::DMatrix::<f32>::zeros(n, m);
556        for v in mat.iter_mut() {
557            *v = rand::random();
558        }
559
560        let sol = &mat * &col;
561
562        let mut ctx = Context::new();
563        let mut eqns = vec![];
564        for row in 0..n {
565            let mut out = Tree::from(-sol[row]);
566            for (col, t) in trees.iter().enumerate() {
567                out += *mat.get((row, col)).unwrap() * t.clone();
568            }
569            for i in 0..n {
570                for j in 0..n {
571                    let index = i * n + j + n;
572                    out += *mat.get((row, index)).unwrap()
573                        * trees[i].clone()
574                        * trees[j].clone();
575                }
576            }
577            let root = ctx.import(&out);
578            let f = VmFunction::new(&ctx, &[root]).unwrap();
579            eqns.push(f);
580        }
581
582        let params = vars.iter().map(|v| (*v, Parameter::Free(0.5))).collect();
583        let out = solve(&eqns, &params).unwrap();
584
585        // It's possible for there to be multiple solutions here, so we'll check
586        // the actual equations.
587        for i in 0..n {
588            col[i] = out[&vars[i]];
589            for j in 0..n {
590                let index = i * n + j + n;
591                col[index] = out[&vars[i]] * out[&vars[j]];
592            }
593        }
594        let sol2 = &mat * &col;
595        let err = (&sol - &sol2).norm_squared();
596        if err >= 1e-3 {
597            return false;
598        }
599        for (a, b) in sol.iter().zip(sol2.iter()) {
600            if !relative_eq!(a, b, epsilon = 1e-2) {
601                return false;
602            }
603        }
604        true
605    }
606
607    // Quadratic functions can get trapped in local minima, so we only require a
608    // certain percent to succeed (hard-coded to 90% right now)
609    fn many_quadratic(size: usize, count: usize) {
610        let mut okay = 0;
611        for _ in 0..count {
612            if one_quadratic(size) {
613                okay += 1;
614            }
615        }
616        assert!(
617            okay >= count * 9 / 10,
618            "too many failures: {okay} / {count}"
619        );
620    }
621
622    #[test]
623    fn small_quadratic() {
624        many_quadratic(2, 1000);
625    }
626
627    #[test]
628    fn medium_quadratic() {
629        many_quadratic(5, 100);
630    }
631
632    #[test]
633    fn large_quadratic() {
634        many_quadratic(10, 50);
635    }
636}