Skip to main content

hpr_analysis/
optimize.rs

1//! Optimization: the design variables that make a model's output as small as it can be.
2//!
3//! **Guide:** [Optimization][guide] runs the optimizer on test functions whose minima are known,
4//! then finds the ballast and body length that send a rocket to 3,048 m, chooses a motor and a
5//! catalog nose cone for it, traces a rocket's trade-off between apogee and stability, and says
6//! how far to trust it.
7//!
8//! [guide]: https://nrdptel.github.io/hpr-sim/optimization.html
9//!
10//! - [`Variable`]: one number the optimizer may change, with where it starts, the size of its
11//!   first steps, and optional bounds. An integer variable ([`Variable::integer`]) takes whole
12//!   numbers only: a count, or a choice from a list, such as a motor or a catalog part.
13//! - [`cmaes`]: the covariance matrix adaptation evolution strategy (CMA-ES), which samples
14//!   candidates around a mean, keeps the better half, and learns from them which way, and how
15//!   far, to step next. It needs only the output's ranking, no derivatives, so it suits flights,
16//!   whose outputs are noisy in their last digits.
17//! - [`nsga2`]: NSGA-II, a genetic algorithm for two or more goals at once (apogee against
18//!   stability, say), which finds the *Pareto front*: the designs where one goal can only be
19//!   bettered by giving up another.
20//! - [`ego`]: efficient global optimization (EGO), for a model so slow that only tens of
21//!   evaluations can be afforded: it fits a surrogate to the points evaluated so far and
22//!   evaluates next where the surrogate expects the most improvement.
23//! - [`Evaluation`]: a value and a constraint violation, for a model with constraints, ranked
24//!   by Deb's feasibility rules ([`cmaes::Run::tell_constrained`]).
25//! - [`benchmark`]: test functions with known minima, and test problems with known fronts, which
26//!   the tests hold the optimizers to.
27//!
28//! A model is minimized; to maximize an output, minimize its negative. To hit a target, minimize
29//! the squared miss, as the guide's example does.
30//!
31//! # Reproducibility
32//!
33//! A run is drawn from a seed. Each CMA-ES candidate has its own random stream
34//! ([`SeededRng::for_stream`](hpr_core::random::SeededRng::for_stream)), keyed by the seed, its
35//! generation and its place in the generation, and each NSGA-II generation has one stream, so a run is bit
36//! for bit the same every time on one platform, however its candidates are evaluated.
37//!
38//! # Left out
39//!
40//! EGO is checked on two and three variables only (on six it often stops at a local minimum),
41//! and optimizing a Monte Carlo run's statistics is a later increment of [M6.2, the optimization milestone][roadmap].
42//!
43//! [roadmap]: https://nrdptel.github.io/hpr-sim/decisions-and-roadmap.html#m6-2
44
45pub mod benchmark;
46pub mod cmaes;
47pub mod ego;
48mod eigen;
49mod normal;
50pub mod nsga2;
51
52use std::collections::BTreeSet;
53
54use serde::{Deserialize, Serialize};
55
56use crate::error::AnalysisError;
57
58/// The most variables an optimizer takes. CMA-ES holds an `n × n` covariance and decomposes it
59/// each generation, `O(n³)` work; 200 variables is far more than a rocket design has.
60pub const MAX_VARIABLES: usize = 200;
61
62/// The largest whole number an integer variable's bound may be, `2⁵²`: every whole number up to
63/// it, and every threshold halfway between two, is an `f64`.
64const MAX_WHOLE: f64 = 4_503_599_627_370_496.0;
65
66/// The whole number nearest `x`, a tie going to the lower, and 0 rather than −0. Exact: below
67/// 2⁵² `⌊x⌋ + 0.5` is an `f64`, and from there on `x` is whole (`x − 0.5` or `x − ⌊x⌋` would
68/// round).
69pub(crate) fn nearest_whole(x: f64) -> f64 {
70    let floor = x.floor();
71    let nearest = if x > floor + 0.5 { floor + 1.0 } else { floor };
72    nearest + 0.0
73}
74
75/// Checks an integer variable's bounds: whole numbers within `±2⁵²`.
76fn check_whole(low: f64, high: f64) -> Result<(), AnalysisError> {
77    let whole = |x: f64| x.fract() == 0.0 && x.abs() <= MAX_WHOLE;
78    if !whole(low) {
79        return Err(AnalysisError::Domain {
80            what: "integer variable's low bound (a whole number within ±2⁵²)",
81            value: low,
82        });
83    }
84    if !whole(high) {
85        return Err(AnalysisError::Domain {
86            what: "integer variable's high bound (a whole number within ±2⁵²)",
87            value: high,
88        });
89    }
90    Ok(())
91}
92
93/// A number the optimizer may change: its name, where it starts, the size of its first steps,
94/// the range it must stay in, and whether it takes only whole numbers. It serializes as its
95/// fields (an unbounded side as `null`; `integer` only when true, and read as `false` when
96/// absent), and reads back through [`Variable::new`], [`Variable::within`] and [`Variable::integer`]'s checks.
97#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
98#[serde(try_from = "VariableData", into = "VariableData")]
99pub struct Variable {
100    name: String,
101    start: f64,
102    step: f64,
103    low: f64,
104    high: f64,
105    integer: bool,
106}
107
108/// The serialized form of a [`Variable`].
109#[derive(Serialize, Deserialize)]
110#[serde(deny_unknown_fields)]
111struct VariableData {
112    name: String,
113    start: f64,
114    step: f64,
115    low: Option<f64>,
116    high: Option<f64>,
117    #[serde(default, skip_serializing_if = "std::ops::Not::not")]
118    integer: bool,
119}
120
121impl TryFrom<VariableData> for Variable {
122    type Error = AnalysisError;
123
124    fn try_from(data: VariableData) -> Result<Self, AnalysisError> {
125        let variable = Self::new(data.name, data.start, data.step)?.within(
126            data.low.unwrap_or(f64::NEG_INFINITY),
127            data.high.unwrap_or(f64::INFINITY),
128        )?;
129        if data.integer {
130            variable.integer()
131        } else {
132            Ok(variable)
133        }
134    }
135}
136
137impl From<Variable> for VariableData {
138    fn from(v: Variable) -> Self {
139        Self {
140            name: v.name,
141            start: v.start,
142            step: v.step,
143            low: v.low.is_finite().then_some(v.low),
144            high: v.high.is_finite().then_some(v.high),
145            integer: v.integer,
146        }
147    }
148}
149
150impl Variable {
151    /// A variable named `name`, starting at `start`, with no bounds. `step` is the size of the
152    /// optimizer's first steps in it: about a quarter to a third of the range its best value is
153    /// expected in, in the variable's own units.
154    ///
155    /// # Errors
156    ///
157    /// [`AnalysisError::Domain`] for a `start` that isn't finite, or a `step` that isn't finite
158    /// and positive.
159    pub fn new(name: impl Into<String>, start: f64, step: f64) -> Result<Self, AnalysisError> {
160        if !start.is_finite() {
161            return Err(AnalysisError::Domain {
162                what: "variable's start",
163                value: start,
164            });
165        }
166        if !(step.is_finite() && step > 0.0) {
167            return Err(AnalysisError::Domain {
168                what: "variable's step (finite, positive)",
169                value: step,
170            });
171        }
172        Ok(Self {
173            name: name.into(),
174            start,
175            step,
176            low: f64::NEG_INFINITY,
177            high: f64::INFINITY,
178            integer: false,
179        })
180    }
181
182    /// The same variable, held to `[low, high]`. Either bound may be infinite, for none.
183    ///
184    /// # Errors
185    ///
186    /// [`AnalysisError::Domain`] for a NaN bound, a `high` not above `low`, a start outside
187    /// the range, or for an integer variable bounds that aren't whole numbers within `±2⁵²`.
188    pub fn within(mut self, low: f64, high: f64) -> Result<Self, AnalysisError> {
189        if low.is_nan() || low == f64::INFINITY {
190            return Err(AnalysisError::Domain {
191                what: "variable's low bound",
192                value: low,
193            });
194        }
195        if high.is_nan() || high <= low {
196            return Err(AnalysisError::Domain {
197                what: "variable's high bound (above the low one)",
198                value: high,
199            });
200        }
201        if !(low..=high).contains(&self.start) {
202            return Err(AnalysisError::Domain {
203                what: "variable's start (within its bounds)",
204                value: self.start,
205            });
206        }
207        if self.integer {
208            check_whole(low, high)?;
209        }
210        self.low = low;
211        self.high = high;
212        Ok(self)
213    }
214
215    /// The same variable, taking only the whole numbers from its low bound to its high one: a
216    /// count, or the place of a choice in a list (a motor, a catalog part). The optimizer still
217    /// draws it as a real number, and the model is given the whole number nearest the draw,
218    /// clamped to the bounds ([`Variable::encode`]); [`cmaes`] keeps every value within reach by
219    /// the *margin* of CMA-ES with margin. Its step is the size of the first steps, in whole
220    /// numbers: 1 is a good start for a handful of choices. Its start may lie between two whole
221    /// numbers, as the draws are centered on it.
222    ///
223    /// # Errors
224    ///
225    /// [`AnalysisError::Domain`] for bounds that aren't whole numbers within `±2⁵²`: set them
226    /// first with [`Variable::within`], which checks them again if called after.
227    pub fn integer(mut self) -> Result<Self, AnalysisError> {
228        check_whole(self.low, self.high)?;
229        self.integer = true;
230        Ok(self)
231    }
232
233    /// Whether it takes only whole numbers ([`Variable::integer`]).
234    pub fn is_integer(&self) -> bool {
235        self.integer
236    }
237
238    /// The value the model is given for a draw `x`: `x` itself, or for an integer variable the
239    /// whole number nearest `x`, a draw halfway between two going to the lower, clamped to the
240    /// bounds. These are the *encoding* of R. Hamano et al., "CMA-ES with Margin", GECCO 2022
241    /// (arXiv:2205.13482, §4.1, p. 5), with thresholds halfway between neighbouring values. NaN
242    /// stays NaN (a run stops before it would draw one).
243    pub fn encode(&self, x: f64) -> f64 {
244        if self.integer {
245            // `+ 0.0` turns a bound written −0 into 0.
246            nearest_whole(x).clamp(self.low, self.high) + 0.0
247        } else {
248            x
249        }
250    }
251
252    /// Its name.
253    pub fn name(&self) -> &str {
254        &self.name
255    }
256
257    /// Where it starts.
258    pub fn start(&self) -> f64 {
259        self.start
260    }
261
262    /// The size of its first steps.
263    pub fn step(&self) -> f64 {
264        self.step
265    }
266
267    /// Its low bound, `−∞` for none.
268    pub fn low(&self) -> f64 {
269        self.low
270    }
271
272    /// Its high bound, `+∞` for none.
273    pub fn high(&self) -> f64 {
274        self.high
275    }
276
277    /// Whether `x` is within its bounds.
278    pub fn contains(&self, x: f64) -> bool {
279        (self.low..=self.high).contains(&x)
280    }
281}
282
283/// What a model gives for one candidate under constraints: its value, and by how much it breaks
284/// the constraints, zero if it keeps them all.
285///
286/// Candidates are ranked by K. Deb's feasibility rules ("An efficient constraint handling method
287/// for genetic algorithms", *Computer Methods in Applied Mechanics and Engineering* 186(2–4),
288/// 311–338 (2000), <https://doi.org/10.1016/S0045-7825(99)00389-8>, §3): a candidate that keeps
289/// every constraint beats one that doesn't; of two that keep them, the smaller value wins; of
290/// two that don't, the smaller violation wins (here ties in violation go to the smaller value).
291/// No penalty weight is needed, as values and violations are never compared with each other.
292/// The rules rank, and CMA-ES uses only ranks.
293///
294/// A candidate the model can't evaluate (a flight that fails) is `value` and `violation` both
295/// `+∞`: it ranks behind every other. Both serialize `+∞` as none, a JSON `null`.
296#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
297#[non_exhaustive]
298pub struct Evaluation {
299    /// The model's value.
300    #[serde(with = "cmaes::infinity_as_none")]
301    pub value: f64,
302    /// The total violation, `Σ max(0, gⱼ)` over constraints written `gⱼ ≤ 0`: zero if the
303    /// candidate keeps them all.
304    #[serde(with = "cmaes::infinity_as_none")]
305    pub violation: f64,
306}
307
308impl Evaluation {
309    /// A value with no constraints to break.
310    pub const fn feasible(value: f64) -> Self {
311        Self {
312            value,
313            violation: 0.0,
314        }
315    }
316
317    /// A candidate the model can't evaluate: value and violation both `+∞`, behind every other.
318    pub const fn failed() -> Self {
319        Self {
320            value: f64::INFINITY,
321            violation: f64::INFINITY,
322        }
323    }
324
325    /// A value under constraints `gⱼ(x) ≤ 0`, given as the numbers `gⱼ`: the violation is
326    /// `Σ max(0, gⱼ)`, Deb's (2000) overall violation. Deb divides each constraint by a constant
327    /// so that they count alike (a margin in calibers and a speed in m/s, say); do the same before
328    /// passing them in. A NaN `gⱼ` gives a NaN violation, which [`cmaes::Run::tell_constrained`]
329    /// refuses.
330    pub fn constrained(value: f64, constraints: &[f64]) -> Self {
331        // Folded from +0: an empty f64 sum is −0, which `total_cmp` would rank first.
332        let violation = constraints
333            .iter()
334            .map(|&g| if g.is_nan() || g > 0.0 { g } else { 0.0 })
335            .fold(0.0, |total, g| total + g);
336        Self { value, violation }
337    }
338
339    /// Whether the candidate keeps every constraint.
340    pub fn is_feasible(&self) -> bool {
341        self.violation == 0.0
342    }
343
344    /// Deb's rules as an ordering: [`Less`](std::cmp::Ordering::Less) if `self` ranks ahead of
345    /// `other`. Violation first, with `−0` equal to `0`, then value by [`f64::total_cmp`]
346    /// (CMA-ES's own ranking of plain values).
347    pub fn rank(&self, other: &Self) -> std::cmp::Ordering {
348        let violation = if self.violation == other.violation {
349            std::cmp::Ordering::Equal
350        } else {
351            self.violation.total_cmp(&other.violation)
352        };
353        violation.then(self.value.total_cmp(&other.value))
354    }
355
356    /// Whether `self` is strictly better than `other` by Deb's rules, comparing as `<` does, so
357    /// `−0` and `0` tie.
358    pub(crate) fn beats(&self, other: &Self) -> bool {
359        self.violation < other.violation
360            || (self.violation == other.violation && self.value < other.value)
361    }
362}
363
364/// Checks that there are between one and [`MAX_VARIABLES`] variables, and no two share a name.
365fn check_variables(variables: &[Variable]) -> Result<(), AnalysisError> {
366    if variables.is_empty() {
367        return Err(AnalysisError::TooFew {
368            what: "variables",
369            count: 0,
370            minimum: 1,
371        });
372    }
373    if variables.len() > MAX_VARIABLES {
374        return Err(AnalysisError::Count {
375            what: "variables",
376            count: variables.len(),
377            limit: MAX_VARIABLES,
378        });
379    }
380    let mut names = BTreeSet::new();
381    for variable in variables {
382        if !names.insert(variable.name()) {
383            return Err(AnalysisError::DuplicateVariable(variable.name().to_owned()));
384        }
385    }
386    Ok(())
387}
388
389#[cfg(test)]
390mod tests {
391    use super::*;
392
393    #[test]
394    fn variable_checks_its_numbers() {
395        assert!(Variable::new("x", f64::NAN, 1.0).is_err());
396        assert!(Variable::new("x", 0.0, 0.0).is_err());
397        assert!(Variable::new("x", 0.0, f64::INFINITY).is_err());
398        let x = Variable::new("x", 0.5, 0.1).unwrap();
399        assert!(x.clone().within(1.0, 2.0).is_err(), "start outside");
400        assert!(x.clone().within(1.0, 1.0).is_err(), "empty range");
401        assert!(x.clone().within(f64::NAN, 1.0).is_err());
402        assert!(x.clone().within(0.0, f64::NAN).is_err());
403        assert!(x.clone().within(f64::INFINITY, f64::INFINITY).is_err());
404        let x = x.within(0.0, f64::INFINITY).unwrap();
405        assert!(x.contains(0.0) && x.contains(1e300) && !x.contains(-1e-300));
406    }
407
408    #[test]
409    fn evaluations_rank_by_deb_rules() {
410        use std::cmp::Ordering::{Greater, Less};
411        let e = Evaluation::constrained(5.0, &[-1.0, 0.0, -3.0]);
412        assert!(e.is_feasible());
413        let broken = Evaluation::constrained(-100.0, &[0.25, -1.0, 0.5]);
414        assert_eq!(broken.violation, 0.75);
415        // Feasible beats infeasible, whatever the values.
416        assert_eq!(e.rank(&broken), Less);
417        // Two infeasible: the smaller violation, whatever the values.
418        assert_eq!(broken.rank(&Evaluation::constrained(-1e9, &[1.0])), Less);
419        // Two feasible: the smaller value.
420        assert_eq!(e.rank(&Evaluation::feasible(4.0)), Greater);
421        assert!(Evaluation::constrained(0.0, &[f64::NAN]).violation.is_nan());
422        // No constraints is +0, not the −0 of an empty sum, so values decide.
423        let none = Evaluation::constrained(10.0, &[]);
424        assert!(none.violation.is_sign_positive());
425        assert_eq!(none.rank(&Evaluation::feasible(1.0)), Greater);
426        let negative_zero = Evaluation {
427            value: 10.0,
428            violation: -0.0,
429        };
430        assert_eq!(negative_zero.rank(&Evaluation::feasible(1.0)), Greater);
431        // A failure ranks behind everything, and reads back from JSON.
432        assert_eq!(broken.rank(&Evaluation::failed()), Less);
433        let json = serde_json::to_string(&Evaluation::failed()).unwrap();
434        assert_eq!(json, r#"{"value":null,"violation":null}"#);
435        assert_eq!(
436            serde_json::from_str::<Evaluation>(&json).unwrap(),
437            Evaluation::failed()
438        );
439    }
440
441    #[test]
442    fn variable_serializes_unbounded_sides_as_null_and_rechecks_on_reading() {
443        let x = Variable::new("ballast", 0.2, 0.1)
444            .unwrap()
445            .within(0.0, f64::INFINITY)
446            .unwrap();
447        let json = serde_json::to_string(&x).unwrap();
448        assert_eq!(
449            json,
450            r#"{"name":"ballast","start":0.2,"step":0.1,"low":0.0,"high":null}"#
451        );
452        assert_eq!(serde_json::from_str::<Variable>(&json).unwrap(), x);
453        // `integer` is written only when true, and may be written false.
454        let explicit =
455            r#"{"name":"ballast","start":0.2,"step":0.1,"low":0.0,"high":null,"integer":false}"#;
456        assert_eq!(serde_json::from_str::<Variable>(explicit).unwrap(), x);
457        let bad = r#"{"name":"x","start":2.0,"step":0.1,"low":0.0,"high":1.0}"#;
458        assert!(serde_json::from_str::<Variable>(bad).is_err());
459        let motor = Variable::new("motor", 2.0, 1.0)
460            .unwrap()
461            .within(0.0, 4.0)
462            .unwrap()
463            .integer()
464            .unwrap();
465        let json = serde_json::to_string(&motor).unwrap();
466        assert_eq!(
467            json,
468            r#"{"name":"motor","start":2.0,"step":1.0,"low":0.0,"high":4.0,"integer":true}"#
469        );
470        assert_eq!(serde_json::from_str::<Variable>(&json).unwrap(), motor);
471        // An integer variable's bounds are checked on reading too.
472        let bad = r#"{"name":"k","start":0.0,"step":1.0,"low":-0.5,"high":3.0,"integer":true}"#;
473        assert!(serde_json::from_str::<Variable>(bad).is_err());
474    }
475
476    #[test]
477    fn integer_variables_need_whole_bounds() {
478        let k = Variable::new("k", 1.5, 1.0).unwrap();
479        for (low, high) in [
480            (f64::NEG_INFINITY, 4.0),
481            (0.0, f64::INFINITY),
482            (0.5, 4.0),
483            (0.0, 3.5),
484            (-1e16, 4.0),
485        ] {
486            let err = k.clone().within(low, high).unwrap().integer().unwrap_err();
487            let AnalysisError::Domain { what, .. } = err else {
488                panic!("{low}, {high}: {err:?}");
489            };
490            assert!(what.starts_with("integer variable's"), "{what}");
491            // Bounds set after `integer` are checked too.
492            let err = k
493                .clone()
494                .within(0.0, 4.0)
495                .unwrap()
496                .integer()
497                .unwrap()
498                .within(low, high)
499                .unwrap_err();
500            let AnalysisError::Domain { what, .. } = err else {
501                panic!("{low}, {high} after: {err:?}");
502            };
503            assert!(what.starts_with("integer variable's"), "{what}");
504        }
505        // A start between two whole numbers is allowed: the draws are centered on it.
506        let k = k.within(0.0, 4.0).unwrap().integer().unwrap();
507        assert!(k.is_integer() && !Variable::new("x", 0.0, 1.0).unwrap().is_integer());
508    }
509
510    /// The nearest whole number, halfway going to the lower, clamped to the bounds; a continuous
511    /// variable's draw is its own value.
512    #[test]
513    fn integer_draws_encode_to_the_nearest_value() {
514        let k = Variable::new("k", 0.0, 1.0)
515            .unwrap()
516            .within(-2.0, 3.0)
517            .unwrap()
518            .integer()
519            .unwrap();
520        for (x, value) in [
521            (0.0, 0.0_f64),
522            (0.5, 0.0),
523            (0.500_000_000_000_1, 1.0),
524            (-0.5, -1.0),
525            (-0.499_999_999_999_9, 0.0),
526            (1e-300, 0.0),
527            (2.5, 2.0),
528            (2.6, 3.0),
529            (40.0, 3.0),
530            (-1.5, -2.0),
531            (-7.2, -2.0),
532            (f64::MAX, 3.0),
533            (-0.0, 0.0),
534            (-0.3, 0.0),
535            // `x + 1` would round to 0.5 here, a tie.
536            (-0.499_999_999_999_999_94, 0.0),
537            (0.499_999_999_999_999_94, 0.0),
538        ] {
539            // Bits, so that −0 isn't taken for 0.
540            assert_eq!(k.encode(x).to_bits(), value.to_bits(), "{x}");
541        }
542        // Up to 2⁵², where `x − 0.5` would round to an even neighbour; past it, clamped.
543        let wide = Variable::new("k", 0.0, 1.0)
544            .unwrap()
545            .within(-MAX_WHOLE, MAX_WHOLE)
546            .unwrap()
547            .integer()
548            .unwrap();
549        for (x, value) in [
550            (MAX_WHOLE - 1.0, MAX_WHOLE - 1.0),
551            (MAX_WHOLE - 0.5, MAX_WHOLE - 1.0),
552            (MAX_WHOLE, MAX_WHOLE),
553            (MAX_WHOLE + 1.0, MAX_WHOLE),
554            (-MAX_WHOLE + 0.5, -MAX_WHOLE),
555        ] {
556            assert_eq!(wide.encode(x), value, "{x}");
557        }
558        for high in [MAX_WHOLE + 1.0, 2.0 * MAX_WHOLE] {
559            let err = Variable::new("k", 0.0, 1.0)
560                .unwrap()
561                .within(0.0, high)
562                .unwrap()
563                .integer()
564                .unwrap_err();
565            let AnalysisError::Domain { what, value } = err else {
566                panic!("{high}: {err:?}");
567            };
568            assert!(what.starts_with("integer variable's high bound"), "{what}");
569            assert_eq!(value, high);
570        }
571        let x = Variable::new("x", 0.0, 1.0).unwrap();
572        assert_eq!(x.encode(0.7), 0.7);
573    }
574
575    #[test]
576    fn variables_need_distinct_names_and_a_count_in_range() {
577        let x = Variable::new("x", 0.0, 1.0).unwrap();
578        assert!(matches!(
579            check_variables(&[]),
580            Err(AnalysisError::TooFew {
581                what: "variables",
582                ..
583            })
584        ));
585        assert!(matches!(
586            check_variables(&[x.clone(), x.clone()]),
587            Err(AnalysisError::DuplicateVariable(name)) if name == "x"
588        ));
589        let many = vec![x; MAX_VARIABLES + 1];
590        assert!(matches!(
591            check_variables(&many),
592            Err(AnalysisError::Count {
593                what: "variables",
594                count,
595                ..
596            }) if count == MAX_VARIABLES + 1
597        ));
598    }
599}