#![cfg(feature = "ndarray")]
use basin::problems::BoothBoxed;
use basin::{BoundedCmaEs, BoundedCmaInject, CmaEsState, Executor, Lbfgsb};
use ndarray::{Array1, Array2};
#[test]
fn converges_on_booth_boxed_slack() {
let lower = Array1::from_elem(2, -5.0);
let upper = Array1::from_elem(2, 5.0);
let problem = BoothBoxed::<Array1<f64>>::new(lower, upper);
let m0 = Array1::from_vec(vec![0.0, 2.0]);
let cma = BoundedCmaEs::<Array1<f64>, Array2<f64>>::new(19);
let solver = BoundedCmaInject::with_inner_solver(cma, Lbfgsb::new())
.with_k(1)
.with_inner_max_iter(50);
let result = Executor::new(
problem,
solver,
CmaEsState::<Array1<f64>, Array2<f64>>::new(m0, 0.5),
)
.max_iter(200)
.run()
.unwrap();
let p = result.param();
let err = (p[0] - 1.0).abs().max((p[1] - 3.0).abs());
assert!(
err <= 1e-6,
"booth-boxed iterate = ({}, {}), expected ≈ (1, 3) within 1e-6 (err = {})",
p[0],
p[1],
err
);
}
#[test]
fn aggregates_lbfgsb_work_into_outer() {
let lower = Array1::from_elem(2, -5.0);
let upper = Array1::from_elem(2, 5.0);
let m0 = Array1::from_vec(vec![0.0, 2.0]);
let outer_iters: u64 = 20;
let inner_iters: u64 = 50;
let k: usize = 1;
let vanilla = Executor::new(
BoothBoxed::<Array1<f64>>::new(lower.clone(), upper.clone()),
BoundedCmaEs::<Array1<f64>, Array2<f64>>::new(29),
CmaEsState::<Array1<f64>, Array2<f64>>::new(m0.clone(), 0.5),
)
.max_iter(outer_iters)
.run()
.unwrap();
let cma = BoundedCmaEs::<Array1<f64>, Array2<f64>>::new(29);
let solver = BoundedCmaInject::with_inner_solver(cma, Lbfgsb::new())
.with_k(k)
.with_inner_max_iter(inner_iters);
let memetic = Executor::new(
BoothBoxed::<Array1<f64>>::new(lower, upper),
solver,
CmaEsState::<Array1<f64>, Array2<f64>>::new(m0, 0.5),
)
.max_iter(outer_iters)
.run()
.unwrap();
let min_extra = (outer_iters.saturating_sub(1)) * (k as u64) * 3;
assert!(
memetic.cost_evals() >= vanilla.cost_evals() + min_extra,
"memetic cost_evals = {} should exceed vanilla {} by at least \
{} (outer iters × k × (L-Bfgs-B init cost + gradient + re-eval))",
memetic.cost_evals(),
vanilla.cost_evals(),
min_extra
);
}