1use core::fmt;
10
11use serde::{Deserialize, Serialize};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
15#[serde(transparent)]
16pub struct VarId(pub u32);
17
18impl VarId {
19 #[must_use]
21 pub const fn index(self) -> usize {
22 self.0 as usize
23 }
24}
25
26#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28#[serde(tag = "op", content = "args", rename_all = "camelCase")]
29pub enum Expr {
30 Const(f64),
32 Var(VarId),
34 Add(Box<Expr>, Box<Expr>),
36 Sub(Box<Expr>, Box<Expr>),
38 Mul(Box<Expr>, Box<Expr>),
40 Div(Box<Expr>, Box<Expr>),
42 Neg(Box<Expr>),
44 Sin(Box<Expr>),
46 Cos(Box<Expr>),
48 Sqrt(Box<Expr>),
50 Abs(Box<Expr>),
52 Atan2(Box<Expr>, Box<Expr>),
54 Hypot(Box<Expr>, Box<Expr>),
56 Min(Box<Expr>, Box<Expr>),
58 Max(Box<Expr>, Box<Expr>),
60}
61
62#[derive(Debug, Clone, PartialEq, Default)]
64pub struct Dual {
65 pub v: f64,
67 pub g: Vec<(u32, f64)>,
69}
70
71fn merge(a: &[(u32, f64)], ka: f64, b: &[(u32, f64)], kb: f64) -> Vec<(u32, f64)> {
72 let mut out = Vec::with_capacity(a.len() + b.len());
73 let (mut i, mut j) = (0, 0);
74 while i < a.len() || j < b.len() {
75 match (a.get(i), b.get(j)) {
76 (Some(&(va, da)), Some(&(vb, db))) if va == vb => {
77 out.push((va, da * ka + db * kb));
78 i += 1;
79 j += 1;
80 }
81 (Some(&(va, da)), Some(&(vb, _))) if va < vb => {
82 out.push((va, da * ka));
83 i += 1;
84 }
85 (Some(_), Some(&(vb, db))) => {
86 out.push((vb, db * kb));
87 j += 1;
88 }
89 (Some(&(va, da)), None) => {
90 out.push((va, da * ka));
91 i += 1;
92 }
93 (None, Some(&(vb, db))) => {
94 out.push((vb, db * kb));
95 j += 1;
96 }
97 (None, None) => break,
98 }
99 }
100 out
101}
102
103fn scale(a: &[(u32, f64)], k: f64) -> Vec<(u32, f64)> {
104 a.iter().map(|&(v, d)| (v, d * k)).collect()
105}
106
107#[allow(clippy::should_implement_trait, clippy::redundant_guards)]
110impl Expr {
111 #[must_use]
113 pub const fn c(v: f64) -> Self {
114 Self::Const(v)
115 }
116
117 #[must_use]
119 pub const fn var(v: VarId) -> Self {
120 Self::Var(v)
121 }
122
123 fn as_const(&self) -> Option<f64> {
124 if let Self::Const(v) = self { Some(*v) } else { None }
125 }
126
127 #[must_use]
129 pub fn add(a: Self, b: Self) -> Self {
130 match (a.as_const(), b.as_const()) {
131 (Some(x), Some(y)) => Self::Const(x + y),
132 (Some(x), None) if x == 0.0 => b,
133 (None, Some(y)) if y == 0.0 => a,
134 _ => Self::Add(Box::new(a), Box::new(b)),
135 }
136 }
137
138 #[must_use]
140 pub fn sub(a: Self, b: Self) -> Self {
141 match (a.as_const(), b.as_const()) {
142 (Some(x), Some(y)) => Self::Const(x - y),
143 (None, Some(y)) if y == 0.0 => a,
144 (Some(x), None) if x == 0.0 => Self::neg(b),
145 _ => Self::Sub(Box::new(a), Box::new(b)),
146 }
147 }
148
149 #[must_use]
151 pub fn mul(a: Self, b: Self) -> Self {
152 match (a.as_const(), b.as_const()) {
153 (Some(x), Some(y)) => Self::Const(x * y),
154 (Some(x), _) | (_, Some(x)) if x == 0.0 => Self::Const(0.0),
155 (Some(x), None) if x == 1.0 => b,
156 (None, Some(y)) if y == 1.0 => a,
157 _ => Self::Mul(Box::new(a), Box::new(b)),
158 }
159 }
160
161 #[must_use]
163 pub fn div(a: Self, b: Self) -> Self {
164 match (a.as_const(), b.as_const()) {
165 (Some(x), Some(y)) => Self::Const(x / y),
166 (None, Some(y)) if y == 1.0 => a,
167 _ => Self::Div(Box::new(a), Box::new(b)),
168 }
169 }
170
171 #[must_use]
173 pub fn neg(a: Self) -> Self {
174 match a {
175 Self::Const(x) => Self::Const(-x),
176 Self::Neg(inner) => *inner,
177 other => Self::Neg(Box::new(other)),
178 }
179 }
180
181 #[must_use]
183 pub fn sin(a: Self) -> Self {
184 a.as_const().map_or_else(|| Self::Sin(Box::new(a)), |x| Self::Const(libm::sin(x)))
185 }
186
187 #[must_use]
189 pub fn cos(a: Self) -> Self {
190 a.as_const().map_or_else(|| Self::Cos(Box::new(a)), |x| Self::Const(libm::cos(x)))
191 }
192
193 #[must_use]
195 pub fn sqrt(a: Self) -> Self {
196 a.as_const().map_or_else(|| Self::Sqrt(Box::new(a)), |x| Self::Const(x.sqrt()))
197 }
198
199 #[must_use]
201 pub fn abs(a: Self) -> Self {
202 a.as_const().map_or_else(|| Self::Abs(Box::new(a)), |x| Self::Const(x.abs()))
203 }
204
205 #[must_use]
207 pub fn atan2(y: Self, x: Self) -> Self {
208 match (y.as_const(), x.as_const()) {
209 (Some(a), Some(b)) => Self::Const(libm::atan2(a, b)),
210 _ => Self::Atan2(Box::new(y), Box::new(x)),
211 }
212 }
213
214 #[must_use]
216 pub fn hypot(x: Self, y: Self) -> Self {
217 match (x.as_const(), y.as_const()) {
218 (Some(a), Some(b)) => Self::Const(libm::hypot(a, b)),
219 _ => Self::Hypot(Box::new(x), Box::new(y)),
220 }
221 }
222
223 #[must_use]
225 pub fn min(a: Self, b: Self) -> Self {
226 match (a.as_const(), b.as_const()) {
227 (Some(x), Some(y)) => Self::Const(x.min(y)),
228 _ => Self::Min(Box::new(a), Box::new(b)),
229 }
230 }
231
232 #[must_use]
234 pub fn max(a: Self, b: Self) -> Self {
235 match (a.as_const(), b.as_const()) {
236 (Some(x), Some(y)) => Self::Const(x.max(y)),
237 _ => Self::Max(Box::new(a), Box::new(b)),
238 }
239 }
240
241 #[must_use]
243 pub fn eval(&self, x: &[f64]) -> f64 {
244 match self {
245 Self::Const(v) => *v,
246 Self::Var(v) => x.get(v.index()).copied().unwrap_or(f64::NAN),
247 Self::Add(a, b) => a.eval(x) + b.eval(x),
248 Self::Sub(a, b) => a.eval(x) - b.eval(x),
249 Self::Mul(a, b) => a.eval(x) * b.eval(x),
250 Self::Div(a, b) => a.eval(x) / b.eval(x),
251 Self::Neg(a) => -a.eval(x),
252 Self::Sin(a) => libm::sin(a.eval(x)),
253 Self::Cos(a) => libm::cos(a.eval(x)),
254 Self::Sqrt(a) => a.eval(x).sqrt(),
255 Self::Abs(a) => a.eval(x).abs(),
256 Self::Atan2(y, xx) => libm::atan2(y.eval(x), xx.eval(x)),
257 Self::Hypot(a, b) => libm::hypot(a.eval(x), b.eval(x)),
258 Self::Min(a, b) => a.eval(x).min(b.eval(x)),
259 Self::Max(a, b) => a.eval(x).max(b.eval(x)),
260 }
261 }
262
263 #[must_use]
268 pub fn eval_dual(&self, x: &[f64]) -> Dual {
269 match self {
270 Self::Const(v) => Dual { v: *v, g: Vec::new() },
271 Self::Var(v) => Dual { v: x.get(v.index()).copied().unwrap_or(f64::NAN), g: vec![(v.0, 1.0)] },
272 Self::Add(a, b) => {
273 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
274 Dual { v: a.v + b.v, g: merge(&a.g, 1.0, &b.g, 1.0) }
275 }
276 Self::Sub(a, b) => {
277 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
278 Dual { v: a.v - b.v, g: merge(&a.g, 1.0, &b.g, -1.0) }
279 }
280 Self::Mul(a, b) => {
281 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
282 Dual { v: a.v * b.v, g: merge(&a.g, b.v, &b.g, a.v) }
283 }
284 Self::Div(a, b) => {
285 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
286 let inv = 1.0 / b.v;
287 Dual { v: a.v * inv, g: merge(&a.g, inv, &b.g, -a.v * inv * inv) }
288 }
289 Self::Neg(a) => {
290 let a = a.eval_dual(x);
291 Dual { v: -a.v, g: scale(&a.g, -1.0) }
292 }
293 Self::Sin(a) => {
294 let a = a.eval_dual(x);
295 Dual { v: libm::sin(a.v), g: scale(&a.g, libm::cos(a.v)) }
296 }
297 Self::Cos(a) => {
298 let a = a.eval_dual(x);
299 Dual { v: libm::cos(a.v), g: scale(&a.g, -libm::sin(a.v)) }
300 }
301 Self::Sqrt(a) => {
302 let a = a.eval_dual(x);
303 let v = a.v.sqrt();
304 let k = if v > 0.0 { 0.5 / v } else { 0.0 };
305 Dual { v, g: scale(&a.g, k) }
306 }
307 Self::Abs(a) => {
308 let a = a.eval_dual(x);
309 let s = if a.v > 0.0 {
310 1.0
311 } else if a.v < 0.0 {
312 -1.0
313 } else {
314 0.0
315 };
316 Dual { v: a.v.abs(), g: scale(&a.g, s) }
317 }
318 Self::Atan2(y, xx) => {
319 let (y, xx) = (y.eval_dual(x), xx.eval_dual(x));
320 let r2 = y.v * y.v + xx.v * xx.v;
321 if r2 > 0.0 {
322 Dual { v: libm::atan2(y.v, xx.v), g: merge(&y.g, xx.v / r2, &xx.g, -y.v / r2) }
323 } else {
324 Dual { v: 0.0, g: Vec::new() }
325 }
326 }
327 Self::Hypot(a, b) => {
328 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
329 let v = libm::hypot(a.v, b.v);
330 if v > 0.0 { Dual { v, g: merge(&a.g, a.v / v, &b.g, b.v / v) } } else { Dual { v, g: a.g } }
331 }
332 Self::Min(a, b) => {
333 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
334 if a.v <= b.v { a } else { b }
335 }
336 Self::Max(a, b) => {
337 let (a, b) = (a.eval_dual(x), b.eval_dual(x));
338 if a.v >= b.v { a } else { b }
339 }
340 }
341 }
342
343 #[must_use]
345 pub fn vars(&self) -> Vec<VarId> {
346 let mut out = Vec::new();
347 self.collect_vars(&mut out);
348 out.sort_unstable();
349 out.dedup();
350 out
351 }
352
353 fn collect_vars(&self, out: &mut Vec<VarId>) {
354 match self {
355 Self::Const(_) => {}
356 Self::Var(v) => out.push(*v),
357 Self::Neg(a) | Self::Sin(a) | Self::Cos(a) | Self::Sqrt(a) | Self::Abs(a) => a.collect_vars(out),
358 Self::Add(a, b)
359 | Self::Sub(a, b)
360 | Self::Mul(a, b)
361 | Self::Div(a, b)
362 | Self::Atan2(a, b)
363 | Self::Hypot(a, b)
364 | Self::Min(a, b)
365 | Self::Max(a, b) => {
366 a.collect_vars(out);
367 b.collect_vars(out);
368 }
369 }
370 }
371
372 #[must_use]
375 pub fn linear_form(&self, x: &[f64], fixed: &dyn Fn(VarId) -> bool) -> Option<LinearForm> {
376 match self {
377 Self::Const(v) => Some(LinearForm { terms: Vec::new(), constant: *v }),
378 Self::Var(v) => {
379 if fixed(*v) {
380 Some(LinearForm { terms: Vec::new(), constant: x.get(v.index()).copied().unwrap_or(f64::NAN) })
381 } else {
382 Some(LinearForm { terms: vec![(*v, 1.0)], constant: 0.0 })
383 }
384 }
385 Self::Add(a, b) => Some(a.linear_form(x, fixed)?.combine(&b.linear_form(x, fixed)?, 1.0)),
386 Self::Sub(a, b) => Some(a.linear_form(x, fixed)?.combine(&b.linear_form(x, fixed)?, -1.0)),
387 Self::Neg(a) => Some(a.linear_form(x, fixed)?.scaled(-1.0)),
388 Self::Mul(a, b) => {
389 let (la, lb) = (a.linear_form(x, fixed)?, b.linear_form(x, fixed)?);
390 if la.terms.is_empty() {
391 Some(lb.scaled(la.constant))
392 } else if lb.terms.is_empty() {
393 Some(la.scaled(lb.constant))
394 } else {
395 None
396 }
397 }
398 Self::Div(a, b) => {
399 let (la, lb) = (a.linear_form(x, fixed)?, b.linear_form(x, fixed)?);
400 (lb.terms.is_empty() && lb.constant != 0.0).then(|| la.scaled(1.0 / lb.constant))
401 }
402 other => {
404 let vars = other.vars();
405 if vars.iter().all(|v| fixed(*v)) {
406 Some(LinearForm { terms: Vec::new(), constant: other.eval(x) })
407 } else {
408 None
409 }
410 }
411 }
412 }
413
414 #[must_use]
417 pub fn substitute(&self, f: &dyn Fn(VarId) -> Self) -> Self {
418 match self {
419 Self::Const(v) => Self::Const(*v),
420 Self::Var(v) => f(*v),
421 Self::Add(a, b) => Self::add(a.substitute(f), b.substitute(f)),
422 Self::Sub(a, b) => Self::sub(a.substitute(f), b.substitute(f)),
423 Self::Mul(a, b) => Self::mul(a.substitute(f), b.substitute(f)),
424 Self::Div(a, b) => Self::div(a.substitute(f), b.substitute(f)),
425 Self::Neg(a) => Self::neg(a.substitute(f)),
426 Self::Sin(a) => Self::sin(a.substitute(f)),
427 Self::Cos(a) => Self::cos(a.substitute(f)),
428 Self::Sqrt(a) => Self::sqrt(a.substitute(f)),
429 Self::Abs(a) => Self::abs(a.substitute(f)),
430 Self::Atan2(a, b) => Self::atan2(a.substitute(f), b.substitute(f)),
431 Self::Hypot(a, b) => Self::hypot(a.substitute(f), b.substitute(f)),
432 Self::Min(a, b) => Self::min(a.substitute(f), b.substitute(f)),
433 Self::Max(a, b) => Self::max(a.substitute(f), b.substitute(f)),
434 }
435 }
436
437 #[must_use]
439 pub fn size(&self) -> usize {
440 match self {
441 Self::Const(_) | Self::Var(_) => 1,
442 Self::Neg(a) | Self::Sin(a) | Self::Cos(a) | Self::Sqrt(a) | Self::Abs(a) => 1 + a.size(),
443 Self::Add(a, b)
444 | Self::Sub(a, b)
445 | Self::Mul(a, b)
446 | Self::Div(a, b)
447 | Self::Atan2(a, b)
448 | Self::Hypot(a, b)
449 | Self::Min(a, b)
450 | Self::Max(a, b) => 1 + a.size() + b.size(),
451 }
452 }
453}
454
455impl fmt::Display for Expr {
456 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
457 match self {
458 Self::Const(v) => write!(f, "{v}"),
459 Self::Var(v) => write!(f, "x{}", v.0),
460 Self::Add(a, b) => write!(f, "({a} + {b})"),
461 Self::Sub(a, b) => write!(f, "({a} - {b})"),
462 Self::Mul(a, b) => write!(f, "({a} * {b})"),
463 Self::Div(a, b) => write!(f, "({a} / {b})"),
464 Self::Neg(a) => write!(f, "-{a}"),
465 Self::Sin(a) => write!(f, "sin({a})"),
466 Self::Cos(a) => write!(f, "cos({a})"),
467 Self::Sqrt(a) => write!(f, "sqrt({a})"),
468 Self::Abs(a) => write!(f, "abs({a})"),
469 Self::Atan2(a, b) => write!(f, "atan2({a}, {b})"),
470 Self::Hypot(a, b) => write!(f, "hypot({a}, {b})"),
471 Self::Min(a, b) => write!(f, "min({a}, {b})"),
472 Self::Max(a, b) => write!(f, "max({a}, {b})"),
473 }
474 }
475}
476
477#[derive(Debug, Clone, PartialEq, Default)]
479pub struct LinearForm {
480 pub terms: Vec<(VarId, f64)>,
482 pub constant: f64,
484}
485
486impl LinearForm {
487 fn scaled(mut self, k: f64) -> Self {
488 for t in &mut self.terms {
489 t.1 *= k;
490 }
491 self.constant *= k;
492 self
493 }
494
495 fn combine(mut self, o: &Self, k: f64) -> Self {
496 for &(v, c) in &o.terms {
497 if let Some(t) = self.terms.iter_mut().find(|t| t.0 == v) {
498 t.1 += c * k;
499 } else {
500 self.terms.push((v, c * k));
501 }
502 }
503 self.constant += o.constant * k;
504 self.terms.sort_by_key(|t| t.0);
505 self.terms.retain(|t| t.1 != 0.0);
506 self
507 }
508}
509
510#[derive(Debug, Clone, PartialEq)]
512pub struct PointExpr {
513 pub x: Expr,
515 pub y: Expr,
517}
518
519impl PointExpr {
520 #[must_use]
522 pub const fn new(x: Expr, y: Expr) -> Self {
523 Self { x, y }
524 }
525
526 #[must_use]
528 pub const fn vars(x: VarId, y: VarId) -> Self {
529 Self { x: Expr::Var(x), y: Expr::Var(y) }
530 }
531
532 #[must_use]
534 pub const fn constant(x: f64, y: f64) -> Self {
535 Self { x: Expr::Const(x), y: Expr::Const(y) }
536 }
537
538 #[must_use]
540 pub fn transformed(&self, m: [f64; 6]) -> Self {
541 let [a, b, c, d, e, f] = m;
542 if m == [1.0, 0.0, 0.0, 1.0, 0.0, 0.0] {
543 return self.clone();
544 }
545 Self {
546 x: Expr::add(
547 Expr::add(Expr::mul(Expr::c(a), self.x.clone()), Expr::mul(Expr::c(c), self.y.clone())),
548 Expr::c(e),
549 ),
550 y: Expr::add(
551 Expr::add(Expr::mul(Expr::c(b), self.x.clone()), Expr::mul(Expr::c(d), self.y.clone())),
552 Expr::c(f),
553 ),
554 }
555 }
556
557 #[must_use]
559 pub fn eval(&self, x: &[f64]) -> (f64, f64) {
560 (self.x.eval(x), self.y.eval(x))
561 }
562}
563
564#[cfg(test)]
565mod tests {
566 use super::*;
567
568 fn v(i: u32) -> Expr {
569 Expr::Var(VarId(i))
570 }
571
572 fn numeric_grad(e: &Expr, x: &[f64]) -> Vec<f64> {
574 (0..x.len())
575 .map(|i| {
576 let h = 1e-6 * (1.0 + x[i].abs());
577 let mut xp = x.to_vec();
578 let mut xm = x.to_vec();
579 xp[i] += h;
580 xm[i] -= h;
581 (e.eval(&xp) - e.eval(&xm)) / (2.0 * h)
582 })
583 .collect()
584 }
585
586 #[test]
587 fn dual_matches_central_differences() {
588 let e = Expr::add(
589 Expr::mul(Expr::sin(v(0)), Expr::hypot(v(1), v(2))),
590 Expr::div(Expr::atan2(v(2), Expr::sub(v(0), Expr::c(3.0))), Expr::add(Expr::sqrt(v(1)), Expr::c(2.0))),
591 );
592 let x = [0.7, 2.3, -1.1];
593 let d = e.eval_dual(&x);
594 assert!((d.v - e.eval(&x)).abs() < 1e-15);
595 let num = numeric_grad(&e, &x);
596 for (i, n) in num.iter().enumerate() {
597 let a = d.g.iter().find(|g| g.0 == i as u32).map_or(0.0, |g| g.1);
598 assert!((a - n).abs() < 1e-7 * (1.0 + n.abs()), "var {i}: {a} vs {n}");
599 }
600 }
601
602 #[test]
603 fn linear_form_detection() {
604 let e = Expr::sub(Expr::add(Expr::mul(Expr::c(2.0), v(0)), v(1)), Expr::div(v(2), Expr::c(4.0)));
605 let lf = e.linear_form(&[0.0; 3], &|_| false).unwrap();
606 assert_eq!(lf.terms, vec![(VarId(0), 2.0), (VarId(1), 1.0), (VarId(2), -0.25)]);
607 let nl = Expr::mul(v(0), v(1));
608 assert!(nl.linear_form(&[1.0, 2.0], &|_| false).is_none());
609 let lf2 = nl.linear_form(&[1.0, 2.0], &|id| id == VarId(1)).unwrap();
611 assert_eq!(lf2.terms, vec![(VarId(0), 2.0)]);
612 let h = Expr::hypot(v(0), v(1));
614 assert_eq!(h.linear_form(&[3.0, 4.0], &|_| true).unwrap().constant, 5.0);
615 }
616
617 #[test]
618 fn simplification() {
619 assert_eq!(Expr::add(Expr::c(1.0), Expr::c(2.0)), Expr::c(3.0));
620 assert_eq!(Expr::mul(Expr::c(1.0), v(3)), v(3));
621 assert_eq!(Expr::mul(Expr::c(0.0), v(3)), Expr::c(0.0));
622 assert_eq!(Expr::neg(Expr::neg(v(1))), v(1));
623 }
624
625 #[test]
626 fn non_smooth_points_are_deterministic() {
627 let h = Expr::hypot(v(0), v(1)).eval_dual(&[0.0, 0.0]);
628 assert_eq!(h.v, 0.0);
629 assert_eq!(h.g, vec![(0, 1.0)]);
630 let s = Expr::sqrt(v(0)).eval_dual(&[0.0]);
631 assert_eq!(s.g, vec![(0, 0.0)]);
632 }
633}