use crate::core::math::Scalar;
use super::{
BracketError, BracketTerminationReason, Sample, Search, evaluate,
result_accessors, search_builders,
};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RootBracketResult<F: Scalar = f64> {
points: [Sample<F>; 2],
iterations: u64,
function_evals: u64,
reason: BracketTerminationReason,
}
impl<F: Scalar> RootBracketResult<F> {
pub fn bracket(&self) -> (F, F) {
(self.points[0].x, self.points[1].x)
}
pub fn values(&self) -> (F, F) {
(self.points[0].value, self.points[1].value)
}
result_accessors!();
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RootBracketer<F: Scalar = f64> {
lower: F,
upper: F,
search: Search<F>,
}
impl<F: Scalar> RootBracketer<F> {
pub fn new(lower: F, upper: F) -> Self {
Self {
lower,
upper,
search: Search::default(),
}
}
search_builders!();
pub fn bracket<C, E>(
&self,
mut function: C,
) -> Result<RootBracketResult<F>, BracketError<E, F>>
where
C: FnMut(F) -> Result<F, E>,
{
self.search.validate(&[self.lower, self.upper])?;
let mut evaluations = 0;
let mut points = [
evaluate(&mut function, self.lower, &mut evaluations)?,
evaluate(&mut function, self.upper, &mut evaluations)?,
];
let mut iterations = 0;
let mut distances = [self.lower - self.upper, self.upper - self.lower];
let anchors = [self.upper, self.lower];
let limits = [self.search.lower, self.search.upper];
let mut stopped = [false; 2];
let reason = if brackets(points) {
BracketTerminationReason::Bracketed
} else {
loop {
for i in 0..2 {
if limits[i] == Some(points[i].x) {
stopped[i] = true;
}
}
if stopped.iter().all(|&s| s) {
break if limits[0] == Some(points[0].x)
&& limits[1] == Some(points[1].x)
{
BracketTerminationReason::BoundsReached
} else {
BracketTerminationReason::NoProgress
};
}
if iterations == self.search.max_iter {
break BracketTerminationReason::MaxIter;
}
let mut found: Option<[Sample<F>; 2]> = None;
let before = evaluations;
for i in 0..2 {
if stopped[i] {
continue;
}
let next = self.search.next(
points[i].x,
anchors[i],
&mut distances[i],
limits[i],
);
let Some(x) = next else {
stopped[i] = true;
continue;
};
let sample = evaluate(&mut function, x, &mut evaluations)?;
let pair = if i == 0 {
[sample, points[i]]
} else {
[points[i], sample]
};
if brackets(pair)
&& found.is_none_or(|old| {
pair[1].x - pair[0].x < old[1].x - old[0].x
})
{
found = Some(pair);
}
points[i] = sample;
}
iterations += u64::from(evaluations != before);
if let Some(pair) = found {
points = pair;
break BracketTerminationReason::Bracketed;
}
}
};
Ok(RootBracketResult {
points,
iterations,
function_evals: evaluations,
reason,
})
}
}
fn brackets<F: Scalar>(points: [Sample<F>; 2]) -> bool {
let [a, b] = points;
a.value == F::zero()
|| b.value == F::zero()
|| (a.value > F::zero()) != (b.value > F::zero())
}