use crate::error::{OptimError, Result};
use scirs2_core::ndarray::ScalarOperand;
use scirs2_core::ndarray_ext::{Array1, ArrayView1};
use scirs2_core::numeric::Float;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NewtonCG<T: Float> {
learning_rate: T,
cg_tolerance: T,
cg_max_iters: usize,
hessian_reg: T,
#[serde(default = "no_trust_region")]
trust_region_radius: Option<T>,
#[serde(default = "default_max_trust_region_radius")]
max_trust_region_radius: T,
#[serde(default = "default_min_trust_region_radius")]
min_trust_region_radius: T,
#[serde(default = "default_trust_region_eta")]
trust_region_eta: T,
#[serde(default = "default_last_step_accepted")]
last_step_accepted: bool,
step_count: usize,
}
fn no_trust_region<T>() -> Option<T> {
None
}
fn default_max_trust_region_radius<T: Float>() -> T {
T::from(10.0).unwrap_or_else(T::one)
}
fn default_min_trust_region_radius<T: Float>() -> T {
T::from(1e-8).unwrap_or_else(T::epsilon)
}
fn default_trust_region_eta<T: Float>() -> T {
T::from(0.1).unwrap_or_else(T::zero)
}
fn default_last_step_accepted() -> bool {
true
}
#[derive(Debug, Clone)]
struct CgSolution<T: Float> {
direction: Array1<T>,
on_boundary: bool,
hd: Array1<T>,
}
fn dot<T: Float>(a: &ArrayView1<T>, b: &ArrayView1<T>) -> T {
a.iter()
.zip(b.iter())
.fold(T::zero(), |acc, (&x, &y)| acc + x * y)
}
fn boundary_tau<T: Float>(d: &Array1<T>, p: &Array1<T>, delta: T) -> Option<T> {
let d_view = d.view();
let p_view = p.view();
let pp = dot(&p_view, &p_view);
if !matches!(
pp.partial_cmp(&T::zero()),
Some(std::cmp::Ordering::Greater)
) {
return None;
}
let dp = dot(&d_view, &p_view);
let dd = dot(&d_view, &d_view);
let c = dd - delta * delta;
let discriminant = dp * dp - pp * c;
if discriminant < T::zero() || !discriminant.is_finite() {
return None;
}
let tau = (-dp + discriminant.sqrt()) / pp;
if tau.is_finite() && tau > T::zero() {
Some(tau)
} else {
None
}
}
impl<T: Float + ScalarOperand> Default for NewtonCG<T> {
fn default() -> Self {
Self {
learning_rate: T::one(),
cg_tolerance: T::from(1e-6).unwrap_or_else(T::epsilon),
cg_max_iters: 100,
hessian_reg: T::from(1e-6).unwrap_or_else(T::epsilon),
trust_region_radius: None,
max_trust_region_radius: default_max_trust_region_radius(),
min_trust_region_radius: default_min_trust_region_radius(),
trust_region_eta: default_trust_region_eta(),
last_step_accepted: true,
step_count: 0,
}
}
}
impl<T: Float + ScalarOperand> NewtonCG<T> {
pub fn new(
learning_rate: T,
cg_tolerance: T,
cg_max_iters: usize,
hessian_reg: T,
) -> Result<Self> {
if !matches!(
learning_rate.partial_cmp(&T::zero()),
Some(std::cmp::Ordering::Greater)
) {
return Err(OptimError::InvalidParameter(
"learning_rate must be positive".to_string(),
));
}
if !matches!(
cg_tolerance.partial_cmp(&T::zero()),
Some(std::cmp::Ordering::Greater)
) {
return Err(OptimError::InvalidParameter(
"cg_tolerance must be positive".to_string(),
));
}
if cg_max_iters == 0 {
return Err(OptimError::InvalidParameter(
"cg_max_iters must be positive".to_string(),
));
}
if hessian_reg < T::zero() {
return Err(OptimError::InvalidParameter(
"hessian_reg must be non-negative".to_string(),
));
}
Ok(Self {
learning_rate,
cg_tolerance,
cg_max_iters,
hessian_reg,
trust_region_radius: None,
max_trust_region_radius: default_max_trust_region_radius(),
min_trust_region_radius: default_min_trust_region_radius(),
trust_region_eta: default_trust_region_eta(),
last_step_accepted: true,
step_count: 0,
})
}
pub fn with_trust_region(mut self, radius: T) -> Result<Self> {
if !matches!(
radius.partial_cmp(&T::zero()),
Some(std::cmp::Ordering::Greater)
) {
return Err(OptimError::InvalidParameter(
"trust region radius must be positive".to_string(),
));
}
self.trust_region_radius = Some(radius);
if self.max_trust_region_radius < radius {
self.max_trust_region_radius = radius;
}
Ok(self)
}
pub fn with_max_trust_region_radius(mut self, radius: T) -> Result<Self> {
if !matches!(
radius.partial_cmp(&T::zero()),
Some(std::cmp::Ordering::Greater)
) {
return Err(OptimError::InvalidParameter(
"max trust region radius must be positive".to_string(),
));
}
self.max_trust_region_radius = radius;
Ok(self)
}
pub fn with_trust_region_eta(mut self, eta: T) -> Result<Self> {
if eta < T::zero() || eta >= T::one() {
return Err(OptimError::InvalidParameter(
"trust region eta must lie in [0, 1)".to_string(),
));
}
self.trust_region_eta = eta;
Ok(self)
}
pub fn trust_region_radius(&self) -> Option<T> {
self.trust_region_radius
}
pub fn last_step_accepted(&self) -> bool {
self.last_step_accepted
}
pub fn step<F>(
&mut self,
params: ArrayView1<T>,
grads: ArrayView1<T>,
hvp_fn: F,
) -> Result<Array1<T>>
where
F: Fn(&[T]) -> Vec<T>,
{
let n = params.len();
if grads.len() != n {
return Err(OptimError::DimensionMismatch(format!(
"Expected gradient size {}, got {}",
n,
grads.len()
)));
}
self.step_count += 1;
self.last_step_accepted = true;
let solution = self.truncated_cg(&grads, &hvp_fn, self.trust_region_radius)?;
Ok(params.to_owned() + &(solution.direction * self.learning_rate))
}
pub fn step_with_loss<F, L>(
&mut self,
params: ArrayView1<T>,
grads: ArrayView1<T>,
hvp_fn: F,
mut loss_fn: L,
) -> Result<Array1<T>>
where
F: Fn(&[T]) -> Vec<T>,
L: FnMut(&Array1<T>) -> T,
{
let n = params.len();
if grads.len() != n {
return Err(OptimError::DimensionMismatch(format!(
"Expected gradient size {}, got {}",
n,
grads.len()
)));
}
self.step_count += 1;
let mut delta = match self.trust_region_radius {
Some(radius) => radius,
None => T::one(),
};
let solution = self.truncated_cg(&grads, &hvp_fn, Some(delta))?;
let direction = solution.direction;
let g_dot_d = dot(&grads, &direction.view());
let d_dot_hd = dot(&direction.view(), &solution.hd.view());
let half = T::from(0.5).unwrap_or_else(|| T::one() / (T::one() + T::one()));
let predicted_reduction = -(g_dot_d + half * d_dot_hd);
let quarter = half * half;
let two = T::one() + T::one();
if !predicted_reduction.is_finite() || predicted_reduction <= T::zero() {
delta = (delta * quarter).max(self.min_trust_region_radius);
self.trust_region_radius = Some(delta);
self.last_step_accepted = false;
return Ok(params.to_owned());
}
let current = params.to_owned();
let candidate = ¤t + &direction;
let f_current = loss_fn(¤t);
let f_candidate = loss_fn(&candidate);
let actual_reduction = f_current - f_candidate;
let rho = if actual_reduction.is_finite() {
actual_reduction / predicted_reduction
} else {
T::neg_infinity()
};
let three_quarters = T::from(0.75).unwrap_or_else(|| T::one() - quarter);
if rho < quarter {
delta = delta * quarter;
} else if rho > three_quarters && solution.on_boundary {
delta = (delta * two).min(self.max_trust_region_radius);
}
delta = delta
.max(self.min_trust_region_radius)
.min(self.max_trust_region_radius);
self.trust_region_radius = Some(delta);
if rho > self.trust_region_eta {
self.last_step_accepted = true;
Ok(candidate)
} else {
self.last_step_accepted = false;
Ok(current)
}
}
fn truncated_cg<F>(
&self,
grads: &ArrayView1<T>,
hvp_fn: &F,
radius: Option<T>,
) -> Result<CgSolution<T>>
where
F: Fn(&[T]) -> Vec<T>,
{
let n = grads.len();
let mut d: Array1<T> = Array1::zeros(n);
let mut r = grads.mapv(|x| -x); let mut p = r.clone();
let mut r_norm_sq = dot(&r.view(), &r.view());
let initial_r_norm_sq = r_norm_sq;
let tol_sq = self.cg_tolerance * self.cg_tolerance;
let residual_target = tol_sq * initial_r_norm_sq;
let mut on_boundary = false;
let tiny = T::from(1e-12).unwrap_or_else(T::epsilon);
if r_norm_sq <= residual_target {
let hd = self.hessian_vector_product(hvp_fn, &d)?;
return Ok(CgSolution {
direction: d,
on_boundary: false,
hd,
});
}
for cg_iter in 0..self.cg_max_iters {
let ap_reg = self.hessian_vector_product(hvp_fn, &p)?;
let p_dot_ap = dot(&p.view(), &ap_reg.view());
if p_dot_ap <= T::zero() {
match radius {
Some(delta) => {
if let Some(tau) = boundary_tau(&d, &p, delta) {
for i in 0..n {
d[i] = d[i] + tau * p[i];
}
}
on_boundary = true;
}
None => {
if cg_iter == 0 {
d = grads.mapv(|x| -x);
}
}
}
break;
}
if p_dot_ap < tiny {
break;
}
let alpha = r_norm_sq / p_dot_ap;
if let Some(delta) = radius {
let mut d_next_norm_sq = T::zero();
for i in 0..n {
let v = d[i] + alpha * p[i];
d_next_norm_sq = d_next_norm_sq + v * v;
}
if d_next_norm_sq >= delta * delta {
if let Some(tau) = boundary_tau(&d, &p, delta) {
for i in 0..n {
d[i] = d[i] + tau * p[i];
}
}
on_boundary = true;
break;
}
}
for i in 0..n {
d[i] = d[i] + alpha * p[i];
}
for i in 0..n {
r[i] = r[i] - alpha * ap_reg[i];
}
let r_norm_sq_new = dot(&r.view(), &r.view());
if r_norm_sq_new <= residual_target {
break;
}
let beta = r_norm_sq_new / r_norm_sq;
r_norm_sq = r_norm_sq_new;
for i in 0..n {
p[i] = r[i] + beta * p[i];
}
}
let hd = self.hessian_vector_product(hvp_fn, &d)?;
Ok(CgSolution {
direction: d,
on_boundary,
hd,
})
}
fn hessian_vector_product<F>(&self, hvp_fn: &F, v: &Array1<T>) -> Result<Array1<T>>
where
F: Fn(&[T]) -> Vec<T>,
{
let n = v.len();
let v_vec: Vec<T> = v.iter().copied().collect();
let hv_vec = hvp_fn(&v_vec);
if hv_vec.len() != n {
return Err(OptimError::DimensionMismatch(format!(
"Hessian-vector product returned wrong size: expected {}, got {}",
n,
hv_vec.len()
)));
}
let mut hv = Array1::from_vec(hv_vec);
for i in 0..n {
hv[i] = hv[i] + self.hessian_reg * v[i];
}
Ok(hv)
}
pub fn step_count(&self) -> usize {
self.step_count
}
pub fn reset(&mut self) {
self.step_count = 0;
self.last_step_accepted = true;
}
pub fn get_learning_rate(&self) -> T {
self.learning_rate
}
pub fn set_learning_rate(&mut self, learning_rate: T) {
self.learning_rate = learning_rate;
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
use scirs2_core::ndarray_ext::array;
#[test]
fn test_newton_cg_creation() {
let optimizer = NewtonCG::<f32>::default();
assert_eq!(optimizer.step_count(), 0);
}
#[test]
fn test_newton_cg_custom_creation() {
let optimizer = NewtonCG::<f32>::new(0.5, 1e-8, 50, 1e-5)
.expect("NewtonCG::<f32>::new succeeds in test_newton_cg_custom_creation");
assert_eq!(optimizer.step_count(), 0);
assert_relative_eq!(optimizer.get_learning_rate(), 0.5);
}
#[test]
fn test_newton_cg_invalid_params() {
assert!(NewtonCG::<f32>::new(-0.1, 1e-6, 100, 1e-6).is_err());
assert!(NewtonCG::<f32>::new(1.0, -1e-6, 100, 1e-6).is_err());
assert!(NewtonCG::<f32>::new(1.0, 1e-6, 0, 1e-6).is_err());
assert!(NewtonCG::<f32>::new(1.0, 1e-6, 100, -1e-6).is_err());
}
#[test]
fn test_newton_cg_quadratic_function() {
let mut optimizer = NewtonCG::<f64>::new(1.0, 1e-8, 50, 0.0)
.expect("NewtonCG::<f64>::new succeeds in test_newton_cg_quadratic_function");
let mut params = array![2.0, 2.0];
let b = array![1.0, 1.0];
let hvp_fn = |v: &[f64]| -> Vec<f64> { vec![2.0 * v[0], 2.0 * v[1]] };
let grads = array![
2.0 * params[0] - b[0], 2.0 * params[1] - b[1] ];
params = optimizer
.step(params.view(), grads.view(), hvp_fn)
.expect("step succeeds in test_newton_cg_quadratic_function");
assert_relative_eq!(params[0], 0.5, epsilon = 0.1);
assert_relative_eq!(params[1], 0.5, epsilon = 0.1);
}
#[test]
fn test_newton_cg_convergence() {
let mut optimizer = NewtonCG::<f64>::new(1.0, 1e-8, 100, 0.0)
.expect("NewtonCG::<f64>::new succeeds in test_newton_cg_convergence");
let mut params = array![5.0, 5.0];
let hvp_fn = |v: &[f64]| -> Vec<f64> { vec![2.0 * v[0], 2.0 * v[1]] };
for _ in 0..10 {
let grads = array![2.0 * params[0], 2.0 * params[1]];
params = optimizer
.step(params.view(), grads.view(), hvp_fn)
.expect("step succeeds in test_newton_cg_convergence");
}
assert!(
params[0].abs() < 0.01,
"Failed to converge, got x = {}",
params[0]
);
assert!(
params[1].abs() < 0.01,
"Failed to converge, got y = {}",
params[1]
);
}
#[test]
fn test_newton_cg_reset() {
let mut optimizer = NewtonCG::<f32>::default();
let params = array![1.0, 2.0, 3.0];
let grads = array![0.1, 0.2, 0.3];
let hvp_fn = |v: &[f32]| -> Vec<f32> { v.to_vec() };
optimizer
.step(params.view(), grads.view(), hvp_fn)
.expect("step succeeds in test_newton_cg_reset");
assert_eq!(optimizer.step_count(), 1);
optimizer.reset();
assert_eq!(optimizer.step_count(), 0);
}
#[test]
fn test_newton_cg_learning_rate() {
let mut optimizer = NewtonCG::<f32>::default();
assert_relative_eq!(optimizer.get_learning_rate(), 1.0);
optimizer.set_learning_rate(0.5);
assert_relative_eq!(optimizer.get_learning_rate(), 0.5);
}
}