use crate::sqp::qp_assembly::Triplet;
use pounce_common::types::{Index, Number};
pub struct DampedBfgs {
n: usize,
b: Vec<Number>,
prev_x: Option<Vec<Number>>,
prev_grad_lag: Option<Vec<Number>>,
sized: bool,
h_irow: Vec<Index>,
h_jcol: Vec<Index>,
}
impl DampedBfgs {
pub fn new(n: usize) -> Self {
let nz = n * (n + 1) / 2;
let mut b = vec![0.0; nz];
let mut h_irow = Vec::with_capacity(nz);
let mut h_jcol = Vec::with_capacity(nz);
for i in 0..n {
for j in 0..=i {
if i == j {
b[i * (i + 1) / 2 + j] = 1.0;
}
h_irow.push((i + 1) as Index);
h_jcol.push((j + 1) as Index);
}
}
Self {
n,
b,
prev_x: None,
prev_grad_lag: None,
sized: false,
h_irow,
h_jcol,
}
}
pub fn has_prev(&self) -> bool {
self.prev_x.is_some()
}
pub fn seed_scale(&mut self, gamma: Number) {
if !gamma.is_finite() || gamma <= 0.0 {
return;
}
for i in 0..self.n {
self.set(i, i, gamma);
}
self.sized = true;
}
pub fn reset_to_scale(&mut self) {
let mut sum = 0.0;
let mut count = 0usize;
for i in 0..self.n {
let d = self.get(i, i);
if d.is_finite() && d > 0.0 {
sum += d;
count += 1;
}
}
let gamma = if count > 0 {
sum / count as Number
} else {
1.0
};
let gamma = if gamma.is_finite() && gamma > 0.0 {
gamma
} else {
1.0
};
for v in self.b.iter_mut() {
*v = 0.0;
}
for i in 0..self.n {
self.set(i, i, gamma);
}
}
fn idx(&self, i: usize, j: usize) -> usize {
debug_assert!(i < self.n && j < self.n);
let (lo, hi) = if i >= j { (j, i) } else { (i, j) };
hi * (hi + 1) / 2 + lo
}
fn get(&self, i: usize, j: usize) -> Number {
self.b[self.idx(i, j)]
}
fn set(&mut self, i: usize, j: usize, v: Number) {
let k = self.idx(i, j);
self.b[k] = v;
}
pub fn update(&mut self, x_new: &[Number], grad_lag_new: &[Number]) {
assert_eq!(x_new.len(), self.n, "BFGS::update: x_new.len() != n");
assert_eq!(
grad_lag_new.len(),
self.n,
"BFGS::update: grad_lag_new.len() != n"
);
if let (Some(prev_x), Some(prev_grad_lag)) = (self.prev_x.take(), self.prev_grad_lag.take())
{
let s: Vec<Number> = x_new
.iter()
.zip(prev_x.iter())
.map(|(a, b)| a - b)
.collect();
let y: Vec<Number> = grad_lag_new
.iter()
.zip(prev_grad_lag.iter())
.map(|(a, b)| a - b)
.collect();
self.update_sy(&s, &y);
}
self.prev_x = Some(x_new.to_vec());
self.prev_grad_lag = Some(grad_lag_new.to_vec());
}
pub fn update_sy(&mut self, s: &[Number], y: &[Number]) {
assert_eq!(s.len(), self.n, "BFGS::update_sy: s.len() != n");
assert_eq!(y.len(), self.n, "BFGS::update_sy: y.len() != n");
{
if !self.sized {
let s_y: Number = s.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let s_s: Number = s.iter().map(|v| v * v).sum();
if s_y > 1e-30 && s_s > 1e-30 {
let gamma = s_y / s_s;
for i in 0..self.n {
self.set(i, i, gamma);
}
}
self.sized = true;
}
let bs: Vec<Number> = (0..self.n)
.map(|i| (0..self.n).map(|j| self.get(i, j) * s[j]).sum())
.collect();
let s_bs: Number = s.iter().zip(bs.iter()).map(|(a, b)| a * b).sum();
let s_y: Number = s.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
let theta = if s_y >= 0.2 * s_bs {
1.0
} else if s_bs - s_y > 1e-14 {
0.8 * s_bs / (s_bs - s_y)
} else {
1.0
};
let y_damp: Vec<Number> = y
.iter()
.zip(bs.iter())
.map(|(yi, bsi)| theta * yi + (1.0 - theta) * bsi)
.collect();
let s_y_damp: Number = s.iter().zip(y_damp.iter()).map(|(a, b)| a * b).sum();
if s_bs > 1e-14 && s_y_damp > 1e-14 {
for i in 0..self.n {
for j in 0..=i {
let new_val = self.get(i, j) - (bs[i] * bs[j]) / s_bs
+ (y_damp[i] * y_damp[j]) / s_y_damp;
self.set(i, j, new_val);
}
}
}
}
}
pub fn as_triplet(&self) -> Triplet {
let mut vals = Vec::with_capacity(self.h_irow.len());
for i in 0..self.n {
for j in 0..=i {
vals.push(self.get(i, j));
}
}
Triplet {
n_rows: self.n,
n_cols: self.n,
irow: self.h_irow.clone(),
jcol: self.h_jcol.clone(),
vals,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn diag(b: &DampedBfgs, i: usize) -> Number {
b.get(i, i)
}
#[test]
fn first_update_sizes_the_identity_seed() {
let mut b = DampedBfgs::new(2);
b.update(&[0.0, 0.0], &[0.0, 0.0]); assert!((diag(&b, 0) - 1.0).abs() < 1e-12, "seed must be I");
assert!((diag(&b, 1) - 1.0).abs() < 1e-12, "seed must be I");
b.update(&[1.0, 0.0], &[9.0, 0.0]); assert!(
(diag(&b, 1) - 9.0).abs() < 1e-9,
"off-axis diagonal should be sized to γ = 9, got {}",
diag(&b, 1)
);
assert!(
(diag(&b, 0) - 9.0).abs() < 1e-9,
"on-axis diagonal should be 9 after sizing + rank-2 update, got {}",
diag(&b, 0)
);
}
#[test]
fn seed_scale_sets_the_diagonal_and_marks_sized() {
let mut b = DampedBfgs::new(3);
b.seed_scale(25.0);
assert!(b.sized, "seeding must suppress the later one-time sizing");
for i in 0..3 {
assert!((diag(&b, i) - 25.0).abs() < 1e-12);
}
let mut c = DampedBfgs::new(2);
for bad in [0.0, -1.0, f64::NAN, f64::INFINITY] {
c.seed_scale(bad);
assert!(!c.sized, "seed_scale({bad}) must be refused");
assert!((diag(&c, 0) - 1.0).abs() < 1e-12, "B must stay I");
}
}
#[test]
fn reset_to_scale_drops_curvature_but_keeps_magnitude() {
let mut b = DampedBfgs::new(2);
b.seed_scale(100.0);
b.update(&[0.0, 0.0], &[0.0, 0.0]);
b.update(&[1.0, 1.0], &[150.0, 40.0]); assert!(
b.get(1, 0).abs() > 1e-9,
"test precondition: expected off-diagonal curvature, got {}",
b.get(1, 0)
);
let mean_diag = (diag(&b, 0) + diag(&b, 1)) / 2.0;
b.reset_to_scale();
assert!(b.get(1, 0).abs() < 1e-12, "off-diagonals must be zeroed");
for i in 0..2 {
assert!(
(diag(&b, i) - mean_diag).abs() < 1e-9,
"diagonal must keep the mean scale {mean_diag}, got {}",
diag(&b, i)
);
}
assert!(
mean_diag > 10.0,
"sanity: the retained scale should reflect the problem, not 1"
);
}
#[test]
fn update_sy_matches_the_x_grad_lag_form() {
let mut via_update = DampedBfgs::new(2);
via_update.update(&[0.0, 0.0], &[1.0, 2.0]);
via_update.update(&[1.0, 3.0], &[4.0, 9.0]);
let mut via_sy = DampedBfgs::new(2);
via_sy.update_sy(&[1.0, 3.0], &[3.0, 7.0]);
for i in 0..2 {
for j in 0..=i {
assert!(
(via_update.get(i, j) - via_sy.get(i, j)).abs() < 1e-12,
"B[{i},{j}]: update={} update_sy={}",
via_update.get(i, j),
via_sy.get(i, j)
);
}
}
}
#[test]
fn sizing_happens_only_once() {
let mut b = DampedBfgs::new(2);
b.update(&[0.0, 0.0], &[0.0, 0.0]);
b.update(&[1.0, 0.0], &[9.0, 0.0]); assert!(b.sized);
b.update(&[1.0, 1.0], &[9.0, 4.0]); assert!(
(diag(&b, 0) - 9.0).abs() < 1e-9,
"second pair must not re-seed; axis-0 diagonal = {}",
diag(&b, 0)
);
}
}