Skip to main content

uncertain_numerics/
active.rs

1//! Active Bayesian quadrature acquisition based on posterior integral-variance reduction.
2
3use crate::{
4    ActiveSelectionError, BayesianQuadratureError, GaussianConditioner, GaussianMeasure,
5    KernelMean, RbfKernel, ScalarKernel,
6};
7
8/// Candidate selected by posterior integral-variance reduction.
9#[derive(Debug, Clone, Copy, PartialEq)]
10pub struct SelectedCandidate {
11    point: f64,
12    variance_reduction: f64,
13    index: usize,
14}
15
16impl SelectedCandidate {
17    /// Return the selected candidate location.
18    #[must_use]
19    pub const fn point(self) -> f64 {
20        self.point
21    }
22
23    /// Return the predicted posterior integral-variance reduction.
24    #[must_use]
25    pub const fn variance_reduction(self) -> f64 {
26        self.variance_reduction
27    }
28
29    /// Return the selected candidate index in the supplied candidate slice.
30    #[must_use]
31    pub const fn index(self) -> usize {
32        self.index
33    }
34}
35
36/// Expected reduction in posterior integral variance from evaluating one candidate node.
37#[derive(Debug, Clone, Copy, PartialEq)]
38pub struct VarianceReductionAcquisition {
39    kernel: RbfKernel,
40    measure: GaussianMeasure,
41    jitter: f64,
42}
43
44impl VarianceReductionAcquisition {
45    /// Construct a variance-reduction acquisition function.
46    #[must_use]
47    pub const fn new(kernel: RbfKernel, measure: GaussianMeasure, jitter: f64) -> Self {
48        Self {
49            kernel,
50            measure,
51            jitter,
52        }
53    }
54
55    /// Return the configured kernel.
56    #[must_use]
57    pub const fn kernel(&self) -> RbfKernel {
58        self.kernel
59    }
60
61    /// Return the configured integration measure.
62    #[must_use]
63    pub const fn measure(&self) -> GaussianMeasure {
64        self.measure
65    }
66
67    /// Return the explicit diagonal jitter.
68    #[must_use]
69    pub const fn jitter(&self) -> f64 {
70        self.jitter
71    }
72
73    /// Evaluate the posterior integral-variance reduction at `candidate`.
74    ///
75    /// # Errors
76    ///
77    /// Returns [`BayesianQuadratureError`] for invalid current nodes, a non-finite
78    /// candidate, or a Gaussian-conditioning failure on the current design.
79    pub fn reduction(&self, nodes: &[f64], candidate: f64) -> Result<f64, BayesianQuadratureError> {
80        validate_nodes(nodes)?;
81        if !candidate.is_finite() {
82            return Err(BayesianQuadratureError::NonFiniteObservationNode);
83        }
84
85        let (conditioner, solved_mean) = self.prepare(nodes)?;
86        self.reduction_with_prepared(nodes, candidate, &conditioner, &solved_mean)
87    }
88
89    /// Select the candidate with the largest predicted variance reduction.
90    ///
91    /// Ties are resolved deterministically by keeping the first maximum in the
92    /// supplied candidate slice.
93    ///
94    /// # Errors
95    ///
96    /// Returns [`ActiveSelectionError`] for empty/non-finite candidate sets or
97    /// when acquisition evaluation fails for the current design.
98    pub fn select_best(
99        &self,
100        nodes: &[f64],
101        candidates: &[f64],
102    ) -> Result<SelectedCandidate, ActiveSelectionError> {
103        if candidates.is_empty() {
104            return Err(ActiveSelectionError::EmptyCandidates);
105        }
106        if candidates.iter().any(|candidate| !candidate.is_finite()) {
107            return Err(ActiveSelectionError::NonFiniteCandidate);
108        }
109
110        validate_nodes(nodes)?;
111        let (conditioner, solved_mean) = self.prepare(nodes)?;
112
113        let mut best = SelectedCandidate {
114            point: candidates[0],
115            variance_reduction: self.reduction_with_prepared(
116                nodes,
117                candidates[0],
118                &conditioner,
119                &solved_mean,
120            )?,
121            index: 0,
122        };
123
124        for (index, &candidate) in candidates.iter().enumerate().skip(1) {
125            let reduction =
126                self.reduction_with_prepared(nodes, candidate, &conditioner, &solved_mean)?;
127            if reduction > best.variance_reduction {
128                best = SelectedCandidate {
129                    point: candidate,
130                    variance_reduction: reduction,
131                    index,
132                };
133            }
134        }
135
136        Ok(best)
137    }
138
139    fn prepare(
140        &self,
141        nodes: &[f64],
142    ) -> Result<(GaussianConditioner, Vec<f64>), BayesianQuadratureError> {
143        let dimension = nodes.len();
144        let mut gram = Vec::with_capacity(dimension * dimension);
145        for &left in nodes {
146            for &right in nodes {
147                gram.push(self.kernel.covariance(left, right));
148            }
149        }
150
151        let kernel_mean: Vec<f64> = nodes
152            .iter()
153            .map(|&node| self.kernel.kernel_mean(&self.measure, node))
154            .collect();
155        let conditioner = GaussianConditioner::new(&gram, dimension, self.jitter)?;
156        let solved_mean = conditioner.solve(&kernel_mean)?;
157        Ok((conditioner, solved_mean))
158    }
159
160    fn reduction_with_prepared(
161        &self,
162        nodes: &[f64],
163        candidate: f64,
164        conditioner: &GaussianConditioner,
165        solved_mean: &[f64],
166    ) -> Result<f64, BayesianQuadratureError> {
167        let candidate_covariance: Vec<f64> = nodes
168            .iter()
169            .map(|&node| self.kernel.covariance(node, candidate))
170            .collect();
171        let solved_candidate = conditioner.solve(&candidate_covariance)?;
172
173        let posterior_integral_covariance = self.kernel.kernel_mean(&self.measure, candidate)
174            - dot(&candidate_covariance, solved_mean);
175        let predictive_variance = self.kernel.covariance(candidate, candidate) + self.jitter
176            - dot(&candidate_covariance, &solved_candidate);
177
178        let tolerance = 64.0 * f64::EPSILON * self.kernel.signal_variance().max(1.0);
179        if predictive_variance <= tolerance {
180            return Ok(0.0);
181        }
182
183        Ok(posterior_integral_covariance * posterior_integral_covariance / predictive_variance)
184    }
185}
186
187fn validate_nodes(nodes: &[f64]) -> Result<(), BayesianQuadratureError> {
188    if nodes.is_empty() {
189        return Err(BayesianQuadratureError::EmptyObservations);
190    }
191    if nodes.iter().any(|value| !value.is_finite()) {
192        return Err(BayesianQuadratureError::NonFiniteObservationNode);
193    }
194    Ok(())
195}
196
197fn dot(left: &[f64], right: &[f64]) -> f64 {
198    debug_assert_eq!(left.len(), right.len());
199    left.iter().zip(right).map(|(x, y)| x * y).sum()
200}
201
202#[cfg(test)]
203#[allow(clippy::float_cmp)] // exact round-trips of constructor inputs are intended
204mod tests {
205    use super::VarianceReductionAcquisition;
206    use crate::{ActiveSelectionError, BayesianQuadrature, GaussianMeasure, RbfKernel};
207
208    const TOLERANCE: f64 = 1.0e-11;
209
210    fn fixture() -> VarianceReductionAcquisition {
211        let kernel = RbfKernel::new(1.0, 1.1).expect("kernel parameters are valid");
212        let measure = GaussianMeasure::new(0.2, 1.0).expect("measure parameters are valid");
213        VarianceReductionAcquisition::new(kernel, measure, 1.0e-12)
214    }
215
216    #[test]
217    fn reduction_is_non_negative() {
218        let acquisition = fixture();
219        let nodes = [-1.0, 0.0, 1.0];
220        for candidate in [-2.0, -0.4, 0.5, 1.7] {
221            assert!(acquisition.reduction(&nodes, candidate).expect("valid") >= 0.0);
222        }
223    }
224
225    #[test]
226    fn existing_node_has_negligible_reduction() {
227        let acquisition = fixture();
228        let reduction = acquisition
229            .reduction(&[-1.0, 0.0, 1.0], 0.0)
230            .expect("valid");
231        assert!(reduction <= 1.0e-10);
232    }
233
234    #[test]
235    fn criterion_matches_actual_posterior_variance_drop() {
236        let acquisition = fixture();
237        let nodes = [-1.25, -0.1, 1.1];
238        let values = [0.4, -0.3, 0.8];
239        let candidate = 0.55;
240        let quadrature = BayesianQuadrature::new(
241            acquisition.kernel(),
242            acquisition.measure(),
243            acquisition.jitter(),
244        );
245        let before = quadrature.posterior(&nodes, &values).expect("valid");
246        let mut augmented_nodes = nodes.to_vec();
247        augmented_nodes.push(candidate);
248        let mut augmented_values = values.to_vec();
249        augmented_values.push(-0.2);
250        let after = quadrature
251            .posterior(&augmented_nodes, &augmented_values)
252            .expect("valid");
253        let predicted_drop = acquisition.reduction(&nodes, candidate).expect("valid");
254        assert!((predicted_drop - (before.variance() - after.variance())).abs() <= TOLERANCE);
255    }
256
257    #[test]
258    fn selector_returns_global_candidate_maximum() {
259        let acquisition = fixture();
260        let nodes = [-1.0, 0.0, 1.0];
261        let candidates = [-2.0, -0.6, 0.4, 1.8];
262        let selected = acquisition.select_best(&nodes, &candidates).expect("valid");
263        for &candidate in &candidates {
264            let reduction = acquisition.reduction(&nodes, candidate).expect("valid");
265            assert!(selected.variance_reduction() + TOLERANCE >= reduction);
266        }
267        assert_eq!(selected.point(), candidates[selected.index()]);
268    }
269
270    #[test]
271    fn selector_uses_first_maximum_for_ties() {
272        let acquisition = fixture();
273        let nodes = [-1.0, 0.0, 1.0];
274        let candidates = [0.0, 0.0, 0.0];
275        let selected = acquisition.select_best(&nodes, &candidates).expect("valid");
276        assert_eq!(selected.index(), 0);
277    }
278
279    #[test]
280    fn selector_rejects_invalid_candidate_sets() {
281        let acquisition = fixture();
282        assert_eq!(
283            acquisition.select_best(&[-1.0, 0.0, 1.0], &[]),
284            Err(ActiveSelectionError::EmptyCandidates)
285        );
286        assert_eq!(
287            acquisition.select_best(&[-1.0, 0.0, 1.0], &[0.5, f64::NAN]),
288            Err(ActiveSelectionError::NonFiniteCandidate)
289        );
290    }
291}