use crate::core::math::Scalar;
use super::{
BracketError, BracketTerminationReason, Sample, Search, evaluate,
result_accessors, search_builders,
};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct MinimumBracketResult<F: Scalar = f64> {
points: [Sample<F>; 3],
iterations: u64,
function_evals: u64,
reason: BracketTerminationReason,
}
impl<F: Scalar> MinimumBracketResult<F> {
pub fn bracket(&self) -> (F, F, F) {
(self.points[0].x, self.points[1].x, self.points[2].x)
}
pub fn values(&self) -> (F, F, F) {
(
self.points[0].value,
self.points[1].value,
self.points[2].value,
)
}
result_accessors!();
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct MinimumBracketer<F: Scalar = f64> {
points: [F; 3],
search: Search<F>,
}
impl<F: Scalar> MinimumBracketer<F> {
pub fn new(left: F, middle: F, right: F) -> Self {
Self {
points: [left, middle, right],
search: Search::default(),
}
}
search_builders!();
pub fn bracket<C, E>(
&self,
mut function: C,
) -> Result<MinimumBracketResult<F>, BracketError<E, F>>
where
C: FnMut(F) -> Result<F, E>,
{
self.search.validate(&self.points)?;
let mut evaluations = 0;
let [a, m, b] = self.points;
let mut points = [
evaluate(&mut function, a, &mut evaluations)?,
evaluate(&mut function, m, &mut evaluations)?,
evaluate(&mut function, b, &mut evaluations)?,
];
let limit = if points[0].value < points[2].value {
points.swap(0, 2);
self.search.lower
} else {
self.search.upper
};
let anchor = points[2].x;
let mut distance = points[2].x - points[1].x;
let mut iterations = 0;
let reason = loop {
let [a, m, b] = points.map(|p| p.value);
if m <= a && m <= b && (m < a || m < b) {
break BracketTerminationReason::Bracketed;
}
if limit == Some(points[2].x) {
break BracketTerminationReason::BoundsReached;
}
if iterations == self.search.max_iter {
break BracketTerminationReason::MaxIter;
}
let Some(x) =
self.search.next(points[2].x, anchor, &mut distance, limit)
else {
break BracketTerminationReason::NoProgress;
};
if !(x - points[1].x).is_finite() {
break BracketTerminationReason::NoProgress;
}
let next = evaluate(&mut function, x, &mut evaluations)?;
points = [points[1], points[2], next];
iterations += 1;
};
if points[0].x > points[2].x {
points.swap(0, 2);
}
Ok(MinimumBracketResult {
points,
iterations,
function_evals: evaluations,
reason,
})
}
}