1use hpr_core::random::SeededRng;
66use serde::{Deserialize, Serialize};
67
68use super::cmaes::Cmaes;
69use super::{Variable, normal};
70use crate::error::AnalysisError;
71
72pub const DESIGNS: usize = 100;
75
76pub const SEARCH_POINTS: usize = 200;
79
80pub const NUGGET: f64 = 1e-8;
83
84pub const MAX_POINTS: usize = 1000;
87
88pub const MAX_EGO_VARIABLES: usize = 50;
91
92const LOG_THETA: (f64, f64) = (-3.0, 3.0);
95
96#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
111#[non_exhaustive]
112pub enum Transform {
113 #[default]
115 None,
116 Log,
119 NegativeLog,
123}
124
125impl Transform {
126 pub fn apply(self, y: f64) -> Option<f64> {
129 match self {
130 _ if y == f64::INFINITY => Some(y),
131 Self::None => Some(y),
132 Self::Log => (y > 0.0).then(|| y.ln()),
133 Self::NegativeLog => (y < 0.0).then(|| -(-y).ln()),
134 }
135 }
136
137 fn domain(self) -> &'static str {
139 match self {
140 Self::None => "EGO value",
141 Self::Log => "EGO value under the ln(y) transform (positive)",
142 Self::NegativeLog => "EGO value under the -ln(-y) transform (negative)",
143 }
144 }
145}
146
147#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
168#[serde(try_from = "EgoData")]
169pub struct Ego {
170 variables: Vec<Variable>,
171 initial: usize,
172 max_evaluations: usize,
173 target: Option<f64>,
174 tolerance_improvement: f64,
175 transform: Transform,
176}
177
178#[derive(Deserialize)]
180#[serde(deny_unknown_fields)]
181struct EgoData {
182 variables: Vec<Variable>,
183 initial: usize,
184 max_evaluations: usize,
185 target: Option<f64>,
186 tolerance_improvement: f64,
187 #[serde(default)]
188 transform: Transform,
189}
190
191impl TryFrom<EgoData> for Ego {
192 type Error = AnalysisError;
193
194 fn try_from(data: EgoData) -> Result<Self, AnalysisError> {
195 let ego = Ego::new(data.variables)?
196 .with_initial(data.initial)?
197 .with_max_evaluations(data.max_evaluations)?
198 .with_tolerance_improvement(data.tolerance_improvement)?
199 .with_transform(data.transform);
200 match data.target {
201 Some(target) => ego.with_target(target),
202 None => Ok(ego),
203 }
204 }
205}
206
207#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
209#[non_exhaustive]
210pub enum Stop {
211 Target,
213 Evaluations,
215 Improvement,
218}
219
220#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
222#[non_exhaustive]
223pub struct Optimum {
224 pub point: Vec<f64>,
226 #[serde(with = "super::cmaes::infinity_as_none")]
229 pub value: f64,
230 pub evaluation: usize,
232 pub evaluations: usize,
234 pub stop: Stop,
236}
237
238fn check_points(what: &'static str, count: usize) -> Result<usize, AnalysisError> {
240 if count > MAX_POINTS {
241 return Err(AnalysisError::Count {
242 what,
243 count,
244 limit: MAX_POINTS,
245 });
246 }
247 Ok(count)
248}
249
250impl Ego {
251 pub fn new(variables: Vec<Variable>) -> Result<Self, AnalysisError> {
261 if variables.is_empty() {
262 return Err(AnalysisError::TooFew {
263 what: "variables",
264 count: 0,
265 minimum: 1,
266 });
267 }
268 if variables.len() > MAX_EGO_VARIABLES {
269 return Err(AnalysisError::Count {
270 what: "EGO variables",
271 count: variables.len(),
272 limit: MAX_EGO_VARIABLES,
273 });
274 }
275 for v in &variables {
276 if !(v.low.is_finite() && v.high.is_finite()) {
277 return Err(AnalysisError::Domain {
278 what: "EGO variable's bounds (both finite)",
279 value: if v.low.is_finite() { v.high } else { v.low },
280 });
281 }
282 if v.integer {
283 return Err(AnalysisError::Domain {
284 what: "EGO variable (continuous only)",
285 value: v.low,
286 });
287 }
288 }
289 let n = variables.len();
290 Ok(Self {
291 variables,
292 initial: 10 * n,
293 max_evaluations: 20 * n,
294 target: None,
295 tolerance_improvement: 0.0,
296 transform: Transform::None,
297 })
298 }
299
300 pub fn with_initial(mut self, points: usize) -> Result<Self, AnalysisError> {
307 check_points("initial design's points", points)?;
308 if points < 2 {
309 return Err(AnalysisError::TooFew {
310 what: "initial design's points",
311 count: points,
312 minimum: 2,
313 });
314 }
315 self.initial = points;
316 Ok(self)
317 }
318
319 pub fn with_max_evaluations(mut self, max: usize) -> Result<Self, AnalysisError> {
328 check_points("evaluations", max)?;
329 if max == 0 {
330 return Err(AnalysisError::TooFew {
331 what: "evaluations",
332 count: 0,
333 minimum: 1,
334 });
335 }
336 self.max_evaluations = max;
337 Ok(self)
338 }
339
340 pub fn with_target(mut self, target: f64) -> Result<Self, AnalysisError> {
346 if !target.is_finite() {
347 return Err(AnalysisError::Domain {
348 what: "target (finite)",
349 value: target,
350 });
351 }
352 self.target = Some(target);
353 Ok(self)
354 }
355
356 pub fn with_tolerance_improvement(mut self, tolerance: f64) -> Result<Self, AnalysisError> {
368 if !(tolerance.is_finite() && tolerance >= 0.0) {
369 return Err(AnalysisError::Domain {
370 what: "improvement tolerance",
371 value: tolerance,
372 });
373 }
374 self.tolerance_improvement = tolerance;
375 Ok(self)
376 }
377
378 #[must_use]
382 pub fn with_transform(mut self, transform: Transform) -> Self {
383 self.transform = transform;
384 self
385 }
386
387 pub fn transform(&self) -> Transform {
389 self.transform
390 }
391
392 pub fn variables(&self) -> &[Variable] {
394 &self.variables
395 }
396
397 pub fn minimize(
409 &self,
410 seed: u64,
411 mut model: impl FnMut(&[f64]) -> f64,
412 ) -> Result<Optimum, AnalysisError> {
413 let n = self.variables.len();
414 let mut points: Vec<Vec<f64>> = Vec::new();
415 let mut values: Vec<f64> = Vec::new();
416 let mut best: Option<usize> = None;
417 let mut evaluate = |unit: Vec<f64>,
418 points: &mut Vec<Vec<f64>>,
419 values: &mut Vec<f64>,
420 best: &mut Option<usize>|
421 -> Result<(), AnalysisError> {
422 let value = model(&self.to_variables(&unit));
423 if value.is_nan() || value == f64::NEG_INFINITY {
424 return Err(AnalysisError::Output {
425 index: values.len(),
426 value,
427 });
428 }
429 if self.transform.apply(value).is_none() {
430 return Err(AnalysisError::Domain {
431 what: self.transform.domain(),
432 value,
433 });
434 }
435 if best.is_none_or(|b| value < values[b]) {
436 *best = Some(values.len());
437 }
438 points.push(unit);
439 values.push(value);
440 Ok(())
441 };
442 let finish = |points: &[Vec<f64>], values: &[f64], best: usize, stop: Stop| Optimum {
443 point: self.to_variables(&points[best]),
444 value: values[best],
445 evaluation: best + 1,
446 evaluations: values.len(),
447 stop,
448 };
449
450 let mut rng = SeededRng::for_stream(seed, &[0]);
451 for unit in latin_hypercube(self.initial.min(self.max_evaluations), n, &mut rng) {
452 evaluate(unit, &mut points, &mut values, &mut best)?;
453 let b = best.unwrap_or(0);
454 if self.target.is_some_and(|t| values[b] <= t) {
455 return Ok(finish(&points, &values, b, Stop::Target));
456 }
457 if values.len() >= self.max_evaluations {
458 return Ok(finish(&points, &values, b, Stop::Evaluations));
459 }
460 }
461
462 let mut log_theta = vec![0.0; n];
463 loop {
464 let b = best.unwrap_or(0);
466 let mut rng = SeededRng::for_stream(seed, &[1, values.len() as u64]);
467 let transformed: Vec<f64> = values
470 .iter()
471 .map(|v| self.transform.apply(*v).unwrap_or(f64::INFINITY))
472 .collect();
473 let worst = transformed
474 .iter()
475 .copied()
476 .filter(|v| v.is_finite())
477 .fold(f64::NEG_INFINITY, f64::max);
478 let worst = if worst.is_finite() { worst } else { 0.0 };
479 let fitted: Vec<f64> = transformed
480 .iter()
481 .map(|v| if v.is_finite() { *v } else { worst })
482 .collect();
483 let (mean, sd) = mean_and_sd(&fitted);
484 let standard: Vec<f64> = fitted.iter().map(|v| (v - mean) / sd).collect();
485 log_theta = fit_log_theta(&points, &standard, &log_theta, rng.next_u64())?;
486 let theta: Vec<f64> = log_theta.iter().map(|l| 10f64.powf(*l)).collect();
487 let next = match Kriging::fit(&points, &standard, &theta) {
488 Some(kriging) => {
489 let f_min = standard[b];
490 let (next, improvement) =
491 kriging.most_improving(f_min, &points[b], &mut rng)?;
492 if improvement * sd < self.tolerance_improvement {
493 return Ok(finish(&points, &values, b, Stop::Improvement));
494 }
495 next
496 }
497 None => (0..n).map(|_| rng.uniform()).collect(),
499 };
500 evaluate(next, &mut points, &mut values, &mut best)?;
501 let b = best.unwrap_or(0);
502 if self.target.is_some_and(|t| values[b] <= t) {
503 return Ok(finish(&points, &values, b, Stop::Target));
504 }
505 if values.len() >= self.max_evaluations {
506 return Ok(finish(&points, &values, b, Stop::Evaluations));
507 }
508 }
509 }
510
511 fn to_variables(&self, unit: &[f64]) -> Vec<f64> {
513 self.variables
514 .iter()
515 .zip(unit)
516 .map(|(v, u)| (v.low + u * (v.high - v.low)).clamp(v.low, v.high))
517 .collect()
518 }
519}
520
521fn mean_and_sd(values: &[f64]) -> (f64, f64) {
523 let count = values.len() as f64;
525 let mean = values.iter().sum::<f64>() / count;
526 let variance = values.iter().map(|v| (v - mean) * (v - mean)).sum::<f64>() / count;
527 let sd = variance.sqrt();
528 (mean, if sd > 0.0 { sd } else { 1.0 })
529}
530
531fn latin_hypercube(points: usize, n: usize, rng: &mut SeededRng) -> Vec<Vec<f64>> {
535 let mut best: (f64, Vec<Vec<f64>>) = (f64::NEG_INFINITY, Vec::new());
536 for _ in 0..DESIGNS {
537 let mut design = vec![vec![0.0; n]; points];
538 for k in 0..n {
539 let mut slices: Vec<usize> = (0..points).collect();
541 for i in (1..points).rev() {
542 let j = (rng.uniform() * (i + 1) as f64) as usize;
544 slices.swap(i, j);
545 }
546 for (row, slice) in design.iter_mut().zip(slices) {
547 row[k] = (slice as f64 + rng.uniform()) / points as f64;
549 }
550 }
551 let mut least = f64::INFINITY;
552 for i in 0..points {
553 for j in 0..i {
554 least = least.min(squared_distance(&design[i], &design[j]));
555 }
556 }
557 if least > best.0 {
558 best = (least, design);
559 }
560 }
561 best.1
562}
563
564fn squared_distance(a: &[f64], b: &[f64]) -> f64 {
566 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
567}
568
569fn fit_log_theta(
571 points: &[Vec<f64>],
572 values: &[f64],
573 start: &[f64],
574 seed: u64,
575) -> Result<Vec<f64>, AnalysisError> {
576 let variables = start
577 .iter()
578 .enumerate()
579 .map(|(k, s)| {
580 Variable::new(format!("log10 theta {k}"), *s, 1.0)?.within(LOG_THETA.0, LOG_THETA.1)
581 })
582 .collect::<Result<Vec<_>, _>>()?;
583 let n = start.len();
584 let optimum = Cmaes::new(variables)?
585 .with_max_evaluations(100 * n)?
586 .with_tolerance_x(1e-3)?
587 .with_tolerance_value(1e-6)?
588 .minimize(seed, |log_theta| {
589 let theta: Vec<f64> = log_theta.iter().map(|l| 10f64.powf(*l)).collect();
590 Kriging::fit(points, values, &theta).map_or(f64::INFINITY, |k| -k.log_likelihood)
591 })?;
592 Ok(if optimum.value.is_finite() {
594 optimum.point
595 } else {
596 start.to_vec()
597 })
598}
599
600#[derive(Debug, Clone)]
602struct Kriging<'a> {
603 points: &'a [Vec<f64>],
604 theta: &'a [f64],
605 factor: Vec<f64>,
607 mu: f64,
609 sigma2: f64,
611 weights: Vec<f64>,
613 inverse_ones: Vec<f64>,
615 ones_inverse_ones: f64,
617 log_likelihood: f64,
620}
621
622impl<'a> Kriging<'a> {
623 fn fit(points: &'a [Vec<f64>], values: &[f64], theta: &'a [f64]) -> Option<Self> {
626 let m = points.len();
627 let mut factor = vec![0.0; m * m];
628 for i in 0..m {
629 for j in 0..i {
630 factor[i * m + j] = correlation(&points[i], &points[j], theta);
631 }
632 factor[i * m + i] = 1.0 + NUGGET;
633 }
634 cholesky(&mut factor, m)?;
635 let inverse_ones = solve(&factor, m, &vec![1.0; m]);
636 let ones_inverse_ones: f64 = inverse_ones.iter().sum();
637 let inverse_values = solve(&factor, m, values);
638 let mu = inverse_values.iter().sum::<f64>() / ones_inverse_ones;
639 let weights: Vec<f64> = inverse_values
640 .iter()
641 .zip(&inverse_ones)
642 .map(|(v, o)| v - mu * o)
643 .collect();
644 let sigma2 = values
646 .iter()
647 .zip(&weights)
648 .map(|(y, w)| (y - mu) * w)
649 .sum::<f64>()
650 / m as f64;
651 let log_determinant: f64 = (0..m).map(|i| 2.0 * factor[i * m + i].ln()).sum();
652 let log_likelihood = -0.5 * m as f64 * sigma2.ln() - 0.5 * log_determinant;
654 (sigma2 > 0.0 && log_likelihood.is_finite()).then_some(Self {
655 points,
656 theta,
657 factor,
658 mu,
659 sigma2,
660 weights,
661 inverse_ones,
662 ones_inverse_ones,
663 log_likelihood,
664 })
665 }
666
667 fn predict(&self, x: &[f64]) -> (f64, f64) {
672 let m = self.points.len();
673 let r: Vec<f64> = self
674 .points
675 .iter()
676 .map(|p| correlation(x, p, self.theta))
677 .collect();
678 let prediction = self.mu + dot(&r, &self.weights);
679 let half = forward(&self.factor, m, &r);
680 let ones = 1.0 - dot(&self.inverse_ones, &r);
681 let variance =
682 self.sigma2 * (1.0 - dot(&half, &half) + ones * ones / self.ones_inverse_ones);
683 (prediction, variance.max(0.0).sqrt())
684 }
685
686 fn expected_improvement(&self, x: &[f64], f_min: f64) -> f64 {
688 let (prediction, error) = self.predict(x);
689 expected_improvement(f_min - prediction, error)
690 }
691
692 fn most_improving(
696 &self,
697 f_min: f64,
698 best: &[f64],
699 rng: &mut SeededRng,
700 ) -> Result<(Vec<f64>, f64), AnalysisError> {
701 let n = best.len();
702 let mut start = (f64::NEG_INFINITY, best.to_vec());
703 for _ in 0..SEARCH_POINTS * n {
704 let x: Vec<f64> = (0..n).map(|_| rng.uniform()).collect();
705 let improvement = self.expected_improvement(&x, f_min);
706 if improvement > start.0 {
707 start = (improvement, x);
708 }
709 }
710 let mut found = (f64::NEG_INFINITY, best.to_vec());
711 for from in [start.1, best.to_vec()] {
712 let variables = from
713 .iter()
714 .enumerate()
715 .map(|(k, s)| Variable::new(format!("x{k}"), *s, 0.1)?.within(0.0, 1.0))
716 .collect::<Result<Vec<_>, _>>()?;
717 let optimum = Cmaes::new(variables)?
718 .with_max_evaluations(200 * n)?
719 .with_tolerance_x(1e-9)?
720 .minimize(rng.next_u64(), |x| -self.expected_improvement(x, f_min))?;
721 if -optimum.value > found.0 {
722 found = (-optimum.value, optimum.point);
723 }
724 }
725 Ok((found.1, found.0))
726 }
727}
728
729fn expected_improvement(d: f64, s: f64) -> f64 {
732 if s > 0.0 {
733 let z = d / s;
734 let density = (-0.5 * z * z).exp() / (2.0 * std::f64::consts::PI).sqrt();
735 (d * normal::cdf(z) + s * density).max(0.0)
736 } else {
737 d.max(0.0)
738 }
739}
740
741fn correlation(a: &[f64], b: &[f64], theta: &[f64]) -> f64 {
743 let exponent: f64 = a
744 .iter()
745 .zip(b)
746 .zip(theta)
747 .map(|((x, y), t)| t * (x - y) * (x - y))
748 .sum();
749 (-exponent).exp()
750}
751
752fn dot(a: &[f64], b: &[f64]) -> f64 {
754 a.iter().zip(b).map(|(x, y)| x * y).sum()
755}
756
757fn cholesky(a: &mut [f64], m: usize) -> Option<()> {
761 for j in 0..m {
762 let mut diagonal = a[j * m + j];
763 for k in 0..j {
764 diagonal -= a[j * m + k] * a[j * m + k];
765 }
766 if diagonal.is_nan() || diagonal <= 0.0 {
767 return None;
768 }
769 let pivot = diagonal.sqrt();
770 a[j * m + j] = pivot;
771 for i in j + 1..m {
772 let mut sum = a[i * m + j];
773 for k in 0..j {
774 sum -= a[i * m + k] * a[j * m + k];
775 }
776 a[i * m + j] = sum / pivot;
777 }
778 for k in j + 1..m {
779 a[j * m + k] = 0.0;
780 }
781 }
782 Some(())
783}
784
785fn forward(factor: &[f64], m: usize, b: &[f64]) -> Vec<f64> {
787 let mut x = b.to_vec();
788 for i in 0..m {
789 for k in 0..i {
790 x[i] -= factor[i * m + k] * x[k];
791 }
792 x[i] /= factor[i * m + i];
793 }
794 x
795}
796
797fn solve(factor: &[f64], m: usize, b: &[f64]) -> Vec<f64> {
799 let mut x = forward(factor, m, b);
800 for i in (0..m).rev() {
801 for k in i + 1..m {
802 x[i] -= factor[k * m + i] * x[k];
803 }
804 x[i] /= factor[i * m + i];
805 }
806 x
807}
808
809#[cfg(test)]
810mod tests {
811 #![allow(clippy::unwrap_used, reason = "tests stop at the failure")]
812
813 use super::*;
814
815 #[test]
817 fn kriging_interpolates() {
818 let points: Vec<Vec<f64>> = (0..8)
819 .map(|i| vec![f64::from(i) / 7.0, (f64::from(i) * 0.37).fract()])
820 .collect();
821 let values: Vec<f64> = points.iter().map(|p| (3.0 * p[0]).sin() + p[1]).collect();
822 let theta = [2.0, 3.0];
823 let kriging = Kriging::fit(&points, &values, &theta).unwrap();
824 for (p, v) in points.iter().zip(&values) {
825 let (prediction, error) = kriging.predict(p);
826 assert!((prediction - v).abs() < 1e-6, "{prediction} against {v}");
827 assert!(error < 1e-3, "{error}");
828 }
829 let (_, between) = kriging.predict(&[0.5, 0.9]);
830 assert!(between > 1e-3, "{between}");
831 }
832
833 #[test]
835 fn likelihood_two_points() {
836 let points = vec![vec![0.0], vec![0.5]];
837 let values = [1.0, -1.0];
838 let theta = [2.0];
839 let kriging = Kriging::fit(&points, &values, &theta).unwrap();
840 let rho = (-0.5f64).exp();
841 let a = 1.0 + NUGGET;
842 let determinant = a * a - rho * rho;
844 let sigma2 = (2.0 * a + 2.0 * rho) / determinant / 2.0;
845 assert!(kriging.mu.abs() < 1e-15);
846 assert!((kriging.sigma2 - sigma2).abs() < 1e-14 * sigma2);
847 let expected = -sigma2.ln() - 0.5 * determinant.ln();
848 assert!((kriging.log_likelihood - expected).abs() < 1e-14);
849 }
850
851 #[test]
854 #[allow(
855 clippy::needless_range_loop,
856 reason = "the adjugate's indices, as written by hand"
857 )]
858 fn kriging_three_points_by_hand() {
859 let points = vec![vec![0.0], vec![0.3], vec![1.0]];
860 let values = [2.0, -1.0, 0.5];
861 let theta = [1.7];
862 let kriging = Kriging::fit(&points, &values, &theta).unwrap();
863 let c = |a: f64, b: f64| (-1.7 * (a - b) * (a - b)).exp();
864 let x = [0.0, 0.3, 1.0];
865 let mut r = [[0.0; 3]; 3];
866 for i in 0..3 {
867 for j in 0..3 {
868 r[i][j] = if i == j { 1.0 + NUGGET } else { c(x[i], x[j]) };
869 }
870 }
871 let det = r[0][0] * (r[1][1] * r[2][2] - r[1][2] * r[2][1])
872 - r[0][1] * (r[1][0] * r[2][2] - r[1][2] * r[2][0])
873 + r[0][2] * (r[1][0] * r[2][1] - r[1][1] * r[2][0]);
874 let mut inv = [[0.0; 3]; 3];
875 for i in 0..3 {
876 for j in 0..3 {
877 let (a, b) = ((j + 1) % 3, (j + 2) % 3);
879 let (p, q) = ((i + 1) % 3, (i + 2) % 3);
880 inv[i][j] = (r[a][p] * r[b][q] - r[a][q] * r[b][p]) / det;
881 }
882 }
883 let times = |v: &[f64; 3]| -> [f64; 3] {
884 [0, 1, 2].map(|i| (0..3).map(|j| inv[i][j] * v[j]).sum())
885 };
886 let ones = times(&[1.0; 3]);
887 let ones_ones: f64 = ones.iter().sum();
888 let inv_y = times(&values);
889 let mu = inv_y.iter().sum::<f64>() / ones_ones;
890 let centered = values.map(|v| v - mu);
891 let weights = times(¢ered);
892 let sigma2 = (0..3).map(|i| centered[i] * weights[i]).sum::<f64>() / 3.0;
893 assert!(
894 (kriging.mu - mu).abs() < 1e-9,
895 "{} against {mu}",
896 kriging.mu
897 );
898 assert!((kriging.sigma2 - sigma2).abs() < 1e-9 * sigma2);
899 let at = 1.8;
900 let rv = x.map(|xi| c(at, xi));
901 let prediction = mu + (0..3).map(|i| rv[i] * weights[i]).sum::<f64>();
902 let inv_r = times(&rv);
903 let r_inv_r: f64 = (0..3).map(|i| rv[i] * inv_r[i]).sum();
904 let ones_r: f64 = (0..3).map(|i| ones[i] * rv[i]).sum();
905 let variance = sigma2 * (1.0 - r_inv_r + (1.0 - ones_r) * (1.0 - ones_r) / ones_ones);
906 let (got, error) = kriging.predict(&[at]);
907 assert!(
908 (got - prediction).abs() < 1e-9,
909 "{got} against {prediction}"
910 );
911 assert!(
912 (error - variance.sqrt()).abs() < 1e-9,
913 "{error} against {}",
914 variance.sqrt()
915 );
916 assert!((1.0 - ones_r).powi(2) / ones_ones > 1e-3);
918 assert!(mu.abs() > 0.1);
919 }
920
921 #[test]
923 fn expected_improvement_values() {
924 let s = 0.3;
925 let at_zero = expected_improvement(0.0, s);
926 assert!((at_zero - s / (2.0 * std::f64::consts::PI).sqrt()).abs() < 1e-16);
927 assert!((expected_improvement(5.0, s) - 5.0).abs() < 1e-12);
928 assert!(expected_improvement(-5.0, s) < 1e-30);
929 assert_eq!(expected_improvement(0.25, 0.0), 0.25);
930 assert_eq!(expected_improvement(-0.25, 0.0), 0.0);
931 }
932
933 #[test]
935 fn latin_hypercube_fills_every_slice() {
936 let mut rng = SeededRng::for_stream(7, &[0]);
937 let design = latin_hypercube(12, 3, &mut rng);
938 for k in 0..3 {
939 let mut slices: Vec<usize> = design.iter().map(|p| (p[k] * 12.0) as usize).collect();
941 slices.sort_unstable();
942 assert_eq!(slices, (0..12).collect::<Vec<_>>());
943 }
944 }
945
946 #[test]
948 fn refuses_unbounded_and_integer() {
949 let open = Variable::new("x", 0.0, 1.0).unwrap();
950 assert!(matches!(
951 Ego::new(vec![open.clone()]),
952 Err(AnalysisError::Domain { what, .. }) if what.starts_with("EGO variable's bounds")
953 ));
954 let whole = open.within(0.0, 4.0).unwrap().integer().unwrap();
955 assert!(matches!(
956 Ego::new(vec![whole]),
957 Err(AnalysisError::Domain { what, .. }) if what.starts_with("EGO variable (continuous")
958 ));
959 }
960}