#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Penalties {
pub domain_wall: f64,
pub copy: f64,
}
impl Default for Penalties {
fn default() -> Self {
Penalties { domain_wall: 1.0, copy: 1.0 }
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Stage {
pub beta: f64,
pub sweeps: usize,
pub penalties: Penalties,
}
#[derive(Clone, Debug, PartialEq, Default)]
pub struct Schedule {
stages: Vec<Stage>,
}
impl Schedule {
pub fn new() -> Self {
Schedule { stages: Vec::new() }
}
pub fn constant(beta: f64, sweeps: usize) -> Self {
Schedule { stages: vec![Stage { beta, sweeps, penalties: Penalties::default() }] }
}
pub fn geometric(beta_min: f64, beta_max: f64, stages: usize, sweeps_per: usize) -> Self {
assert!(beta_min > 0.0, "beta_min must be positive; a geometric ladder cannot start at 0");
assert!(beta_max > beta_min, "need beta_max > beta_min");
assert!(stages >= 2, "a ladder needs at least 2 rungs");
let r = (beta_max / beta_min).powf(1.0 / (stages - 1) as f64);
Schedule {
stages: (0..stages)
.map(|i| Stage {
beta: beta_min * r.powi(i as i32),
sweeps: sweeps_per,
penalties: Penalties::default(),
})
.collect(),
}
}
pub fn ramp_domain_wall(mut self, start: f64, end: f64) -> Self {
for (i, s) in ramp(start, end, self.stages.len()).into_iter().enumerate() {
self.stages[i].penalties.domain_wall = s;
}
self
}
pub fn ramp_copy(mut self, start: f64, end: f64) -> Self {
for (i, s) in ramp(start, end, self.stages.len()).into_iter().enumerate() {
self.stages[i].penalties.copy = s;
}
self
}
pub fn push(&mut self, stage: Stage) {
self.stages.push(stage);
}
pub fn stages(&self) -> &[Stage] {
&self.stages
}
pub fn len(&self) -> usize {
self.stages.len()
}
pub fn is_empty(&self) -> bool {
self.stages.is_empty()
}
pub fn total_sweeps(&self) -> u64 {
self.stages.iter().map(|s| s.sweeps as u64).sum()
}
pub fn node_updates(&self, n: usize) -> u64 {
self.total_sweeps() * n as u64
}
}
fn ramp(start: f64, end: f64, n: usize) -> Vec<f64> {
if n == 0 {
return Vec::new();
}
if n == 1 || start == end {
return vec![end; n];
}
if start > 0.0 && end > 0.0 {
let r = (end / start).powf(1.0 / (n - 1) as f64);
(0..n).map(|i| start * r.powi(i as i32)).collect()
} else {
let step = (end - start) / (n - 1) as f64;
(0..n).map(|i| start + step * i as f64).collect()
}
}
impl From<&[(f64, usize)]> for Schedule {
fn from(v: &[(f64, usize)]) -> Self {
Schedule {
stages: v
.iter()
.map(|&(beta, sweeps)| Stage { beta, sweeps, penalties: Penalties::default() })
.collect(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn geometric_ladder_hits_both_ends() {
let s = Schedule::geometric(0.05, 4.0, 40, 10);
assert_eq!(s.len(), 40);
assert!((s.stages()[0].beta - 0.05).abs() < 1e-12);
assert!((s.stages()[39].beta - 4.0).abs() < 1e-12);
let st = s.stages();
let r0 = st[1].beta / st[0].beta;
for w in st.windows(2) {
assert!(w[1].beta > w[0].beta);
assert!((w[1].beta / w[0].beta - r0).abs() < 1e-12);
}
}
#[test]
fn penalties_ramp_independently_of_temperature() {
let s = Schedule::geometric(0.1, 2.0, 10, 5).ramp_domain_wall(0.5, 8.0).ramp_copy(1.0, 4.0);
let st = s.stages();
assert!((st[0].penalties.domain_wall - 0.5).abs() < 1e-12);
assert!((st[9].penalties.domain_wall - 8.0).abs() < 1e-12);
assert!((st[0].penalties.copy - 1.0).abs() < 1e-12);
assert!((st[9].penalties.copy - 4.0).abs() < 1e-12);
assert!((st[0].beta - 0.1).abs() < 1e-12);
assert!((st[9].beta - 2.0).abs() < 1e-12);
}
#[test]
fn a_ramp_through_zero_does_not_produce_nan() {
let s = Schedule::geometric(0.1, 1.0, 5, 1).ramp_domain_wall(0.0, 4.0);
for st in s.stages() {
assert!(st.penalties.domain_wall.is_finite(), "{:?}", st.penalties);
}
}
#[test]
fn sizing_a_run_before_starting_it() {
let s = Schedule::geometric(0.1, 2.0, 40, 25);
assert_eq!(s.total_sweeps(), 1000);
assert_eq!(s.node_updates(4900), 4_900_000);
}
#[test]
fn degenerate_ladders_are_rejected_loudly() {
assert!(std::panic::catch_unwind(|| Schedule::geometric(0.0, 1.0, 4, 1)).is_err());
assert!(std::panic::catch_unwind(|| Schedule::geometric(1.0, 0.5, 4, 1)).is_err());
assert!(std::panic::catch_unwind(|| Schedule::geometric(0.1, 1.0, 1, 1)).is_err());
}
}