uncertain_numerics/
kernel_mean.rs1use crate::{GaussianMeasure, RbfKernel};
4
5pub trait KernelMean<M> {
17 #[must_use]
19 fn kernel_mean(&self, measure: &M, x: f64) -> f64;
20}
21
22impl KernelMean<GaussianMeasure> for RbfKernel {
23 fn kernel_mean(&self, measure: &GaussianMeasure, x: f64) -> f64 {
24 let length_scale_squared = self.length_scale() * self.length_scale();
25 let combined_variance = length_scale_squared + measure.variance();
26 let centered = x - measure.mean();
27 let scale = self.signal_variance() * (length_scale_squared / combined_variance).sqrt();
28 let exponent = -(centered * centered) / (2.0 * combined_variance);
29
30 scale * exponent.exp()
31 }
32}
33
34#[cfg(test)]
35mod tests {
36 use super::KernelMean;
37 use crate::{ContinuousProbabilityMeasure, GaussianMeasure, RbfKernel, ScalarKernel};
38
39 const TOLERANCE: f64 = 1.0e-12;
40 const QUADRATURE_TOLERANCE: f64 = 1.0e-9;
41
42 fn assert_close(actual: f64, expected: f64, tolerance: f64) {
43 let scale = expected.abs().max(1.0);
44 assert!(
45 (actual - expected).abs() <= tolerance * scale,
46 "expected {expected:.16e}, got {actual:.16e}"
47 );
48 }
49
50 fn simpson_integral<F>(function: F, lower: f64, upper: f64, intervals: u32) -> f64
51 where
52 F: Fn(f64) -> f64,
53 {
54 assert!(intervals > 0);
55 assert_eq!(intervals % 2, 0);
56
57 let step = (upper - lower) / f64::from(intervals);
58 let mut weighted_sum = function(lower) + function(upper);
59
60 for index in 1..intervals {
61 let x = lower + f64::from(index) * step;
62 let weight = if index % 2 == 0 { 2.0 } else { 4.0 };
63 weighted_sum += weight * function(x);
64 }
65
66 weighted_sum * step / 3.0
67 }
68
69 #[test]
70 fn kernel_mean_at_measure_mean_has_closed_form_scale() {
71 let kernel = RbfKernel::new(2.5, 0.75).expect("kernel parameters are valid");
72 let measure = GaussianMeasure::new(-1.25, 1.6).expect("measure parameters are valid");
73 let length_scale_squared = kernel.length_scale() * kernel.length_scale();
74 let expected = kernel.signal_variance()
75 * (length_scale_squared / (length_scale_squared + measure.variance())).sqrt();
76
77 assert_close(
78 kernel.kernel_mean(&measure, measure.mean()),
79 expected,
80 TOLERANCE,
81 );
82 }
83
84 #[test]
85 fn kernel_mean_is_symmetric_around_measure_mean() {
86 let kernel = RbfKernel::new(1.7, 0.9).expect("kernel parameters are valid");
87 let measure = GaussianMeasure::new(2.0, 1.3).expect("measure parameters are valid");
88
89 assert_close(
90 kernel.kernel_mean(&measure, 0.75),
91 kernel.kernel_mean(&measure, 3.25),
92 TOLERANCE,
93 );
94 }
95
96 #[test]
97 fn kernel_mean_is_translation_invariant() {
98 let kernel = RbfKernel::new(1.2, 1.8).expect("kernel parameters are valid");
99 let original = GaussianMeasure::new(-0.5, 0.7).expect("measure parameters are valid");
100 let shifted = GaussianMeasure::new(9.5, 0.7).expect("measure parameters are valid");
101
102 assert_close(
103 kernel.kernel_mean(&original, 1.25),
104 kernel.kernel_mean(&shifted, 11.25),
105 TOLERANCE,
106 );
107 }
108
109 #[test]
110 fn kernel_mean_is_positive_and_bounded_by_signal_variance() {
111 let kernel = RbfKernel::new(3.0, 0.6).expect("kernel parameters are valid");
112 let measure = GaussianMeasure::new(0.0, 2.0).expect("measure parameters are valid");
113
114 for x in [-8.0, -2.0, 0.0, 1.5, 10.0] {
115 let mean = kernel.kernel_mean(&measure, x);
116 assert!(mean > 0.0);
117 assert!(mean <= kernel.signal_variance());
118 }
119 }
120
121 #[test]
122 fn analytic_kernel_mean_matches_independent_numerical_quadrature() {
123 let cases = [
124 (1.0, 1.0, 0.0, 1.0, 0.0),
125 (2.5, 0.4, -1.0, 2.0, 0.75),
126 (0.7, 3.0, 4.0, 0.25, 5.2),
127 (1.8, 0.8, 1.5, 3.5, -2.0),
128 ];
129
130 for (signal_variance, length_scale, measure_mean, measure_variance, x) in cases {
131 let kernel =
132 RbfKernel::new(signal_variance, length_scale).expect("kernel parameters are valid");
133 let measure = GaussianMeasure::new(measure_mean, measure_variance)
134 .expect("measure parameters are valid");
135 let standard_deviation = measure.standard_deviation();
136 let lower = measure.mean() - 10.0 * standard_deviation;
137 let upper = measure.mean() + 10.0 * standard_deviation;
138
139 let numerical = simpson_integral(
140 |integration_point| {
141 kernel.covariance(x, integration_point) * measure.density(integration_point)
142 },
143 lower,
144 upper,
145 20_000,
146 );
147 let analytic = kernel.kernel_mean(&measure, x);
148
149 assert_close(analytic, numerical, QUADRATURE_TOLERANCE);
150 }
151 }
152}