use crate::core::*;
#[derive(Debug, Copy, Clone, PartialEq)]
pub enum LineSearchAlgorithm {
MoreThuente,
BacktrackingArmijo,
BacktrackingStrongWolfe,
BacktrackingWolfe,
}
impl Default for LineSearchAlgorithm {
fn default() -> Self {
LineSearchAlgorithm::MoreThuente
}
}
#[derive(Debug, Copy, Clone)]
pub struct LineSearch {
pub algorithm: LineSearchAlgorithm,
pub ftol: f64,
pub gtol: f64,
pub xtol: f64,
pub min_step: f64,
pub max_step: f64,
pub max_linesearch: usize,
pub gradient_only: bool,
}
impl Default for LineSearch {
fn default() -> Self {
LineSearch {
ftol: 1e-4,
gtol: 0.9,
xtol: 1e-16,
min_step: 1e-20,
max_step: 1e20,
max_linesearch: 40,
gradient_only: false,
algorithm: LineSearchAlgorithm::default(),
}
}
}
impl LineSearch {
fn validate_step(&self, step: f64) -> Result<()> {
if step < self.min_step {
bail!("The line-search step became smaller than LineSearch::min_step.");
}
if step > self.max_step {
bail!("The line-search step became larger than LineSearch::max_step.");
}
Ok(())
}
}
use crate::lbfgs::Problem;
impl LineSearch {
pub fn find<E>(&self, prb: &mut Problem<E>, step: &mut f64) -> Result<usize>
where
E: FnMut(&[f64], &mut [f64]) -> Result<f64>,
{
ensure!(
step.is_sign_positive(),
"A logic error (negative line-search step) occurred."
);
let ls = if self.algorithm == MoreThuente && !prb.orthantwise() {
if !self.gradient_only {
line_search_morethuente(prb, step, &self)
} else {
bail!("Gradient only optimization is incompatible with MoreThuente line search.");
}
} else {
line_search_backtracking(prb, step, &self)
}
.unwrap_or_else(|err| {
error!("line search failed, revert to the previous point!");
prb.revert();
println!("{:?}", err);
0
});
Ok(ls)
}
}
use crate::math::*;
pub fn line_search_morethuente<E>(
prb: &mut Problem<E>,
stp: &mut f64, param: &LineSearch, ) -> Result<usize>
where
E: FnMut(&[f64], &mut [f64]) -> Result<f64>,
{
let dginit = prb.dginit()?;
let mut brackt = false;
let mut stage1 = 1;
let mut uinfo = 0;
let finit = prb.fx;
let dgtest = param.ftol * dginit;
let mut width = param.max_step - param.min_step;
let mut prev_width = 2.0 * width;
let (mut stx, mut sty) = (0.0, 0.0);
let mut fx = finit;
let mut fy = finit;
let mut dgy = dginit;
let mut dgx = dgy;
for count in 0..param.max_linesearch {
let (stmin, stmax) = if brackt {
(if stx <= sty { stx } else { sty }, if stx >= sty { stx } else { sty })
} else {
(stx, *stp + 4.0 * (*stp - stx))
};
if *stp < param.min_step {
*stp = param.min_step
}
if param.max_step < *stp {
*stp = param.max_step
}
if brackt && (*stp <= stmin || stmax <= *stp || param.max_linesearch <= count + 1 || uinfo != 0)
|| brackt && stmax - stmin <= param.xtol * stmax
{
*stp = stx
}
prb.take_line_step(*stp);
prb.evaluate()?;
let f = prb.fx;
let dg = prb.dg_unchecked();
let ftest1 = finit + *stp * dgtest;
if brackt && (*stp <= stmin || stmax <= *stp || uinfo != 0i32) {
bail!(
"A rounding error occurred; alternatively, no line-search step
satisfies the sufficient decrease and curvature conditions."
);
}
if brackt && stmax - stmin <= param.xtol * stmax {
bail!("Relative width of the interval of uncertainty is at most xtol.");
}
if *stp == param.max_step && f <= ftest1 && dg <= dgtest {
bail!("The line-search step became larger than LineSearch::max_step.");
}
if *stp == param.min_step && (ftest1 < f || dgtest <= dg) {
bail!("The line-search step became smaller than LineSearch::min_step.");
}
if dg.abs() <= param.gtol * -dginit {
return Ok(count);
} else if f <= ftest1 && dg.abs() <= param.gtol * -dginit {
return Ok(count);
} else {
if 0 != stage1 && f <= ftest1 && param.ftol.min(param.gtol) * dginit <= dg {
stage1 = 0;
}
if 0 != stage1 && ftest1 < f && f <= fx {
let fm = f - *stp * dgtest;
let mut fxm = fx - stx * dgtest;
let mut fym = fy - sty * dgtest;
let dgm = dg - dgtest;
let mut dgxm = dgx - dgtest;
let mut dgym = dgy - dgtest;
uinfo = mcstep::update_trial_interval(
&mut stx,
&mut fxm,
&mut dgxm,
&mut sty,
&mut fym,
&mut dgym,
&mut *stp,
fm,
dgm,
stmin,
stmax,
&mut brackt,
)?;
fx = fxm + stx * dgtest;
fy = fym + sty * dgtest;
dgx = dgxm + dgtest;
dgy = dgym + dgtest
} else {
uinfo = mcstep::update_trial_interval(
&mut stx,
&mut fx,
&mut dgx,
&mut sty,
&mut fy,
&mut dgy,
&mut *stp,
f,
dg,
stmin,
stmax,
&mut brackt,
)?;
}
if !(brackt) {
continue;
}
if 0.66 * prev_width <= (sty - stx).abs() {
*stp = stx + 0.5 * (sty - stx)
}
prev_width = width;
width = (sty - stx).abs()
}
}
info!("The line-search routine reaches the maximum number of evaluations.");
Ok(param.max_linesearch)
}
mod mcstep {
use super::{cubic_minimizer, cubic_minimizer2, quard_minimizer, quard_minimizer2};
use crate::core::*;
pub(crate) fn update_trial_interval(
x: &mut f64,
fx: &mut f64,
dx: &mut f64,
y: &mut f64,
fy: &mut f64,
dy: &mut f64,
t: &mut f64,
ft: f64,
dt: f64,
tmin: f64,
tmax: f64,
brackt: &mut bool,
) -> Result<i32> {
let dsign = dt * (*dx / (*dx).abs()) < 0.0;
let mut mc = 0.;
let mut mq = 0.;
let mut newt = 0.;
if *brackt {
if *t <= x.min(*y) || x.max(*y) <= *t {
bail!("The line-search step went out of the interval of uncertainty.");
} else if 0.0 <= *dx * (*t - *x) {
bail!("The current search direction increases the objective function value.");
} else if tmax < tmin {
bail!("A logic error occurred; alternatively, the interval of uncertainty became too small.");
}
}
let bound = if *fx < ft {
*brackt = true;
cubic_minimizer(&mut mc, *x, *fx, *dx, *t, ft, dt);
quard_minimizer(&mut mq, *x, *fx, *dx, *t, ft);
if (mc - *x).abs() < (mq - *x).abs() {
newt = mc
} else {
newt = mc + 0.5 * (mq - mc)
}
1
} else if dsign {
*brackt = true;
cubic_minimizer(&mut mc, *x, *fx, *dx, *t, ft, dt);
quard_minimizer2(&mut mq, *x, *dx, *t, dt);
if (mc - *t).abs() > (mq - *t).abs() {
newt = mc
} else {
newt = mq
}
0
} else if dt.abs() < (*dx).abs() {
cubic_minimizer2(&mut mc, *x, *fx, *dx, *t, ft, dt, tmin, tmax);
quard_minimizer2(&mut mq, *x, *dx, *t, dt);
if *brackt {
if (*t - mc).abs() < (*t - mq).abs() {
newt = mc
} else {
newt = mq
}
} else if (*t - mc).abs() > (*t - mq).abs() {
newt = mc
} else {
newt = mq
}
1
} else {
if *brackt {
cubic_minimizer(&mut newt, *t, ft, dt, *y, *fy, *dy);
} else if *x < *t {
newt = tmax
} else {
newt = tmin
}
0
};
if *fx < ft {
*y = *t;
*fy = ft;
*dy = dt
} else {
if dsign {
*y = *x;
*fy = *fx;
*dy = *dx
}
*x = *t;
*fx = ft;
*dx = dt
}
if tmax < newt {
newt = tmax
}
if newt < tmin {
newt = tmin
}
if *brackt && 0 != bound {
mq = *x + 0.66 * (*y - *x);
if *x < *y {
if mq < newt {
newt = mq
}
} else if newt < mq {
newt = mq
}
}
*t = newt;
Ok(0)
}
}
#[inline]
fn cubic_minimizer(cm: &mut f64, u: f64, fu: f64, du: f64, v: f64, fv: f64, dv: f64) {
let d = v - u;
let theta = (fu - fv) * 3.0 / d + du + dv;
let mut p = theta.abs();
let mut q = du.abs();
let mut r = dv.abs();
let s = (p.max(q)).max(r); let a = theta / s;
let mut gamma = s * (a * a - du / s * (dv / s)).sqrt();
if v < u {
gamma = -gamma
}
p = gamma - du + theta;
q = gamma - du + gamma + dv;
r = p / q;
*cm = u + r * d;
}
#[inline]
fn cubic_minimizer2(cm: &mut f64, u: f64, fu: f64, du: f64, v: f64, fv: f64, dv: f64, xmin: f64, xmax: f64) {
let d = v - u;
let theta = (fu - fv) * 3.0 / d + du + dv;
let mut p = theta.abs();
let mut q = du.abs();
let mut r = dv.abs();
let s = (p.max(q)).max(r); let a = theta / s;
let mut gamma = s * (0f64.max(a * a - du / s * (dv / s)).sqrt());
if u < v {
gamma = -gamma
}
p = gamma - dv + theta;
q = gamma - dv + gamma + du;
r = p / q;
if r < 0.0 && gamma != 0.0 {
*cm = v - r * d;
} else if v > u {
*cm = xmax;
} else {
*cm = xmin;
}
}
#[inline]
fn quard_minimizer(qm: &mut f64, u: f64, fu: f64, du: f64, v: f64, fv: f64) {
let a = v - u;
*qm = u + du / ((fu - fv) / a + du) / 2.0 * a;
}
#[inline]
fn quard_minimizer2(qm: &mut f64, u: f64, du: f64, v: f64, dv: f64) {
let a = u - v;
*qm = v + dv / (dv - du) * a;
}
use self::LineSearchAlgorithm::*;
pub fn line_search_backtracking<E>(
prb: &mut Problem<E>,
stp: &mut f64, param: &LineSearch, ) -> Result<usize>
where
E: FnMut(&[f64], &mut [f64]) -> Result<f64>,
{
let dginit = prb.dginit()?;
let dec: f64 = 0.5;
let inc: f64 = 2.1;
let orthantwise = prb.orthantwise();
let finit = prb.fx;
let dgtest = param.ftol * dginit;
let mut width: f64;
for count in 0..param.max_linesearch {
prb.take_line_step(*stp);
prb.evaluate()?;
if prb.fx > finit + *stp * dgtest {
width = dec;
} else if param.algorithm == BacktrackingArmijo || orthantwise {
return Ok(count);
} else {
let dg = prb.dg_unchecked();
if dg < param.gtol * dginit {
width = inc
} else if param.algorithm == BacktrackingWolfe {
return Ok(count);
} else if dg > -param.gtol * dginit {
width = dec
} else {
return Ok(count);
}
}
if param.gradient_only {
info!("allow energy rises");
let dg = prb.dg_unchecked();
if dg.abs() <= -param.gtol * dginit.abs() {
return Ok(count);
}
}
param.validate_step(*stp)?;
*stp *= width
}
info!("The line-search routine reaches the maximum number of evaluations.");
Ok(param.max_linesearch)
}