1use std::f64::consts::PI;
17
18#[must_use]
20pub fn norm_cdf(z: f64) -> f64 {
21 f64::midpoint(1.0, erf(z / std::f64::consts::SQRT_2))
22}
23
24#[must_use]
25pub fn norm_pdf(z: f64) -> f64 {
26 (-0.5 * z * z).exp() / (2.0 * PI).sqrt()
27}
28
29fn erf(x: f64) -> f64 {
30 let sign = if x < 0.0 { -1.0 } else { 1.0 };
32 let x = x.abs();
33 let t = 1.0 / 0.327_591_1f64.mul_add(x, 1.0);
34 let y = (1.061_405_429f64
35 .mul_add(t, -1.453_152_027)
36 .mul_add(t, 1.421_413_741)
37 .mul_add(t, -0.284_496_736)
38 .mul_add(t, 0.254_829_592)
39 * t)
40 .mul_add(-(-x * x).exp(), 1.0);
41 sign * y
42}
43
44fn rbf(a: &[f64], b: &[f64], length_scale: f64) -> f64 {
46 let mut sq = 0.0;
47 for (ai, bi) in a.iter().zip(b.iter()) {
48 let d = ai - bi;
49 sq = d.mul_add(d, sq);
50 }
51 (-sq / (2.0 * length_scale * length_scale)).exp()
52}
53
54#[derive(Debug, Clone)]
59pub struct GaussianProcess {
60 xs: Vec<Vec<f64>>,
62 ys: Vec<f64>,
64 pub length_scale: f64,
66 pub signal_variance: f64,
68 pub noise_variance: f64,
70 l: Option<Vec<f64>>,
72 alpha: Vec<f64>,
74 min_eig: f64,
76}
77
78impl Default for GaussianProcess {
79 fn default() -> Self {
80 Self::new(1.0, 1.0, 1e-6)
81 }
82}
83
84impl GaussianProcess {
85 #[must_use]
87 pub const fn new(length_scale: f64, signal_variance: f64, noise_variance: f64) -> Self {
88 Self {
89 xs: Vec::new(),
90 ys: Vec::new(),
91 length_scale: length_scale.max(1e-6),
92 signal_variance: signal_variance.max(1e-9),
93 noise_variance: noise_variance.max(1e-12),
94 l: None,
95 alpha: Vec::new(),
96 min_eig: f64::INFINITY,
97 }
98 }
99
100 #[must_use]
102 pub fn n_samples(&self) -> usize {
103 self.xs.len()
104 }
105
106 pub fn add_sample(&mut self, x: Vec<f64>, y: f64) {
108 self.xs.push(x);
109 self.ys.push(y);
110 }
111
112 pub fn set_hyperparameters(
114 &mut self,
115 length_scale: f64,
116 signal_variance: f64,
117 noise_variance: f64,
118 ) {
119 self.length_scale = length_scale.max(1e-6);
120 self.signal_variance = signal_variance.max(1e-9);
121 self.noise_variance = noise_variance.max(1e-12);
122 self.l = None;
124 self.alpha.clear();
125 }
126
127 pub fn log_marginal_likelihood(&self) -> Result<f64, String> {
133 let n = self.xs.len();
134 let l = self
135 .l
136 .as_ref()
137 .ok_or_else(|| "GP not fitted — call fit() first".to_string())?;
138 if n == 0 || self.alpha.len() != n {
139 return Err("GP not fitted — call fit() first".to_string());
140 }
141 let quad = self
143 .alpha
144 .iter()
145 .zip(self.ys.iter())
146 .map(|(a, y)| a * y)
147 .sum::<f64>();
148 let mut log_det = 0.0;
150 for i in 0..n {
151 log_det += l[i * n + i].ln();
152 }
153 log_det *= 2.0;
154 Ok((0.5 * n as f64).mul_add(-(2.0 * PI).ln(), (-0.5f64).mul_add(quad, -(0.5 * log_det))))
155 }
156
157 pub fn fit_hyperparameters(
164 &mut self,
165 n_initial: usize,
166 iterations: usize,
167 n_candidates: usize,
168 seed: u64,
169 ) -> Result<(), String> {
170 let n = self.xs.len();
171 if n < 3 {
172 return Err(format!(
173 "need at least 3 samples for hyperparameter fitting, got {n}"
174 ));
175 }
176 let bounds = [(-6.0, 4.0), (-8.0, 6.0), (-12.0, 2.0)]; let gp_cell = std::cell::RefCell::new(&mut *self);
179 let mut opt = BayesianOptimizer::new(
180 |h: &[f64]| {
181 let mut gp = gp_cell.borrow_mut();
182 gp.set_hyperparameters(h[0].exp(), h[1].exp(), h[2].exp());
183 match gp.fit() {
184 Ok(()) => gp.log_marginal_likelihood().unwrap_or(f64::NEG_INFINITY),
185 Err(_) => f64::NEG_INFINITY,
186 }
187 },
188 seed,
189 );
190 let (_, (best, _)) = opt
191 .optimize(&bounds, n_initial.max(3), iterations, n_candidates, 0.01)
192 .map_err(|e| format!("hyperparameter optimization failed: {e}"))?;
193 self.set_hyperparameters(best[0].exp(), best[1].exp(), best[2].exp());
195 self.fit()?;
196 Ok(())
197 }
198
199 pub fn fit(&mut self) -> Result<(), String> {
203 let n = self.xs.len();
204 if n < 2 {
205 return Err(format!(
206 "need at least 2 training samples to fit a GP, got {n}"
207 ));
208 }
209 let mut k = vec![0.0_f64; n * n];
211 for i in 0..n {
212 for j in 0..n {
213 k[i * n + j] =
214 self.signal_variance * rbf(&self.xs[i], &self.xs[j], self.length_scale);
215 }
216 k[i * n + i] += self.noise_variance;
217 }
218 let mut l = vec![0.0_f64; n * n];
220 self.min_eig = f64::INFINITY;
221 for i in 0..n {
222 for j in 0..=i {
223 let mut sum = k[i * n + j];
224 for kk in 0..j {
225 sum = l[i * n + kk].mul_add(-l[j * n + kk], sum);
226 }
227 if i == j {
228 if sum <= 0.0 {
229 sum = sum.max(1e-10);
231 }
232 let sqrt = sum.sqrt();
233 l[i * n + i] = sqrt;
234 self.min_eig = self.min_eig.min(sqrt * sqrt);
235 } else {
236 l[i * n + j] = sum / l[j * n + j];
237 }
238 }
239 }
240 let mut z = vec![0.0_f64; n];
242 for i in 0..n {
243 let mut sum = self.ys[i];
244 for j in 0..i {
245 sum = l[i * n + j].mul_add(-z[j], sum);
246 }
247 z[i] = sum / l[i * n + i];
248 }
249 let mut alpha = vec![0.0_f64; n];
250 for i in (0..n).rev() {
251 let mut sum = z[i];
252 for j in (i + 1)..n {
253 sum = l[j * n + i].mul_add(-alpha[j], sum);
254 }
255 alpha[i] = sum / l[i * n + i];
256 }
257 self.l = Some(l);
258 self.alpha = alpha;
259 Ok(())
260 }
261
262 pub fn predict(&self, x: &[f64]) -> Result<(f64, f64), String> {
267 let n = self.xs.len();
268 let l = self
269 .l
270 .as_ref()
271 .ok_or_else(|| "GP not fitted — call fit() first".to_string())?;
272 if x.len() != self.xs[0].len() {
273 return Err(format!(
274 "query dimension {} != training dimension {}",
275 x.len(),
276 self.xs[0].len()
277 ));
278 }
279 let mut kx = vec![0.0_f64; n];
281 for (i, xi) in self.xs.iter().enumerate() {
282 kx[i] = self.signal_variance * rbf(xi, x, self.length_scale);
283 }
284 let mut v = vec![0.0_f64; n];
286 for i in 0..n {
287 let mut sum = kx[i];
288 for j in 0..i {
289 sum = l[i * n + j].mul_add(-v[j], sum);
290 }
291 v[i] = sum / l[i * n + i];
292 }
293 let mean = self.alpha.iter().zip(kx.iter()).map(|(a, k)| a * k).sum();
294 let var =
296 (self.signal_variance + self.noise_variance) - v.iter().map(|vi| vi * vi).sum::<f64>();
297 Ok((mean, var.max(1e-12)))
298 }
299
300 #[must_use]
302 pub const fn min_eigenvalue(&self) -> f64 {
303 self.min_eig
304 }
305}
306
307pub fn expected_improvement(
312 gp: &GaussianProcess,
313 x: &[f64],
314 best_so_far: f64,
315 exploration: f64,
316) -> Result<f64, String> {
317 let (mean, var) = gp.predict(x)?;
318 let sigma = var.sqrt();
319 let diff = mean - best_so_far - exploration;
320 if sigma < 1e-12 {
321 return Ok(diff.max(0.0));
322 }
323 let z = diff / sigma;
324 Ok(diff.mul_add(norm_cdf(z), sigma * norm_pdf(z)))
325}
326
327#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
329pub struct OptimizationStep {
330 pub iteration: usize,
332 pub params: Vec<f64>,
334 pub fitness: f64,
336 pub surrogate_mean: f64,
338 pub surrogate_std: f64,
340}
341
342pub type OptimizeResult = (Vec<OptimizationStep>, (Vec<f64>, f64));
344
345pub struct BayesianOptimizer<F>
352where
353 F: Fn(&[f64]) -> f64,
354{
355 fitness: F,
356 rng: u64,
357}
358
359impl<F> BayesianOptimizer<F>
360where
361 F: Fn(&[f64]) -> f64,
362{
363 #[must_use]
365 pub const fn new(fitness: F, seed: u64) -> Self {
366 Self { fitness, rng: seed }
367 }
368
369 pub fn optimize(
374 &mut self,
375 bounds: &[(f64, f64)],
376 n_initial: usize,
377 n_iterations: usize,
378 n_candidates: usize,
379 exploration: f64,
380 ) -> Result<OptimizeResult, String> {
381 let dim = bounds.len();
382 if dim == 0 {
383 return Err("at least one parameter dimension required".into());
384 }
385 for (lo, hi) in bounds {
386 if lo > hi {
387 return Err(format!("invalid bounds [{lo}, {hi}]"));
388 }
389 }
390
391 let mut steps = Vec::new();
392 for i in 0..n_initial.max(1) {
394 let params = (0..dim)
395 .map(|d| {
396 let (lo, hi) = bounds[d];
397 (hi - lo).mul_add(rand_u01(&mut self.rng), lo)
398 })
399 .collect::<Vec<_>>();
400 let f = (self.fitness)(¶ms);
401 steps.push(OptimizationStep {
402 iteration: i,
403 surrogate_mean: f,
404 surrogate_std: 0.0,
405 fitness: f,
406 params,
407 });
408 }
409
410 let mut gp = GaussianProcess::default();
411 for s in &steps {
412 gp.add_sample(s.params.clone(), s.fitness);
413 }
414 gp.fit()?;
415
416 let mut best = steps
417 .iter()
418 .max_by(|a, b| {
419 a.fitness
420 .partial_cmp(&b.fitness)
421 .unwrap_or(std::cmp::Ordering::Equal)
422 })
423 .cloned()
424 .ok_or_else(|| "no initial samples".to_string())?;
425
426 for iter in 0..n_iterations {
428 let mut best_ei = f64::NEG_INFINITY;
429 let mut best_candidate = vec![0.0; dim];
430 for _ in 0..n_candidates.max(1) {
431 let candidate = (0..dim)
432 .map(|d| {
433 let (lo, hi) = bounds[d];
434 (hi - lo).mul_add(rand_u01(&mut self.rng), lo)
435 })
436 .collect::<Vec<_>>();
437 let ei =
438 expected_improvement(&gp, &candidate, best.fitness, exploration).unwrap_or(0.0);
439 if ei > best_ei {
440 best_ei = ei;
441 best_candidate = candidate;
442 }
443 }
444
445 let f = (self.fitness)(&best_candidate);
446 let (mean, var) = gp.predict(&best_candidate).unwrap_or((f, 1.0));
447 let step = OptimizationStep {
448 iteration: n_initial + iter,
449 surrogate_mean: mean,
450 surrogate_std: var.sqrt(),
451 fitness: f,
452 params: best_candidate,
453 };
454 if step.fitness > best.fitness {
455 best = step.clone();
456 }
457 gp.add_sample(step.params.clone(), step.fitness);
458 gp.fit()?;
459 steps.push(step);
460 }
461
462 Ok((steps, (best.params.clone(), best.fitness)))
463 }
464}
465
466pub(crate) fn rand_u01(state: &mut u64) -> f64 {
469 *state = state.wrapping_add(0x9E3779B97F4A7C15);
470 let mut z = *state;
471 z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
472 z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
473 z ^= z >> 31;
474 (z >> 11) as f64 / (1u64 << 53) as f64
475}
476
477#[derive(Debug, Clone)]
485pub struct Expr {
486 tokens: Vec<Token>,
487}
488
489#[derive(Debug, Clone, PartialEq)]
490enum Token {
491 Num(f64),
492 Var(usize),
493 Op(char),
494 LParen,
495 RParen,
496 Fn(String),
497 Comma,
498}
499
500impl Expr {
501 pub fn parse(src: &str) -> Result<Self, String> {
503 let mut tokens = Vec::new();
504 let chars: Vec<char> = src.chars().filter(|c| !c.is_whitespace()).collect();
505 let mut i = 0;
506 while i < chars.len() {
507 let c = chars[i];
508 match c {
509 '0'..='9' | '.' => {
510 let start = i;
511 while i < chars.len() && (chars[i].is_ascii_digit() || chars[i] == '.') {
512 i += 1;
513 }
514 let s: String = chars[start..i].iter().collect();
515 let v: f64 = s.parse().map_err(|_| format!("invalid number '{s}'"))?;
516 tokens.push(Token::Num(v));
517 }
518 'x' => {
519 if i + 1 >= chars.len() || chars[i + 1] != '[' {
521 return Err("expected 'x[i]' variable syntax".into());
522 }
523 let mut j = i + 2;
524 let mut idx = String::new();
525 while j < chars.len() && chars[j].is_ascii_digit() {
526 idx.push(chars[j]);
527 j += 1;
528 }
529 if j >= chars.len() || chars[j] != ']' {
530 return Err("unterminated 'x[i]' index".into());
531 }
532 let idx: usize = idx
533 .parse()
534 .map_err(|_| "x[] index must be a non-negative integer".to_string())?;
535 tokens.push(Token::Var(idx));
536 i = j + 1;
537 }
538 '(' => {
539 tokens.push(Token::LParen);
540 i += 1;
541 }
542 ')' => {
543 tokens.push(Token::RParen);
544 i += 1;
545 }
546 ',' => {
547 tokens.push(Token::Comma);
548 i += 1;
549 }
550 '+' | '-' | '*' | '/' | '^' => {
551 tokens.push(Token::Op(c));
552 i += 1;
553 }
554 c if c.is_alphabetic() => {
555 let start = i;
556 while i < chars.len() && chars[i].is_alphabetic() {
557 i += 1;
558 }
559 let name: String = chars[start..i].iter().collect();
560 match name.as_str() {
561 "pi" => tokens.push(Token::Num(PI)),
562 "e" => tokens.push(Token::Num(std::f64::consts::E)),
563 "sin" | "cos" | "tan" | "exp" | "log" | "sqrt" | "abs" => {
564 tokens.push(Token::Fn(name));
565 }
566 _ => return Err(format!("unknown function or constant '{name}'")),
567 }
568 }
569 other => return Err(format!("unexpected character '{other}'")),
570 }
571 }
572 if tokens.is_empty() {
573 return Err("empty expression".into());
574 }
575 let mut normalized = Vec::with_capacity(tokens.len());
578 for (i, tok) in tokens.iter().enumerate() {
579 let prev_is_value = i > 0
580 && matches!(
581 tokens[i - 1],
582 Token::Num(_) | Token::Var(_) | Token::RParen | Token::Fn(_)
583 );
584 let next_is_value = i + 1 < tokens.len()
585 && matches!(
586 tokens[i + 1],
587 Token::Num(_) | Token::Var(_) | Token::LParen | Token::Fn(_) | Token::Op('-')
588 );
589 match tok {
590 Token::Op('-') if !prev_is_value && next_is_value => {
591 normalized.push(Token::Op('~'));
592 }
593 Token::Op(c) if !prev_is_value || !next_is_value => {
594 return Err(format!("operator '{c}' in invalid position"));
595 }
596 other => normalized.push(other.clone()),
597 }
598 }
599 Ok(Self { tokens: normalized })
600 }
601
602 pub fn evaluate(&self, x: &[f64]) -> Result<f64, String> {
604 let mut stack: Vec<f64> = Vec::new();
605 let mut ops: Vec<Token> = Vec::new();
606 let mut i = 0;
607 while i < self.tokens.len() {
608 let tok = &self.tokens[i];
609 match tok {
610 Token::Num(v) => stack.push(*v),
611 Token::Var(idx) => {
612 let v = x
613 .get(*idx)
614 .ok_or_else(|| format!("x[{idx}] out of range (dim {})", x.len()))?;
615 stack.push(*v);
616 }
617 Token::Fn(name) => ops.push(Token::Fn(name.clone())),
618 Token::Op('~') => ops.push(Token::Op('~')),
619 Token::Op(op) => {
620 while let Some(top) = ops.last() {
621 if precedence(*op) <= precedence_from_token(top) {
622 if !reduce(&mut stack, top)? {
623 break;
624 }
625 ops.pop();
626 } else {
627 break;
628 }
629 }
630 ops.push(Token::Op(*op));
631 }
632 Token::LParen => ops.push(Token::LParen),
633 Token::RParen => {
634 while let Some(top) = ops.pop() {
635 if top == Token::LParen {
636 break;
637 }
638 if !reduce(&mut stack, &top)? {
639 return Err(String::from("expression error"));
640 }
641 }
642 if let Some(Token::Fn(name)) = ops.last() {
645 let name = name.clone();
646 ops.pop();
647 let arg = stack
648 .pop()
649 .ok_or_else(|| format!("'{name}' needs an argument"))?;
650 stack.push(apply_fn(&name, arg)?);
651 }
652 }
653 Token::Comma => {}
654 }
655 i += 1;
656 }
657 while let Some(top) = ops.pop() {
658 if top == Token::LParen {
659 return Err(String::from("unbalanced parentheses"));
660 }
661 reduce(&mut stack, &top)?;
662 }
663 stack.pop().ok_or_else(|| String::from("empty expression"))
664 }
665}
666
667fn reduce(stack: &mut Vec<f64>, op: &Token) -> Result<bool, String> {
670 match op {
671 Token::Op('~') => {
672 let a = stack
673 .pop()
674 .ok_or_else(|| String::from("expression error"))?;
675 stack.push(-a);
676 Ok(true)
677 }
678 Token::Fn(name) => {
679 let arg = stack
680 .pop()
681 .ok_or_else(|| format!("'{name}' needs an argument"))?;
682 stack.push(apply_fn(name, arg)?);
683 Ok(true)
684 }
685 Token::Op(c) => {
686 if stack.len() < 2 {
687 return Ok(false);
688 }
689 let b = stack
690 .pop()
691 .ok_or_else(|| String::from("expression error"))?;
692 let a = stack
693 .pop()
694 .ok_or_else(|| String::from("expression error"))?;
695 stack.push(apply_op(*c, a, b)?);
696 Ok(true)
697 }
698 _ => Err(String::from("internal parse error")),
699 }
700}
701
702const fn precedence(op: char) -> u8 {
703 match op {
704 '+' | '-' => 1,
705 '*' | '/' => 2,
706 '^' => 5, '~' => 4,
708 _ => 0,
709 }
710}
711
712const fn precedence_from_token(t: &Token) -> u8 {
713 match t {
714 Token::Op(c) => precedence(*c),
715 Token::Fn(_) => 4,
716 _ => 0,
717 }
718}
719
720fn apply_op(op: char, a: f64, b: f64) -> Result<f64, String> {
721 match op {
722 '+' => Ok(a + b),
723 '-' => Ok(a - b),
724 '*' => Ok(a * b),
725 '/' => {
726 if b.abs() < 1e-300 {
727 Err("division by zero".into())
728 } else {
729 Ok(a / b)
730 }
731 }
732 '^' => Ok(a.powf(b)),
733 _ => Err("unknown operator".into()),
734 }
735}
736
737fn apply_fn(name: &str, arg: f64) -> Result<f64, String> {
738 match name {
739 "sin" => Ok(arg.sin()),
740 "cos" => Ok(arg.cos()),
741 "tan" => Ok(arg.tan()),
742 "exp" => Ok(arg.exp()),
743 "log" => {
744 if arg <= 0.0 {
745 Err("log of non-positive value".into())
746 } else {
747 Ok(arg.ln())
748 }
749 }
750 "sqrt" => {
751 if arg < 0.0 {
752 Err("sqrt of negative value".into())
753 } else {
754 Ok(arg.sqrt())
755 }
756 }
757 "abs" => Ok(arg.abs()),
758 _ => Err(format!("unknown function '{name}'")),
759 }
760}
761
762#[cfg(test)]
763mod tests {
764 #![allow(clippy::suboptimal_flops)] use super::*;
766
767 #[test]
768 fn gp_fits_and_predicts_linear_trend() {
769 let mut gp = GaussianProcess::new(0.5, 1.0, 1e-6);
770 for i in 0..8 {
771 let x = f64::from(i);
772 gp.add_sample(vec![x], 2.0f64.mul_add(x, 1.0));
773 }
774 gp.fit().unwrap();
775 let (mean, var) = gp.predict(&[4.0]).unwrap();
776 assert!((mean - 9.0).abs() < 1.5, "mean {mean} near 9");
777 assert!(var >= 0.0);
778 assert_eq!(gp.n_samples(), 8);
779 }
780
781 #[test]
782 fn gp_needs_two_samples() {
783 let mut gp = GaussianProcess::default();
784 gp.add_sample(vec![0.0], 0.0);
785 assert!(gp.fit().is_err());
786 }
787
788 #[test]
789 fn gp_uncertainty_is_low_near_data() {
790 let mut gp = GaussianProcess::new(1.0, 1.0, 1e-8);
791 for i in 0..5 {
792 gp.add_sample(vec![f64::from(i)], f64::from(i).sin());
793 }
794 gp.fit().unwrap();
795 let (_, var_near) = gp.predict(&[2.0]).unwrap();
796 let (_, var_far) = gp.predict(&[100.0]).unwrap();
797 assert!(var_near < var_far, "variance grows away from data");
798 }
799
800 #[test]
801 fn optimizer_finds_optimum_of_parabola() {
802 let mut opt = BayesianOptimizer::new(|x: &[f64]| -(x[0] - 3.0).powi(2) + 5.0, 42);
803 let (steps, (best_params, best_fitness)) =
804 opt.optimize(&[(0.0, 10.0)], 5, 10, 200, 0.01).unwrap();
805 assert!(!steps.is_empty());
806 assert!(
807 (best_params[0] - 3.0).abs() < 0.5,
808 "best x = {}",
809 best_params[0]
810 );
811 assert!((best_fitness - 5.0).abs() < 0.5, "best f = {best_fitness}");
812 }
813
814 #[test]
815 fn optimizer_two_dimensions() {
816 let mut opt = BayesianOptimizer::new(
817 |x: &[f64]| (x[1] + 2.0).mul_add(-(x[1] + 2.0), -(x[0] - 1.0).powi(2)),
818 7,
819 );
820 let (_, (params, f)) = opt
821 .optimize(&[(0.0, 2.0), (-4.0, 0.0)], 5, 8, 150, 0.01)
822 .unwrap();
823 assert!((params[0] - 1.0).abs() < 0.5);
824 assert!((params[1] + 2.0).abs() < 0.5);
825 assert!(f > -0.6, "f = {f}");
826 }
827
828 #[test]
829 fn optimizer_rejects_invalid_bounds() {
830 let mut opt = BayesianOptimizer::new(|x: &[f64]| x[0], 1);
831 assert!(opt.optimize(&[(5.0, 1.0)], 3, 1, 10, 0.01).is_err());
832 assert!(opt.optimize(&[], 3, 1, 10, 0.01).is_err());
833 }
834
835 #[test]
836 fn expr_arithmetic() {
837 let e = Expr::parse("2 * x[0] + 1").unwrap();
838 assert!((e.evaluate(&[3.0]).unwrap() - 7.0).abs() < 1e-12);
839 let e = Expr::parse("x[0] ^ 2 + x[1] ^ 2").unwrap();
840 assert!((e.evaluate(&[3.0, 4.0]).unwrap() - 25.0).abs() < 1e-12);
841 }
842
843 #[test]
844 fn expr_functions_and_constants() {
845 let e = Expr::parse("sin(x[0]) + cos(x[0]) + pi").unwrap();
846 let v = e.evaluate(&[0.0]).unwrap();
847 assert!((v - (0.0 + 1.0 + PI)).abs() < 1e-12);
848 let e = Expr::parse("sqrt(abs(x[0]))").unwrap();
849 assert!((e.evaluate(&[-9.0]).unwrap() - 3.0).abs() < 1e-12);
850 let e = Expr::parse("exp(log(x[0]))").unwrap();
851 assert!((e.evaluate(&[7.0]).unwrap() - 7.0).abs() < 1e-9);
852 }
853
854 #[test]
855 fn expr_errors_are_safe() {
856 assert!(Expr::parse("").is_err());
857 assert!(Expr::parse("foo(1)").is_err());
858 assert!(Expr::parse("x[0] +").is_err());
859 let e = Expr::parse("1 / (x[0] - x[0])").unwrap();
860 assert!(e.evaluate(&[1.0]).is_err());
861 let e = Expr::parse("x[5]").unwrap();
862 assert!(e.evaluate(&[1.0]).is_err());
863 let e = Expr::parse("log(-1)").unwrap();
864 assert!(e.evaluate(&[0.0]).is_err());
865 }
866
867 #[test]
868 fn expr_nested_parentheses() {
869 let e = Expr::parse("(x[0] + 2) * (x[1] - 3)").unwrap();
870 assert!((e.evaluate(&[3.0, 5.0]).unwrap() - 10.0).abs() < 1e-12);
871 }
872
873 #[test]
874 fn expr_unary_minus() {
875 let e = Expr::parse("-x[0] + 1").unwrap();
876 assert!((e.evaluate(&[3.0]).unwrap() - -2.0).abs() < 1e-12);
877 let e = Expr::parse("2 * -x[0]").unwrap();
878 assert!((e.evaluate(&[4.0]).unwrap() - -8.0).abs() < 1e-12);
879 let e = Expr::parse("x[0] - -2").unwrap();
880 assert!((e.evaluate(&[3.0]).unwrap() - 5.0).abs() < 1e-12);
881 }
882
883 #[test]
884 fn expr_function_composition() {
885 let e = Expr::parse("sqrt(abs(x[0]))").unwrap();
886 assert!((e.evaluate(&[-16.0]).unwrap() - 4.0).abs() < 1e-12);
887 let e = Expr::parse("2 * sin(x[0]) + 1").unwrap();
888 let v = e.evaluate(&[0.0]).unwrap();
889 assert!((v - 1.0).abs() < 1e-12);
890 let e = Expr::parse("sin(x[0]) + cos(x[0])").unwrap();
891 assert!((e.evaluate(&[0.0]).unwrap() - 1.0).abs() < 1e-12);
892 }
893
894 #[test]
895 fn expr_exponent_binds_tighter_than_unary_minus() {
896 let e = Expr::parse("-x[0]^2").unwrap();
898 assert!((e.evaluate(&[3.0]).unwrap() - -9.0).abs() < 1e-12);
899 let e = Expr::parse("2 ^ 3 ^ 2").unwrap();
902 assert!((e.evaluate(&[]).unwrap() - 64.0).abs() < 1e-12);
903 }
904
905 #[test]
906 fn expr_leading_operator_rejected() {
907 assert!(Expr::parse("* x[0]").is_err());
908 assert!(Expr::parse("/ 2").is_err());
909 assert!(Expr::parse("^ x[0]").is_err());
910 }
911
912 #[test]
913 fn norm_cdf_bounds() {
914 assert!(norm_cdf(0.0) > 0.499 && norm_cdf(0.0) < 0.501);
915 assert!(norm_cdf(3.0) > 0.998);
916 assert!(norm_cdf(-3.0) < 0.002);
917 }
918
919 #[test]
920 fn expected_improvement_zero_variance() {
921 let mut gp = GaussianProcess::default();
922 gp.add_sample(vec![0.0], 1.0);
923 gp.add_sample(vec![1.0], 2.0);
924 gp.fit().unwrap();
925 let ei = expected_improvement(&gp, &[0.0], 5.0, 0.0).unwrap();
927 assert!(ei >= 0.0);
928 }
929
930 #[test]
931 fn lml_requires_fit() {
932 let gp = GaussianProcess::default();
933 assert!(gp.log_marginal_likelihood().is_err());
934 }
935
936 #[test]
937 fn lml_prefers_correct_length_scale() {
938 let xs: Vec<f64> = (0..30).map(|i| f64::from(i) * 0.1).collect();
941 let ys: Vec<f64> = xs.iter().map(|x| (6.0_f64 * x).sin()).collect();
942
943 let mut short = GaussianProcess::new(0.3, 1.0, 1e-3);
944 let mut long = GaussianProcess::new(5.0, 1.0, 1e-3);
945 for (x, y) in xs.iter().zip(ys.iter()) {
946 short.add_sample(vec![*x], *y);
947 long.add_sample(vec![*x], *y);
948 }
949 short.fit().unwrap();
950 long.fit().unwrap();
951 let lml_short = short.log_marginal_likelihood().unwrap();
952 let lml_long = long.log_marginal_likelihood().unwrap();
953 assert!(
954 lml_short > lml_long,
955 "ℓ=0.3 ({lml_short}) should beat ℓ=5.0 ({lml_long}) on high-frequency data"
956 );
957 }
958
959 #[test]
960 fn fit_hyperparameters_recovers_signal() {
961 let mut gp = GaussianProcess::new(1.0, 1.0, 0.01);
962 let xs: Vec<f64> = (0..25).map(|i| f64::from(i) * 0.15).collect();
963 let ys: Vec<f64> = xs.iter().map(|x| (2.5_f64 * x).sin()).collect();
964 for (x, y) in xs.iter().zip(ys.iter()) {
965 gp.add_sample(vec![*x], *y);
966 }
967 gp.fit_hyperparameters(6, 8, 120, 42).unwrap();
968 let (mean, _) = gp.predict(&[1.8]).unwrap();
970 let truth = (2.5_f64 * 1.8).sin();
971 assert!(
972 (mean - truth).abs() < 0.15,
973 "post-fit mean {mean} vs truth {truth}"
974 );
975 }
976
977 #[test]
978 fn fit_hyperparameters_needs_three_samples() {
979 let mut gp = GaussianProcess::default();
980 gp.add_sample(vec![0.0], 0.0);
981 gp.add_sample(vec![1.0], 1.0);
982 assert!(gp.fit_hyperparameters(3, 2, 50, 1).is_err());
983 }
984}