use core::f32;
use anyhow::Result;
use ndarray::{Array, Array1, Array2};
use ndarray_linalg::Eig;
use crate::params::CmaesParams;
#[derive(Debug, Clone)]
pub struct CmaesState {
pub z: Array2<f32>,
pub y: Array2<f32>,
pub best_y: Array1<f32>,
pub best_y_fit: Array1<f32>,
pub cov: Array2<f32>,
pub eig_vecs: Array2<f32>,
pub eig_vals: Array1<f32>,
pub inv_sqrt: Array2<f32>,
pub mean: Array1<f32>,
pub sigma: f32,
pub g: i32,
pub evals_count: i32,
pub ps: Array1<f32>,
pub pc: Array1<f32>,
}
impl CmaesState {
pub fn init_state(params: &CmaesParams) -> Result<Self> {
let z: Array2<f32> = Array2::zeros((params.popsize as usize, params.xstart.len()));
let y: Array2<f32> = Array2::zeros((params.popsize as usize, params.xstart.len()));
let best_y: Array1<f32> = Array1::zeros(params.xstart.len());
let best_y_fit: Array1<f32> = Array1::from_elem(1, f32::MAX);
let cov: Array2<f32> = Array2::eye(params.xstart.len());
let inv_sqrt: Array2<f32> = Array2::eye(params.xstart.len());
let eig_vecs: Array2<f32> = Array::eye(params.xstart.len());
let eig_vals: Array1<f32> = Array::from_elem((params.xstart.len(),), 1.0);
let mean: Array1<f32> = Array1::from_vec(params.xstart.clone());
let sigma: f32 = params.sigma;
let g: i32 = 0;
let evals_count = 0;
let ps: Array1<f32> = Array1::zeros(params.xstart.len());
let pc: Array1<f32> = Array1::zeros(params.xstart.len());
Ok(CmaesState {
z,
y,
best_y,
best_y_fit,
cov,
eig_vecs,
eig_vals,
inv_sqrt,
mean,
sigma,
g,
evals_count,
ps,
pc,
})
}
pub fn prepare_ask(&mut self) -> Result<()> {
self.eigen_decomposition();
Ok(())
}
fn eigen_decomposition(&mut self) {
self.cov = (&self.cov + &self.cov.t()) / 2.0;
let (eig_vals, eig_vecs) = self.cov.eig().unwrap();
let mut eig_vals: Array1<f32> = eig_vals.mapv(|eig| eig.re);
let eig_vecs: Array2<f32> = eig_vecs.mapv(|vec| vec.re);
eig_vals.map_inplace(|elem| {
if *elem < 0.0 {
*elem = 0.1 } else if *elem > 10. {
*elem = 10.; }
});
self.inv_sqrt = Array2::from_diag(&eig_vals.mapv(|elem| elem.powf(-0.5)));
self.inv_sqrt = eig_vecs.dot(&self.inv_sqrt).dot(&eig_vecs.t());
self.eig_vecs = eig_vecs;
self.eig_vals = eig_vals;
}
}