Skip to main content

uncertain_numerics/
kernel_mean.rs

1//! Analytic kernel means for supported kernel and probability-measure pairs.
2
3use crate::{GaussianMeasure, RbfKernel};
4
5/// Analytic mean embedding of a scalar kernel under a probability measure.
6///
7/// For a kernel `k` and probability measure `p`, the kernel mean at `x` is
8///
9/// ```text
10/// z(x) = integral k(x, u) p(u) du.
11/// ```
12///
13/// The trait is implemented only for kernel/measure pairs with an explicit
14/// analytic contract. Unsupported pairs therefore fail at compile time rather
15/// than silently falling back to an unrelated numerical approximation.
16pub trait KernelMean<M> {
17    /// Evaluate the analytic kernel mean at `x`.
18    #[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}