Skip to main content

uncertain_numerics/
bayesian_quadrature.rs

1//! End-to-end one-dimensional Bayesian quadrature for the first supported pair.
2
3use crate::{
4    BayesianQuadratureError, GaussianConditioner, GaussianMeasure, KernelIntegral, KernelMean,
5    RbfKernel, ScalarKernel, ScalarNormalPosterior,
6};
7
8/// Bayesian quadrature with an RBF covariance kernel and Gaussian integration measure.
9///
10/// The current implementation assumes a zero Gaussian-process prior mean. This is
11/// intentionally explicit rather than hidden behind a generic prior-mean abstraction.
12#[derive(Debug, Clone, Copy, PartialEq)]
13pub struct BayesianQuadrature {
14    kernel: RbfKernel,
15    measure: GaussianMeasure,
16    jitter: f64,
17}
18
19impl BayesianQuadrature {
20    /// Construct the first supported Bayesian quadrature configuration.
21    #[must_use]
22    pub const fn new(kernel: RbfKernel, measure: GaussianMeasure, jitter: f64) -> Self {
23        Self {
24            kernel,
25            measure,
26            jitter,
27        }
28    }
29
30    /// Return the RBF kernel.
31    #[must_use]
32    pub const fn kernel(&self) -> RbfKernel {
33        self.kernel
34    }
35
36    /// Return the Gaussian integration measure.
37    #[must_use]
38    pub const fn measure(&self) -> GaussianMeasure {
39        self.measure
40    }
41
42    /// Return the fixed diagonal jitter used by Gaussian conditioning.
43    #[must_use]
44    pub const fn jitter(&self) -> f64 {
45        self.jitter
46    }
47
48    /// Compute the posterior distribution of the integral from observed function values.
49    ///
50    /// For observations `y = f(X)` and zero prior mean,
51    ///
52    /// ```text
53    /// posterior_mean = z^T (K + jitter I)^(-1) y
54    /// posterior_var  = kappa - z^T (K + jitter I)^(-1) z
55    /// ```
56    ///
57    /// The inverse is never formed explicitly; both systems are solved from one
58    /// reusable Cholesky factorization.
59    ///
60    /// # Errors
61    ///
62    /// Returns [`BayesianQuadratureError`] for invalid observations, conditioning
63    /// failures, or an invalid posterior variance.
64    pub fn posterior(
65        &self,
66        nodes: &[f64],
67        values: &[f64],
68    ) -> Result<ScalarNormalPosterior, BayesianQuadratureError> {
69        validate_observations(nodes, values)?;
70
71        let dimension = nodes.len();
72        let mut gram = Vec::with_capacity(dimension * dimension);
73        for &left in nodes {
74            for &right in nodes {
75                gram.push(self.kernel.covariance(left, right));
76            }
77        }
78
79        let kernel_mean: Vec<f64> = nodes
80            .iter()
81            .map(|&node| self.kernel.kernel_mean(&self.measure, node))
82            .collect();
83
84        let conditioner = GaussianConditioner::new(&gram, dimension, self.jitter)?;
85        let alpha = conditioner.solve(values)?;
86        let v = conditioner.solve(&kernel_mean)?;
87
88        let posterior_mean = dot(&kernel_mean, &alpha);
89        let prior_integral_variance = self.kernel.kernel_integral(&self.measure);
90        let raw_variance = prior_integral_variance - dot(&kernel_mean, &v);
91        let posterior_variance =
92            non_negative_roundoff_variance(raw_variance, prior_integral_variance)?;
93
94        Ok(ScalarNormalPosterior::new(
95            posterior_mean,
96            posterior_variance,
97        )?)
98    }
99}
100
101fn validate_observations(nodes: &[f64], values: &[f64]) -> Result<(), BayesianQuadratureError> {
102    if nodes.is_empty() {
103        return Err(BayesianQuadratureError::EmptyObservations);
104    }
105    if nodes.len() != values.len() {
106        return Err(BayesianQuadratureError::ObservationLengthMismatch);
107    }
108    if nodes.iter().any(|value| !value.is_finite()) {
109        return Err(BayesianQuadratureError::NonFiniteObservationNode);
110    }
111    if values.iter().any(|value| !value.is_finite()) {
112        return Err(BayesianQuadratureError::NonFiniteObservationValue);
113    }
114
115    Ok(())
116}
117
118fn dot(left: &[f64], right: &[f64]) -> f64 {
119    debug_assert_eq!(left.len(), right.len());
120    left.iter().zip(right).map(|(x, y)| x * y).sum()
121}
122
123fn non_negative_roundoff_variance(value: f64, scale: f64) -> Result<f64, BayesianQuadratureError> {
124    if value >= 0.0 {
125        return Ok(value);
126    }
127
128    let tolerance = 64.0 * f64::EPSILON * scale.abs().max(1.0);
129    if value >= -tolerance {
130        Ok(0.0)
131    } else {
132        Err(BayesianQuadratureError::MateriallyNegativePosteriorVariance { value, tolerance })
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::{BayesianQuadrature, non_negative_roundoff_variance};
139    use crate::{
140        BayesianQuadratureError, ConditioningError, GaussianMeasure, KernelIntegral, KernelMean,
141        RbfKernel,
142    };
143
144    const TOLERANCE: f64 = 1.0e-11;
145
146    fn assert_close(actual: f64, expected: f64) {
147        let scale = expected.abs().max(1.0);
148        assert!(
149            (actual - expected).abs() <= TOLERANCE * scale,
150            "expected {expected:.16e}, got {actual:.16e}"
151        );
152    }
153
154    fn fixture() -> BayesianQuadrature {
155        let kernel = RbfKernel::new(1.0, 1.2).expect("kernel parameters are valid");
156        let measure = GaussianMeasure::new(0.0, 1.0).expect("measure parameters are valid");
157        BayesianQuadrature::new(kernel, measure, 1.0e-12)
158    }
159
160    #[test]
161    fn one_observation_matches_closed_form_conditioning() {
162        let quadrature = fixture();
163        let posterior = quadrature
164            .posterior(&[0.0], &[2.0])
165            .expect("one-node posterior is valid");
166
167        let z = quadrature.kernel().kernel_mean(&quadrature.measure(), 0.0);
168        let kappa = quadrature.kernel().kernel_integral(&quadrature.measure());
169        let denominator = quadrature.kernel().signal_variance() + quadrature.jitter();
170
171        assert_close(posterior.mean(), z * 2.0 / denominator);
172        assert_close(posterior.variance(), kappa - z * z / denominator);
173    }
174
175    #[test]
176    fn zero_observations_produce_zero_posterior_mean() {
177        let quadrature = fixture();
178        let posterior = quadrature
179            .posterior(&[-1.0, 0.0, 1.0], &[0.0, 0.0, 0.0])
180            .expect("posterior is valid");
181
182        assert_close(posterior.mean(), 0.0);
183        assert!(posterior.variance() >= 0.0);
184    }
185
186    #[test]
187    fn posterior_variance_does_not_depend_on_observed_values() {
188        let quadrature = fixture();
189        let first = quadrature
190            .posterior(&[-0.75, 0.25, 1.5], &[1.0, 2.0, -1.0])
191            .expect("posterior is valid");
192        let second = quadrature
193            .posterior(&[-0.75, 0.25, 1.5], &[10.0, -5.0, 3.0])
194            .expect("posterior is valid");
195
196        assert_close(first.variance(), second.variance());
197    }
198
199    #[test]
200    fn posterior_variance_is_no_greater_than_prior_integral_variance() {
201        let quadrature = fixture();
202        let posterior = quadrature
203            .posterior(&[-2.0, -0.5, 0.5, 2.0], &[1.0, 0.5, -0.5, -1.0])
204            .expect("posterior is valid");
205        let prior_variance = quadrature.kernel().kernel_integral(&quadrature.measure());
206
207        assert!(posterior.variance() >= 0.0);
208        assert!(posterior.variance() <= prior_variance + TOLERANCE);
209    }
210
211    #[test]
212    fn rejects_invalid_observations() {
213        let quadrature = fixture();
214
215        assert_eq!(
216            quadrature.posterior(&[], &[]),
217            Err(BayesianQuadratureError::EmptyObservations)
218        );
219        assert_eq!(
220            quadrature.posterior(&[0.0], &[1.0, 2.0]),
221            Err(BayesianQuadratureError::ObservationLengthMismatch)
222        );
223        assert_eq!(
224            quadrature.posterior(&[f64::NAN], &[1.0]),
225            Err(BayesianQuadratureError::NonFiniteObservationNode)
226        );
227        assert_eq!(
228            quadrature.posterior(&[0.0], &[f64::INFINITY]),
229            Err(BayesianQuadratureError::NonFiniteObservationValue)
230        );
231    }
232
233    #[test]
234    fn duplicate_nodes_fail_without_jitter() {
235        let kernel = RbfKernel::new(1.0, 1.0).expect("kernel parameters are valid");
236        let measure = GaussianMeasure::new(0.0, 1.0).expect("measure parameters are valid");
237        let quadrature = BayesianQuadrature::new(kernel, measure, 0.0);
238
239        let result = quadrature.posterior(&[0.0, 0.0], &[1.0, 1.0]);
240
241        assert!(matches!(
242            result,
243            Err(BayesianQuadratureError::Conditioning(
244                ConditioningError::NotPositiveDefinite
245            ))
246        ));
247    }
248
249    #[test]
250    fn duplicate_nodes_can_be_regularized_with_explicit_jitter() {
251        let kernel = RbfKernel::new(1.0, 1.0).expect("kernel parameters are valid");
252        let measure = GaussianMeasure::new(0.0, 1.0).expect("measure parameters are valid");
253        let quadrature = BayesianQuadrature::new(kernel, measure, 1.0e-8);
254
255        let posterior = quadrature
256            .posterior(&[0.0, 0.0], &[1.0, 1.0])
257            .expect("jitter regularizes duplicate nodes");
258
259        assert!(posterior.mean().is_finite());
260        assert!(posterior.variance() >= 0.0);
261    }
262
263    #[test]
264    fn tiny_negative_variance_is_clamped_to_zero() {
265        let scale = 2.0;
266        let tiny_negative = -32.0 * f64::EPSILON * scale;
267
268        assert_eq!(
269            non_negative_roundoff_variance(tiny_negative, scale),
270            Ok(0.0)
271        );
272    }
273
274    #[test]
275    fn material_negative_variance_is_rejected() {
276        let result = non_negative_roundoff_variance(-1.0e-6, 1.0);
277
278        assert!(matches!(
279            result,
280            Err(BayesianQuadratureError::MateriallyNegativePosteriorVariance { .. })
281        ));
282    }
283}