use basin::problems::{Rosenbrock, Sphere};
use basin::{
CmaEs, CmaEsState, CmaInject, DenseMatrix, Executor, RhoTolerance, SolisWets, SolisWetsState,
State, StepOutcome, TerminationReason,
};
#[test]
fn same_seed_yields_identical_trajectory() {
let run = || {
Executor::from_start(
Sphere::<Vec<f64>>::new(),
SolisWets::new(42),
vec![2.0, -1.5],
)
.max_iter(200)
.run()
.unwrap()
};
let result_a = run();
let result_b = run();
assert_eq!(result_a.cost(), result_b.cost());
assert_eq!(result_a.param(), result_b.param());
}
#[test]
fn different_seeds_yield_different_trajectories() {
let run = |seed| {
Executor::from_start(
Sphere::<Vec<f64>>::new(),
SolisWets::new(seed),
vec![2.0, -1.5],
)
.max_iter(50)
.run()
.unwrap()
};
assert_ne!(run(1).param(), run(2).param());
}
#[test]
fn converges_on_sphere_5d_via_rho_tolerance() {
let result = Executor::from_start(
Sphere::<Vec<f64>>::new(),
SolisWets::new(7),
vec![2.0, -1.0, 1.5, 0.5, -2.0],
)
.terminate_on(RhoTolerance::new(1e-8))
.max_iter(100_000)
.run()
.unwrap();
assert_eq!(result.reason, TerminationReason::RhoTolerance);
assert!(
result.cost() < 1e-6,
"sphere 5-D cost = {} (expected < 1e-6)",
result.cost()
);
}
#[test]
fn makes_progress_on_rosenbrock_2d() {
let result = Executor::from_start(
Rosenbrock::<Vec<f64>>::new(),
SolisWets::new(3),
vec![-1.2, 1.0],
)
.terminate_on(RhoTolerance::new(1e-10))
.max_iter(50_000)
.run()
.unwrap();
assert!(
result.cost() < 1e-2,
"rosenbrock 2-D cost = {} (expected < 1e-2)",
result.cost()
);
}
#[test]
fn cost_is_monotone_nonincreasing() {
let mut stepper = Executor::from_start(
Sphere::<Vec<f64>>::new(),
SolisWets::new(99),
vec![3.0, -2.0, 1.0],
)
.max_iter(500)
.into_stepper()
.unwrap();
let mut prev = stepper.state().cost();
while let StepOutcome::Continue = stepper.step().unwrap() {
let current = stepper.state().cost();
assert!(
current <= prev,
"cost increased: prev = {prev}, current = {current}"
);
prev = current;
}
}
#[test]
fn from_start_matches_explicit_state() {
let via_seed =
Executor::from_start(Sphere::<Vec<f64>>::new(), SolisWets::new(5), vec![1.0, 2.0])
.max_iter(100)
.run()
.unwrap();
let via_state = Executor::new(
Sphere::<Vec<f64>>::new(),
SolisWets::new(5),
SolisWetsState::new(vec![1.0, 2.0], 1.0),
)
.max_iter(100)
.run()
.unwrap();
assert_eq!(via_seed.cost(), via_state.cost());
assert_eq!(via_seed.param(), via_state.param());
}
#[test]
fn works_as_cma_inject_inner() {
let cma = CmaEs::<Vec<f64>, DenseMatrix>::new(17);
let solver = CmaInject::with_inner_solver(cma, SolisWets::new(23))
.with_k(1)
.with_inner_max_iter(30);
let result = Executor::new(
Sphere::<Vec<f64>>::new(),
solver,
CmaEsState::<Vec<f64>, DenseMatrix>::new(vec![2.0, -1.5], 0.5),
)
.max_iter(100)
.run()
.unwrap();
assert!(
result.cost() < 1e-6,
"cma-inject(solis-wets) sphere cost = {} (expected < 1e-6)",
result.cost()
);
}