use scirs2_core::ndarray::{Array, Array1, Dimension, ScalarOperand};
use scirs2_core::numeric::Float;
use std::collections::VecDeque;
use std::fmt::Debug;
use crate::error::{OptimError, Result};
use crate::optimizers::Optimizer;
#[derive(Debug, Clone)]
pub struct LBFGS<A: Float + ScalarOperand + Debug> {
learning_rate: A,
history_size: usize,
tolerance_grad: A,
c1: A,
c2: A,
max_ls: usize,
ls_contraction: A,
old_dirs: VecDeque<Array1<A>>,
old_stps: VecDeque<Array1<A>>,
ro: VecDeque<A>,
prev_params: Option<Array1<A>>,
prev_grad: Option<Array1<A>>,
h_diag: A,
n_iter: usize,
alpha: Vec<A>,
}
impl<A: Float + ScalarOperand + Debug + Send + Sync> LBFGS<A> {
pub fn new(learning_rate: A) -> Self {
Self::new_with_config(
learning_rate,
100, A::from(1e-7).unwrap_or_else(A::epsilon), A::from(1e-4).unwrap_or_else(A::epsilon), A::from(0.9).unwrap_or_else(|| A::one()), 25, )
}
pub fn new_with_config(
learning_rate: A,
history_size: usize,
tolerance_grad: A,
c1: A,
c2: A,
max_ls: usize,
) -> Self {
let ls_contraction = A::from(0.5).unwrap_or_else(|| A::one() / (A::one() + A::one()));
Self {
learning_rate,
history_size,
tolerance_grad,
c1,
c2,
max_ls,
ls_contraction,
old_dirs: VecDeque::with_capacity(history_size),
old_stps: VecDeque::with_capacity(history_size),
ro: VecDeque::with_capacity(history_size),
prev_params: None,
prev_grad: None,
h_diag: A::one(),
n_iter: 0,
alpha: vec![A::zero(); history_size],
}
}
pub fn learning_rate(&self) -> A {
self.learning_rate
}
pub fn set_lr(&mut self, lr: A) {
self.learning_rate = lr;
}
pub fn c1(&self) -> A {
self.c1
}
pub fn c2(&self) -> A {
self.c2
}
pub fn max_ls(&self) -> usize {
self.max_ls
}
pub fn history_len(&self) -> usize {
self.old_stps.len()
}
pub fn last_curvature_pair(&self) -> Option<(&Array1<A>, &Array1<A>)> {
match (self.old_stps.back(), self.old_dirs.back()) {
(Some(s), Some(y)) => Some((s, y)),
_ => None,
}
}
pub fn initial_hessian_scale(&self) -> A {
self.h_diag
}
pub fn set_line_search_contraction(&mut self, rho: A) -> Result<()> {
if rho <= A::zero() || rho >= A::one() {
return Err(OptimError::InvalidConfig(
"line search contraction factor must lie in (0, 1)".to_string(),
));
}
self.ls_contraction = rho;
Ok(())
}
pub fn reset(&mut self) {
self.old_dirs.clear();
self.old_stps.clear();
self.ro.clear();
self.prev_params = None;
self.prev_grad = None;
self.h_diag = A::one();
self.n_iter = 0;
self.alpha.fill(A::zero());
}
fn compute_direction(&mut self, gradient: &Array1<A>) -> Array1<A> {
let num_old = self.old_dirs.len();
if num_old == 0 {
return gradient.mapv(|x| -x);
}
let mut q = gradient.mapv(|x| -x);
for i in (0..num_old).rev() {
self.alpha[i] = self.old_stps[i].dot(&q) * self.ro[i];
q = &q - &self.old_dirs[i] * self.alpha[i];
}
let mut r = q * self.h_diag;
for i in 0..num_old {
let beta = self.old_dirs[i].dot(&r) * self.ro[i];
r = &r + &self.old_stps[i] * (self.alpha[i] - beta);
}
r
}
fn update_history(&mut self, y: Array1<A>, s: Array1<A>) -> bool {
if self.history_size == 0 || y.len() != s.len() {
return false;
}
let ys = y.dot(&s);
let eps = A::from(1e-10).unwrap_or_else(A::epsilon);
let threshold = eps * s.dot(&s).sqrt() * y.dot(&y).sqrt();
if !(ys.is_finite() && ys > A::zero() && ys > threshold) {
return false;
}
while self.old_dirs.len() >= self.history_size {
self.old_dirs.pop_front();
self.old_stps.pop_front();
self.ro.pop_front();
}
let yy = y.dot(&y);
self.old_dirs.push_back(y);
self.old_stps.push_back(s);
self.ro.push_back(A::one() / ys);
if yy > A::zero() {
self.h_diag = ys / yy;
}
true
}
fn flatten<D: Dimension>(array: &Array<A, D>, what: &str) -> Result<Array1<A>> {
array
.to_owned()
.into_shape_with_order(array.len())
.map_err(|e| {
OptimError::DimensionMismatch(format!(
"failed to flatten {} of shape {:?}: {}",
what,
array.shape(),
e
))
})
}
fn unflatten<D: Dimension>(flat: Array1<A>, like: &Array<A, D>) -> Result<Array<A, D>> {
flat.into_shape_with_order(like.raw_dim()).map_err(|e| {
OptimError::DimensionMismatch(format!(
"failed to reshape update into {:?}: {}",
like.shape(),
e
))
})
}
fn prepare_direction(
&mut self,
params_flat: &Array1<A>,
gradients_flat: &Array1<A>,
) -> Array1<A> {
if let (Some(prev_params), Some(prev_grad)) = (&self.prev_params, &self.prev_grad) {
if prev_params.len() == params_flat.len() && prev_grad.len() == gradients_flat.len() {
let s = params_flat - prev_params;
let y = gradients_flat - prev_grad;
let _accepted = self.update_history(y, s);
}
}
self.compute_direction(gradients_flat)
}
pub fn step_with_loss<D, F>(
&mut self,
params: &Array<A, D>,
gradients: &Array<A, D>,
mut loss_fn: F,
) -> Result<Array<A, D>>
where
D: Dimension,
F: FnMut(&Array<A, D>) -> A,
{
if params.shape() != gradients.shape() {
return Err(OptimError::DimensionMismatch(format!(
"parameters have shape {:?} but gradients have shape {:?}",
params.shape(),
gradients.shape()
)));
}
let params_flat = Self::flatten(params, "parameters")?;
let gradients_flat = Self::flatten(gradients, "gradients")?;
let grad_norm = gradients_flat.dot(&gradients_flat).sqrt();
if grad_norm <= self.tolerance_grad {
self.prev_params = Some(params_flat);
self.prev_grad = Some(gradients_flat);
return Ok(params.clone());
}
let mut direction = self.prepare_direction(¶ms_flat, &gradients_flat);
let mut gtd = gradients_flat.dot(&direction);
if !matches!(gtd.partial_cmp(&A::zero()), Some(std::cmp::Ordering::Less)) {
direction = gradients_flat.mapv(|x| -x);
gtd = -gradients_flat.dot(&gradients_flat);
}
let f0 = loss_fn(params);
if !f0.is_finite() {
return Err(OptimError::InvalidConfig(
"objective is not finite at the current parameters".to_string(),
));
}
let mut alpha = self.learning_rate;
let mut best: Option<(A, A, Array<A, D>)> = None; let mut accepted: Option<(A, Array<A, D>)> = None;
for _ in 0..self.max_ls.max(1) {
let candidate_flat = ¶ms_flat + &(&direction * alpha);
let candidate = Self::unflatten(candidate_flat, params)?;
let f = loss_fn(&candidate);
if f.is_finite() && f <= f0 + self.c1 * alpha * gtd {
accepted = Some((alpha, candidate));
break;
}
if f.is_finite() && f < f0 {
let improves = match &best {
Some((best_f, _, _)) => f < *best_f,
None => true,
};
if improves {
best = Some((f, alpha, candidate));
}
}
alpha = alpha * self.ls_contraction;
if !matches!(
alpha.partial_cmp(&A::zero()),
Some(std::cmp::Ordering::Greater)
) {
break;
}
}
let (step_size, new_params) = match accepted {
Some((a, candidate)) => (a, candidate),
None => match best {
Some((_, a, candidate)) => (a, candidate),
None => (A::zero(), params.clone()),
},
};
self.prev_params = Some(params_flat);
self.prev_grad = Some(gradients_flat);
if step_size > A::zero() {
self.n_iter += 1;
}
Ok(new_params)
}
}
impl<A, D> Optimizer<A, D> for LBFGS<A>
where
A: Float + ScalarOperand + Debug + Send + Sync,
D: Dimension,
{
fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
if params.shape() != gradients.shape() {
return Err(OptimError::DimensionMismatch(format!(
"parameters have shape {:?} but gradients have shape {:?}",
params.shape(),
gradients.shape()
)));
}
let params_flat = Self::flatten(params, "parameters")?;
let gradients_flat = Self::flatten(gradients, "gradients")?;
let grad_norm = gradients_flat.dot(&gradients_flat).sqrt();
if grad_norm <= self.tolerance_grad {
self.prev_params = Some(params_flat);
self.prev_grad = Some(gradients_flat);
return Ok(params.clone());
}
let direction = self.prepare_direction(¶ms_flat, &gradients_flat);
let step_size = if self.old_stps.is_empty() {
self.learning_rate / (A::one() + grad_norm)
} else {
self.learning_rate
};
let new_params_flat = ¶ms_flat + &(&direction * step_size);
self.prev_params = Some(params_flat);
self.prev_grad = Some(gradients_flat);
self.n_iter += 1;
Self::unflatten(new_params_flat, params)
}
fn get_learning_rate(&self) -> A {
self.learning_rate
}
fn set_learning_rate(&mut self, learning_rate: A) {
self.learning_rate = learning_rate;
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_abs_diff_eq;
use scirs2_core::ndarray::Array1;
#[test]
fn test_lbfgs_basic_creation() {
let optimizer: LBFGS<f64> = LBFGS::new(1.0);
assert_abs_diff_eq!(optimizer.learning_rate(), 1.0);
assert_eq!(optimizer.history_size, 100);
assert_abs_diff_eq!(optimizer.tolerance_grad, 1e-7);
}
#[test]
fn test_lbfgs_convergence() {
let mut optimizer: LBFGS<f64> = LBFGS::new(0.1);
let mut params = Array1::from_vec(vec![10.0]);
for _ in 0..50 {
let gradients = Array1::from_vec(vec![2.0 * params[0]]);
params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_lbfgs_convergence");
}
assert!(params[0].abs() < 0.1);
}
#[test]
fn test_lbfgs_2d() {
let mut optimizer: LBFGS<f64> = LBFGS::new(0.1);
let mut params = Array1::from_vec(vec![5.0, 3.0]);
for _ in 0..50 {
let gradients = Array1::from_vec(vec![2.0 * params[0], 2.0 * params[1]]);
params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_lbfgs_2d");
}
assert!(params[0].abs() < 0.1);
assert!(params[1].abs() < 0.1);
}
#[test]
fn test_lbfgs_reset() {
let mut optimizer: LBFGS<f64> = LBFGS::new(0.1);
let mut params = Array1::from_vec(vec![1.0]);
let gradients = Array1::from_vec(vec![2.0]);
params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_lbfgs_reset");
let gradients2 = Array1::from_vec(vec![1.5]);
params = optimizer
.step(¶ms, &gradients2)
.expect("optimizer.step succeeds in test_lbfgs_reset");
let gradients3 = Array1::from_vec(vec![1.0]);
let _ = optimizer
.step(¶ms, &gradients3)
.expect("optimizer.step succeeds in test_lbfgs_reset");
assert!(!optimizer.old_dirs.is_empty());
assert!(optimizer.n_iter > 0);
optimizer.reset();
assert!(optimizer.old_dirs.is_empty());
assert!(optimizer.old_stps.is_empty());
assert!(optimizer.ro.is_empty());
assert!(optimizer.prev_grad.is_none());
assert_eq!(optimizer.n_iter, 0);
}
}