Skip to main content

wm_simulation/
bayesian.rs

1//! Gaussian Process surrogates and Bayesian optimization.
2//!
3//! Pure-Rust implementation (no external linear algebra):
4//!
5//! - [`GaussianProcess`] — RBF-kernel GP regression with Cholesky solve;
6//!   posterior mean + variance for any query point.
7//! - [`BayesianOptimizer`] — sequential optimization using Expected
8//!   Improvement over a GP surrogate.
9//! - [`Expr`] — a tiny safe expression evaluator (`x[0] * sin(x[1]) + 1`)
10//!   used as the fitness function, mirroring v26's `mc.optimize` tool.
11//!
12//! The surrogate gives the Dream cycle / Homeostatic loop cheap
13//! response-surface models, and `mc.optimize` replaces grid search with
14//! sample-efficient Bayesian exploration.
15
16use std::f64::consts::PI;
17
18/// Numerically safe inverse-normal CDF (Acklam's algorithm) and CDF (erf).
19#[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    // Abramowitz-Stegun 7.1.26 approximation (|err| < 1.5e-7)
31    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
44/// RBF (squared-exponential) kernel distance.
45fn 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/// Gaussian Process regression with an RBF kernel.
55///
56/// Hyperparameters are settable; the kernel is
57/// `k(x, x') = σ_f² · exp(−||x−x'||² / 2ℓ²)` with observation noise `σ_n`.
58#[derive(Debug, Clone)]
59pub struct GaussianProcess {
60    /// Training inputs (row = sample, col = dimension).
61    xs: Vec<Vec<f64>>,
62    /// Training outputs.
63    ys: Vec<f64>,
64    /// Length scale ℓ.
65    pub length_scale: f64,
66    /// Signal variance σ_f².
67    pub signal_variance: f64,
68    /// Observation noise variance σ_n².
69    pub noise_variance: f64,
70    /// Lower-triangular Cholesky factor L of (K + σ_n² I) (post-fit).
71    l: Option<Vec<f64>>,
72    /// K⁻¹ y (post-fit) — n×1 solved by forward/back substitution.
73    alpha: Vec<f64>,
74    /// Smallest Cholesky eigenvalue observed (stability diagnostic).
75    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    /// Create a GP with the given kernel hyperparameters.
86    #[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    /// Number of training samples.
101    #[must_use]
102    pub fn n_samples(&self) -> usize {
103        self.xs.len()
104    }
105
106    /// Add a training sample (x, y).
107    pub fn add_sample(&mut self, x: Vec<f64>, y: f64) {
108        self.xs.push(x);
109        self.ys.push(y);
110    }
111
112    /// Set kernel hyperparameters (all clamped to positive).
113    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        // Invalidate the previous fit
123        self.l = None;
124        self.alpha.clear();
125    }
126
127    /// Log marginal likelihood of the training data under the current
128    /// hyperparameters: `−½ yᵀK⁻¹y − ½ log|K| − ½ n log 2π`.
129    ///
130    /// Requires a successful [`fit`](Self::fit). Used for hyperparameter
131    /// optimization — higher is better.
132    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        // yᵀK⁻¹y = αᵀy (α = K⁻¹y)
142        let quad = self
143            .alpha
144            .iter()
145            .zip(self.ys.iter())
146            .map(|(a, y)| a * y)
147            .sum::<f64>();
148        // log|K| = 2·Σ log(L_ii)
149        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    /// Optimize the kernel hyperparameters (length scale, signal variance,
158    /// noise variance) by maximizing the log marginal likelihood.
159    ///
160    /// Uses the crate's own [`BayesianOptimizer`] over log-space
161    /// hyperparameters (dogfooding). `iterations` GP refits after the
162    /// initial random search; `n_candidates` EI candidates per iteration.
163    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        // Bounds on log-hyperparameters (generous)
177        let bounds = [(-6.0, 4.0), (-8.0, 6.0), (-12.0, 2.0)]; // ln ℓ, ln σ_f², ln σ_n²
178        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        // Apply the best hyperparameters
194        self.set_hyperparameters(best[0].exp(), best[1].exp(), best[2].exp());
195        self.fit()?;
196        Ok(())
197    }
198
199    /// Fit the GP: compute L = cholesky(K + σ_n² I) and α = K⁻¹ y.
200    ///
201    /// Returns an error if fewer than 2 samples are present.
202    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        // Kernel matrix
210        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        // Cholesky with jitter for numerical stability
219        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                        // Near-singular — add jitter and retry this diagonal
230                        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        // Solve L Lᵀ α = y  →  L z = y (forward), Lᵀ α = z (back)
241        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    /// Predict mean and variance at a query point.
263    ///
264    /// Returns `(mean, variance)`; variance includes observation noise.
265    /// Errors if not fitted.
266    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        // k(x) = kernel between x and all training points
280        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        // v = L⁻¹ k(x) — forward substitution
285        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        // var = k(x,x) − vᵀv  (+ noise for the predictive distribution)
295        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    /// Minimum Cholesky eigenvalue during fit (diagnostic).
301    #[must_use]
302    pub const fn min_eigenvalue(&self) -> f64 {
303        self.min_eig
304    }
305}
306
307/// Expected Improvement acquisition at a query point.
308///
309/// `z = (μ − f_best − ξ) / σ`; `EI = (μ − f_best − ξ)Φ(z) + σφ(z)`.
310/// `ξ > 0` biases exploration.
311pub 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/// One step of the Bayesian optimization trace.
328#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
329pub struct OptimizationStep {
330    /// Iteration index (0 = best of the initial random samples).
331    pub iteration: usize,
332    /// Evaluated parameters.
333    pub params: Vec<f64>,
334    /// Observed fitness.
335    pub fitness: f64,
336    /// Surrogate mean at the chosen point.
337    pub surrogate_mean: f64,
338    /// Surrogate std at the chosen point.
339    pub surrogate_std: f64,
340}
341
342/// Result of a Bayesian optimization run: full trace + best (params, fitness).
343pub type OptimizeResult = (Vec<OptimizationStep>, (Vec<f64>, f64));
344
345/// Sequential Bayesian optimization over a box.
346///
347/// 1. Sample `n_initial` random points, evaluate, take the best.
348/// 2. Fit a GP to all evaluations.
349/// 3. For each iteration: score `n_candidates` random points by Expected
350///    Improvement, evaluate the argmax, refit, repeat.
351pub 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    /// Create an optimizer over the fitness function with a PRNG seed.
364    #[must_use]
365    pub const fn new(fitness: F, seed: u64) -> Self {
366        Self { fitness, rng: seed }
367    }
368
369    /// Run the optimization.
370    ///
371    /// `bounds`: `[(min, max), ...]` per dimension (all required).
372    /// Returns the full trace and the best (params, fitness).
373    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        // Phase 1: random initialization
393        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)(&params);
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        // Phase 2: EI-guided search
427        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
466/// SplitMix64 PRNG — identical to the one in `monte_carlo.rs` so seeds
467/// behave consistently across modules.
468pub(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/// A tiny safe expression evaluator for fitness functions.
478///
479/// Grammar: numbers, `x[0]`, `x[i]`, `+ - * / ^ ( )`, unary minus,
480/// `sin cos tan exp log sqrt abs`, constants `pi`, `e`.
481///
482/// No eval, no I/O, no recursion depth risk (operand stack bounded by
483/// expression length) — safe to run on untrusted tool input.
484#[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    /// Parse an expression string.
502    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                    // x[0], x[1], ...
520                    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        // Structural validation: no leading/trailing/double binary operators.
576        // Rewrites prefix '-' into a dedicated negation token ('~').
577        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    /// Evaluate the expression at point `x`.
603    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                    // A function directly before the paren group applies to
643                    // the group's result: sin(x), log(-1), sqrt(abs(x))
644                    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
667/// Apply the top-of-ops operation to the operand stack. Returns `false`
668/// (without popping) when the operator cannot be applied yet.
669fn 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, // exponentiation binds tighter than unary minus: -x^2 = -(x^2)
707        '~' => 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)] // test data expressions, not hot paths
765    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        // -x^2 must be -(x^2), not (-x)^2
897        let e = Expr::parse("-x[0]^2").unwrap();
898        assert!((e.evaluate(&[3.0]).unwrap() - -9.0).abs() < 1e-12);
899        // x^2 ^ 3 is right-associative in math but left-to-right here;
900        // just verify the chain is deterministic
901        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        // At a known point EI ≈ 0 when far below best
926        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        // Data with a short length scale (high frequency): a too-long ℓ
939        // should give a lower LML than the true ℓ.
940        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        // After fitting, prediction should be accurate at a held-out point
969        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}