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}