Skip to main content

active_quadrature/
active_quadrature.rs

1// Sequential active Bayesian quadrature over a finite candidate grid.
2//
3// Run with: cargo run --example active_quadrature
4
5use uncertain_numerics::{
6    ActiveBayesianQuadrature, BayesianQuadrature, GaussianMeasure, RbfKernel,
7};
8
9// A Gaussian bump centred at 0.5. Against p(x) = N(0, 1) its integral has a closed form.
10fn integrand(x: f64) -> f64 {
11    let centered = x - 0.5;
12    (-centered * centered).exp()
13}
14
15fn main() -> Result<(), Box<dyn std::error::Error>> {
16    let kernel = RbfKernel::new(1.0, 1.0)?;
17    let measure = GaussianMeasure::new(0.0, 1.0)?;
18    let quadrature = BayesianQuadrature::new(kernel, measure, 1.0e-10);
19    let active = ActiveBayesianQuadrature::new(quadrature);
20
21    // Start from three evaluations; allow up to six more from a fixed candidate grid,
22    // stopping early once the posterior variance of the integral drops below 1e-6.
23    let initial_nodes = [-1.0, 0.0, 1.0];
24    let initial_values: Vec<f64> = initial_nodes.iter().copied().map(integrand).collect();
25    let candidates: Vec<f64> = (0..=23).map(|i| -2.875 + 0.25 * f64::from(i)).collect();
26
27    let initial = quadrature.posterior(&initial_nodes, &initial_values)?;
28    println!("initial posterior variance = {:.3e}", initial.variance());
29
30    let result = active.run(
31        &initial_nodes,
32        &initial_values,
33        &candidates,
34        6,
35        1.0e-6,
36        integrand,
37    )?;
38
39    for step in result.steps() {
40        println!(
41            "x = {:+.3}  predicted reduction = {:.3e}  posterior variance = {:.3e}",
42            step.point(),
43            step.predicted_variance_reduction(),
44            step.posterior_variance(),
45        );
46    }
47
48    let exact = (1.0_f64 / 3.0).sqrt() * (-0.25_f64 / 3.0).exp();
49    let posterior = result.posterior();
50    println!("stopped because: {:?}", result.termination());
51    println!(
52        "E[I | y] = {:.6} ± {:.3e}   (exact {exact:.6})",
53        posterior.mean(),
54        posterior.standard_deviation(),
55    );
56    Ok(())
57}