use crate::linear_algebra::Vector;
use crate::linear_algebra::qr::{DampedLeastSquares, enorm, max, min};
use crate::scalar::Numeric;
#[derive(Debug, Clone, Copy)]
pub(crate) struct LmParameter<const N: usize, T = f64> {
pub lambda: T,
pub step: Vector<N, T>,
}
pub(crate) fn determine_lambda_and_parameter_update<const N: usize, T: Numeric>(
dls: &DampedLeastSquares<N, T>,
diag: &[T; N],
delta: T,
initial_lambda: T,
) -> LmParameter<N, T> {
let dwarf = T::MIN_POSITIVE;
let p1 = T::from_f64(0.1);
let p001 = T::from_f64(0.001);
let scale = |v: &Vector<N, T>| -> [T; N] { core::array::from_fn(|j| diag[j] * v[j]) };
let (gauss_newton, _) = dls.solve_with_zero_diagonal();
let full_rank = dls.is_non_singular();
let mut dxnorm = enorm(&scale(&gauss_newton));
let mut fp = dxnorm - delta;
if fp <= p1 * delta {
return LmParameter {
lambda: T::ZERO,
step: gauss_newton,
};
}
let mut parl = T::ZERO;
if full_rank {
let scaled = scale(&gauss_newton);
let mut w: [T; N] = core::array::from_fn(|j| {
diag[dls.permutation[j]] * (scaled[dls.permutation[j]] / dxnorm)
});
for j in 0..N {
let mut sum = T::ZERO;
for (i, &wi) in w.iter().enumerate().take(j) {
sum += dls.r[(i, j)] * wi;
}
w[j] = (w[j] - sum) / dls.r[(j, j)];
}
let temp = enorm(&w);
parl = ((fp / delta) / temp) / temp;
}
let w: [T; N] = core::array::from_fn(|j| {
let mut sum = T::ZERO;
for i in 0..=j {
sum += dls.r[(i, j)] * dls.qt_b[i];
}
sum / diag[dls.permutation[j]]
});
let gnorm = enorm(&w);
let mut paru = gnorm / delta;
if paru == T::ZERO {
paru = dwarf / min(delta, p1);
}
let mut par = min(max(initial_lambda, parl), paru);
if par == T::ZERO {
par = gnorm / dxnorm;
}
let mut step = gauss_newton;
for iter in 1..=10 {
if par == T::ZERO {
par = max(dwarf, p001 * paru);
}
let sqrt_par = par.sqrt();
let scaled_diag: [T; N] = core::array::from_fn(|j| sqrt_par * diag[j]);
let (x, cholesky) = dls.solve_with_diagonal(&scaled_diag);
step = x;
let scaled_step = scale(&step);
dxnorm = enorm(&scaled_step);
let previous_fp = fp;
fp = dxnorm - delta;
if fp.abs() <= p1 * delta
|| (parl == T::ZERO && fp <= previous_fp && previous_fp < T::ZERO)
|| iter == 10
{
break;
}
let rhs: [T; N] = core::array::from_fn(|j| {
diag[dls.permutation[j]] * (scaled_step[dls.permutation[j]] / dxnorm)
});
let temp = enorm(&cholesky.solve(rhs));
if temp == T::ZERO {
break;
}
let parc = ((fp / delta) / temp) / temp;
if fp > T::ZERO {
parl = max(parl, par);
} else {
paru = min(paru, par);
}
par = max(parl, par + parc);
}
LmParameter { lambda: par, step }
}