1use crate::{
4 ActiveSelectionError, BayesianQuadratureError, GaussianConditioner, GaussianMeasure,
5 KernelMean, RbfKernel, ScalarKernel,
6};
7
8#[derive(Debug, Clone, Copy, PartialEq)]
10pub struct SelectedCandidate {
11 point: f64,
12 variance_reduction: f64,
13 index: usize,
14}
15
16impl SelectedCandidate {
17 #[must_use]
19 pub const fn point(self) -> f64 {
20 self.point
21 }
22
23 #[must_use]
25 pub const fn variance_reduction(self) -> f64 {
26 self.variance_reduction
27 }
28
29 #[must_use]
31 pub const fn index(self) -> usize {
32 self.index
33 }
34}
35
36#[derive(Debug, Clone, Copy, PartialEq)]
38pub struct VarianceReductionAcquisition {
39 kernel: RbfKernel,
40 measure: GaussianMeasure,
41 jitter: f64,
42}
43
44impl VarianceReductionAcquisition {
45 #[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 #[must_use]
57 pub const fn kernel(&self) -> RbfKernel {
58 self.kernel
59 }
60
61 #[must_use]
63 pub const fn measure(&self) -> GaussianMeasure {
64 self.measure
65 }
66
67 #[must_use]
69 pub const fn jitter(&self) -> f64 {
70 self.jitter
71 }
72
73 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 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)] mod 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}