1use crate::{
4 BayesianQuadratureError, GaussianConditioner, GaussianMeasure, KernelIntegral, KernelMean,
5 RbfKernel, ScalarKernel, ScalarNormalPosterior,
6};
7
8#[derive(Debug, Clone, Copy, PartialEq)]
13pub struct BayesianQuadrature {
14 kernel: RbfKernel,
15 measure: GaussianMeasure,
16 jitter: f64,
17}
18
19impl BayesianQuadrature {
20 #[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 #[must_use]
32 pub const fn kernel(&self) -> RbfKernel {
33 self.kernel
34 }
35
36 #[must_use]
38 pub const fn measure(&self) -> GaussianMeasure {
39 self.measure
40 }
41
42 #[must_use]
44 pub const fn jitter(&self) -> f64 {
45 self.jitter
46 }
47
48 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}