1#![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#[derive(Copy, Clone, Debug)]
12pub enum Parameter {
13 Free(f32),
15 Fixed(f32),
17}
18
19#[derive(thiserror::Error, Debug)]
21#[error("could not solve for matrix pseudo-inverse: {0}")]
22pub struct SingularMatrix(&'static str);
23
24struct Solver<'a, F: Function> {
26 vars: &'a HashMap<Var, Parameter>,
28
29 grad_tapes: Vec<<F::GradSliceEval as BulkEvaluator>::Tape>,
31
32 point_tapes: Vec<<F::PointEval as TracingEvaluator>::Tape>,
34
35 grad_eval: F::GradSliceEval,
37
38 point_eval: F::PointEval,
40
41 input_grad: Vec<Vec<Grad>>,
43
44 input_point: Vec<f32>,
46
47 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 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 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 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 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 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 let out = self.grad_eval.eval(tape, &self.input_grad).unwrap();
139
140 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 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 let (out, _t) =
172 self.point_eval.eval(tape, &self.input_point).unwrap();
173 err += out[0].powi(2); }
175 err
176 }
177}
178
179pub 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 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 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 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 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 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 damping *= 1.5;
254 } else {
255 damping /= 3.0;
257 break (err, delta);
258 }
259 };
260
261 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 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 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 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 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, ¶ms).unwrap();
498
499 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 let mut values = nalgebra::DVector::<f32>::zeros(n);
538 for v in values.iter_mut() {
539 *v = rand::random();
540 }
541
542 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, ¶ms).unwrap();
584
585 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 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}