Skip to main content

dotloom_constraints/
expr.rs

1//! Scalar expression trees with exact derivatives (forward-mode automatic
2//! differentiation on sparse gradients).
3//!
4//! Every constraint row is an [`Expr`]. Built-in geometric rules and plugin
5//! expressions compile to the same representation, so one Jacobian implementation
6//! serves both. Gradients are exact (not finite differences); tests compare them
7//! with central differences.
8
9use core::fmt;
10
11use serde::{Deserialize, Serialize};
12
13/// Index of a solver variable within a [`crate::Problem`].
14#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
15#[serde(transparent)]
16pub struct VarId(pub u32);
17
18impl VarId {
19    /// Index into value vectors.
20    #[must_use]
21    pub const fn index(self) -> usize {
22        self.0 as usize
23    }
24}
25
26/// A scalar expression.
27#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28#[serde(tag = "op", content = "args", rename_all = "camelCase")]
29pub enum Expr {
30    /// Constant.
31    Const(f64),
32    /// Variable.
33    Var(VarId),
34    /// `a + b`.
35    Add(Box<Expr>, Box<Expr>),
36    /// `a - b`.
37    Sub(Box<Expr>, Box<Expr>),
38    /// `a * b`.
39    Mul(Box<Expr>, Box<Expr>),
40    /// `a / b`.
41    Div(Box<Expr>, Box<Expr>),
42    /// `-a`.
43    Neg(Box<Expr>),
44    /// `sin a`.
45    Sin(Box<Expr>),
46    /// `cos a`.
47    Cos(Box<Expr>),
48    /// `sqrt a` (domain `a ≥ 0`; negative inputs evaluate to NaN).
49    Sqrt(Box<Expr>),
50    /// `|a|`.
51    Abs(Box<Expr>),
52    /// `atan2(y, x)`.
53    Atan2(Box<Expr>, Box<Expr>),
54    /// `hypot(x, y)`.
55    Hypot(Box<Expr>, Box<Expr>),
56    /// `min(a, b)`.
57    Min(Box<Expr>, Box<Expr>),
58    /// `max(a, b)`.
59    Max(Box<Expr>, Box<Expr>),
60}
61
62/// Value plus sparse gradient, sorted by variable index.
63#[derive(Debug, Clone, PartialEq, Default)]
64pub struct Dual {
65    /// Value.
66    pub v: f64,
67    /// Non-zero partial derivatives `(var, ∂/∂var)`, sorted by var.
68    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// `add`/`sub`/... are simplifying *constructors* taking two expressions (no `self`
108// receiver); implementing the operator traits instead would hide the constant folding.
109#[allow(clippy::should_implement_trait, clippy::redundant_guards)]
110impl Expr {
111    /// Constant expression.
112    #[must_use]
113    pub const fn c(v: f64) -> Self {
114        Self::Const(v)
115    }
116
117    /// Variable expression.
118    #[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    /// Simplifying constructor for `a + b`.
128    #[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    /// Simplifying constructor for `a - b`.
139    #[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    /// Simplifying constructor for `a * b`.
150    #[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    /// Simplifying constructor for `a / b`.
162    #[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    /// Simplifying constructor for `-a`.
172    #[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    /// `sin a`.
182    #[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    /// `cos a`.
188    #[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    /// `sqrt a`.
194    #[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    /// `|a|`.
200    #[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    /// `atan2(y, x)`.
206    #[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    /// `hypot(x, y)`.
215    #[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    /// `min(a, b)`.
224    #[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    /// `max(a, b)`.
233    #[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    /// Evaluate with variable values `x` (missing variables evaluate to NaN).
242    #[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    /// Evaluate value and exact gradient.
264    ///
265    /// Non-smooth points use a deterministic one-sided derivative: `|0|' = 0`,
266    /// `sqrt'(0) = 0`, `hypot'(0, 0) = (1, 0)`, ties of `min/max` take the first argument.
267    #[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    /// Variables referenced by the expression (sorted, unique).
344    #[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    /// Affine form `Σ cᵢ·xᵢ + k` if the expression is exactly linear in the variables
373    /// whose `fixed[i]` is false (fixed variables are folded into the constant).
374    #[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            // Non-linear operators are linear only when fully constant.
403            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    /// Replace every variable by the expression returned by `f` (constant folding is
415    /// applied by the simplifying constructors).
416    #[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    /// Number of nodes (used for complexity limits).
438    #[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/// Affine form `Σ cᵢ·xᵢ + constant`.
478#[derive(Debug, Clone, PartialEq, Default)]
479pub struct LinearForm {
480    /// Coefficients (sorted by var, no duplicates after normalization).
481    pub terms: Vec<(VarId, f64)>,
482    /// Constant term.
483    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/// A 2D point as a pair of expressions.
511#[derive(Debug, Clone, PartialEq)]
512pub struct PointExpr {
513    /// X expression.
514    pub x: Expr,
515    /// Y expression.
516    pub y: Expr,
517}
518
519impl PointExpr {
520    /// Point from two expressions.
521    #[must_use]
522    pub const fn new(x: Expr, y: Expr) -> Self {
523        Self { x, y }
524    }
525
526    /// Point from two variables.
527    #[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    /// Constant point.
533    #[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    /// Apply an affine transform `[a b c d e f]` (SVG convention).
539    #[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    /// Evaluate.
558    #[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    /// Central-difference gradient for validation.
573    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        // Becomes linear when one factor is fixed.
610        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        // hypot of fixed vars folds to a constant.
613        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}