use crate::prelude::*;
use rand::prelude::*;
use rand_xorshift::XorShiftRng;
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Serialize, Deserialize, Debug)]
pub enum SATempFunc<F> {
TemperatureFast,
Boltzmann,
Exponential(F),
}
impl<F> std::default::Default for SATempFunc<F> {
fn default() -> Self {
SATempFunc::Boltzmann
}
}
#[derive(Clone, Serialize, Deserialize)]
pub struct SimulatedAnnealing<F> {
init_temp: F,
temp_func: SATempFunc<F>,
temp_iter: u64,
stall_iter_accepted: u64,
stall_iter_accepted_limit: u64,
stall_iter_best: u64,
stall_iter_best_limit: u64,
reanneal_fixed: u64,
reanneal_iter_fixed: u64,
reanneal_accepted: u64,
reanneal_iter_accepted: u64,
reanneal_best: u64,
reanneal_iter_best: u64,
cur_temp: F,
rng: XorShiftRng,
}
impl<F> SimulatedAnnealing<F>
where
F: ArgminFloat,
{
pub fn new(init_temp: F) -> Result<Self, Error> {
if init_temp <= F::from_f64(0.0).unwrap() {
Err(ArgminError::InvalidParameter {
text: "Initial temperature must be > 0.".to_string(),
}
.into())
} else {
Ok(SimulatedAnnealing {
init_temp,
temp_func: SATempFunc::TemperatureFast,
temp_iter: 0,
stall_iter_accepted: 0,
stall_iter_accepted_limit: std::u64::MAX,
stall_iter_best: 0,
stall_iter_best_limit: std::u64::MAX,
reanneal_fixed: std::u64::MAX,
reanneal_iter_fixed: 0,
reanneal_accepted: std::u64::MAX,
reanneal_iter_accepted: 0,
reanneal_best: std::u64::MAX,
reanneal_iter_best: 0,
cur_temp: init_temp,
rng: XorShiftRng::from_entropy(),
})
}
}
pub fn temp_func(mut self, temperature_func: SATempFunc<F>) -> Self {
self.temp_func = temperature_func;
self
}
pub fn stall_accepted(mut self, iter: u64) -> Self {
self.stall_iter_accepted_limit = iter;
self
}
pub fn stall_best(mut self, iter: u64) -> Self {
self.stall_iter_best_limit = iter;
self
}
pub fn reannealing_fixed(mut self, iter: u64) -> Self {
self.reanneal_fixed = iter;
self
}
pub fn reannealing_accepted(mut self, iter: u64) -> Self {
self.reanneal_accepted = iter;
self
}
pub fn reannealing_best(mut self, iter: u64) -> Self {
self.reanneal_best = iter;
self
}
fn update_temperature(&mut self) {
self.cur_temp = match self.temp_func {
SATempFunc::TemperatureFast => {
self.init_temp / F::from_u64(self.temp_iter + 1).unwrap()
}
SATempFunc::Boltzmann => self.init_temp / F::from_u64(self.temp_iter + 1).unwrap().ln(),
SATempFunc::Exponential(x) => {
self.init_temp * x.powf(F::from_u64(self.temp_iter + 1).unwrap())
}
};
}
fn reanneal(&mut self) -> (bool, bool, bool) {
let out = (
self.reanneal_iter_fixed >= self.reanneal_fixed,
self.reanneal_iter_accepted >= self.reanneal_accepted,
self.reanneal_iter_best >= self.reanneal_best,
);
if out.0 || out.1 || out.2 {
self.reanneal_iter_fixed = 0;
self.reanneal_iter_accepted = 0;
self.reanneal_iter_best = 0;
self.cur_temp = self.init_temp;
self.temp_iter = 0;
}
out
}
fn update_stall_and_reanneal_iter(&mut self, accepted: bool, new_best: bool) {
self.stall_iter_accepted = if accepted {
0
} else {
self.stall_iter_accepted + 1
};
self.reanneal_iter_accepted = if accepted {
0
} else {
self.reanneal_iter_accepted + 1
};
self.stall_iter_best = if new_best {
0
} else {
self.stall_iter_best + 1
};
self.reanneal_iter_best = if new_best {
0
} else {
self.reanneal_iter_best + 1
};
}
}
impl<O, F> Solver<O> for SimulatedAnnealing<F>
where
O: ArgminOp<Output = F, Float = F>,
F: ArgminFloat,
{
const NAME: &'static str = "Simulated Annealing";
fn init(
&mut self,
_op: &mut OpWrapper<O>,
_state: &IterState<O>,
) -> Result<Option<ArgminIterData<O>>, Error> {
Ok(Some(ArgminIterData::new().kv(make_kv!(
"initial_temperature" => self.init_temp;
"stall_iter_accepted_limit" => self.stall_iter_accepted_limit;
"stall_iter_best_limit" => self.stall_iter_best_limit;
"reanneal_fixed" => self.reanneal_fixed;
"reanneal_accepted" => self.reanneal_accepted;
"reanneal_best" => self.reanneal_best;
))))
}
fn next_iter(
&mut self,
op: &mut OpWrapper<O>,
state: &IterState<O>,
) -> Result<ArgminIterData<O>, Error> {
let prev_param = state.get_param();
let prev_cost = state.get_cost();
let new_param = op.modify(&prev_param, self.cur_temp)?;
let new_cost = op.apply(&new_param)?;
let prob: f64 = self.rng.gen();
let prob = F::from_f64(prob).unwrap();
let accepted = (new_cost < state.get_prev_cost())
|| (F::from_f64(1.0).unwrap()
/ (F::from_f64(1.0).unwrap()
+ ((new_cost - state.get_prev_cost()) / self.cur_temp).exp())
> prob);
self.update_stall_and_reanneal_iter(accepted, new_cost <= state.get_best_cost());
let (r_fixed, r_accepted, r_best) = self.reanneal();
self.temp_iter += 1;
self.reanneal_iter_fixed += 1;
self.update_temperature();
Ok(if accepted {
ArgminIterData::new().param(new_param).cost(new_cost)
} else {
ArgminIterData::new().param(prev_param).cost(prev_cost)
}
.kv(make_kv!(
"t" => self.cur_temp;
"new_be" => new_cost <= state.get_best_cost();
"acc" => accepted;
"st_i_be" => self.stall_iter_best;
"st_i_ac" => self.stall_iter_accepted;
"ra_i_fi" => self.reanneal_iter_fixed;
"ra_i_be" => self.reanneal_iter_best;
"ra_i_ac" => self.reanneal_iter_accepted;
"ra_fi" => r_fixed;
"ra_be" => r_best;
"ra_ac" => r_accepted;
)))
}
fn terminate(&mut self, _state: &IterState<O>) -> TerminationReason {
if self.stall_iter_accepted > self.stall_iter_accepted_limit {
return TerminationReason::AcceptedStallIterExceeded;
}
if self.stall_iter_best > self.stall_iter_best_limit {
return TerminationReason::BestStallIterExceeded;
}
TerminationReason::NotTerminated
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_trait_impl;
test_trait_impl!(sa, SimulatedAnnealing<f64>);
}