use crate::error::{MinuetError, Result};
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum Temperature {
#[default]
Soft,
Hard,
Beta(f64),
Annealed {
start: f64,
end: f64,
steps: usize,
},
}
impl Temperature {
#[must_use]
pub fn soft() -> Self {
Self::Soft
}
#[must_use]
pub fn hard() -> Self {
Self::Hard
}
pub fn beta(beta: f64) -> Result<Self> {
if beta <= 0.0 {
return Err(MinuetError::InvalidTemperature {
beta,
min: 0.0,
max: f64::INFINITY,
});
}
Ok(Self::Beta(beta))
}
pub fn annealed(start: f64, end: f64, steps: usize) -> Result<Self> {
if start <= 0.0 {
return Err(MinuetError::InvalidTemperature {
beta: start,
min: 0.0,
max: f64::INFINITY,
});
}
if end <= 0.0 {
return Err(MinuetError::InvalidTemperature {
beta: end,
min: 0.0,
max: f64::INFINITY,
});
}
if steps == 0 {
return Err(MinuetError::InvalidQuery(
"Annealing steps must be positive".into(),
));
}
Ok(Self::Annealed { start, end, steps })
}
#[must_use]
pub fn beta_at(&self, iteration: usize) -> f64 {
match self {
Self::Soft => 1.0,
Self::Hard => f64::MAX, Self::Beta(b) => *b,
Self::Annealed { start, end, steps } => {
if iteration >= *steps {
*end
} else {
let progress = iteration as f64 / *steps as f64;
let log_start = start.ln();
let log_end = end.ln();
(log_start + progress * (log_end - log_start)).exp()
}
}
}
}
#[must_use]
pub fn is_hard(&self) -> bool {
matches!(self, Self::Hard)
}
#[must_use]
pub fn is_annealed(&self) -> bool {
matches!(self, Self::Annealed { .. })
}
#[must_use]
pub fn num_steps(&self) -> usize {
match self {
Self::Annealed { steps, .. } => *steps,
_ => 1,
}
}
#[must_use]
pub fn to_resonator_config(&self) -> amari_holographic::ResonatorConfig {
const HARD_BETA: f64 = f64::MAX;
let mut config = amari_holographic::ResonatorConfig::default();
match self {
Self::Soft => {
config.initial_beta = 1.0;
config.final_beta = 1.0;
}
Self::Hard => {
config.initial_beta = HARD_BETA;
config.final_beta = HARD_BETA;
}
Self::Beta(b) => {
config.initial_beta = *b;
config.final_beta = *b;
}
Self::Annealed { start, end, steps } => {
config.initial_beta = *start;
config.final_beta = *end;
config.max_iterations = *steps;
}
}
config
}
}
#[derive(Debug, Clone)]
pub struct TemperatureSchedule {
temperatures: Vec<f64>,
current_step: usize,
}
impl TemperatureSchedule {
#[must_use]
pub fn constant(beta: f64, steps: usize) -> Self {
Self {
temperatures: vec![beta; steps],
current_step: 0,
}
}
#[must_use]
pub fn linear(start: f64, end: f64, steps: usize) -> Self {
let temperatures: Vec<f64> = (0..steps)
.map(|i| {
let t = i as f64 / (steps - 1).max(1) as f64;
start + t * (end - start)
})
.collect();
Self {
temperatures,
current_step: 0,
}
}
#[must_use]
pub fn exponential(start: f64, end: f64, steps: usize) -> Self {
let log_start = start.ln();
let log_end = end.ln();
let temperatures: Vec<f64> = (0..steps)
.map(|i| {
let t = i as f64 / (steps - 1).max(1) as f64;
(log_start + t * (log_end - log_start)).exp()
})
.collect();
Self {
temperatures,
current_step: 0,
}
}
#[must_use]
pub fn cosine(start: f64, end: f64, steps: usize) -> Self {
use std::f64::consts::PI;
let temperatures: Vec<f64> = (0..steps)
.map(|i| {
let t = i as f64 / (steps - 1).max(1) as f64;
let cos_factor = (1.0 - (PI * t).cos()) / 2.0;
start + cos_factor * (end - start)
})
.collect();
Self {
temperatures,
current_step: 0,
}
}
#[must_use]
pub fn current(&self) -> f64 {
self.temperatures
.get(self.current_step)
.copied()
.unwrap_or(*self.temperatures.last().unwrap_or(&1.0))
}
pub fn step(&mut self) {
if self.current_step < self.temperatures.len() {
self.current_step += 1;
}
}
pub fn reset(&mut self) {
self.current_step = 0;
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.current_step >= self.temperatures.len()
}
#[must_use]
pub fn temperatures(&self) -> &[f64] {
&self.temperatures
}
#[must_use]
pub fn len(&self) -> usize {
self.temperatures.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.temperatures.is_empty()
}
}
impl Iterator for TemperatureSchedule {
type Item = f64;
fn next(&mut self) -> Option<Self::Item> {
if self.current_step < self.temperatures.len() {
let temp = self.temperatures[self.current_step];
self.current_step += 1;
Some(temp)
} else {
None
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn soft_temperature() {
let temp = Temperature::soft();
assert!((temp.beta_at(0) - 1.0).abs() < 1e-10);
assert!((temp.beta_at(100) - 1.0).abs() < 1e-10);
}
#[test]
fn explicit_beta() {
let temp = Temperature::beta(2.5).unwrap();
assert!((temp.beta_at(0) - 2.5).abs() < 1e-10);
}
#[test]
fn invalid_beta_rejected() {
assert!(Temperature::beta(-1.0).is_err());
assert!(Temperature::beta(0.0).is_err());
}
#[test]
fn annealing_schedule() {
let temp = Temperature::annealed(1.0, 100.0, 10).unwrap();
let beta_0 = temp.beta_at(0);
let beta_5 = temp.beta_at(5);
let beta_9 = temp.beta_at(9);
assert!(beta_0 < beta_5);
assert!(beta_5 < beta_9);
assert!((beta_0 - 1.0).abs() < 0.1);
assert!((temp.beta_at(10) - 100.0).abs() < 1.0);
}
#[test]
fn schedule_iteration() {
let mut schedule = TemperatureSchedule::exponential(1.0, 10.0, 5);
let collected: Vec<f64> = schedule.by_ref().collect();
assert_eq!(collected.len(), 5);
for window in collected.windows(2) {
assert!(window[0] < window[1]);
}
}
#[test]
fn cosine_schedule() {
let schedule = TemperatureSchedule::cosine(1.0, 10.0, 10);
let temps = schedule.temperatures();
assert!((temps[0] - 1.0).abs() < 1e-10);
assert!((temps[9] - 10.0).abs() < 1e-10);
let early_delta = temps[1] - temps[0];
let mid_delta = temps[5] - temps[4];
assert!(mid_delta > early_delta);
}
#[cfg(feature = "serde")]
#[test]
fn temperature_serde_roundtrip() {
let temps = vec![
Temperature::Soft,
Temperature::Hard,
Temperature::Beta(2.5),
Temperature::annealed(1.0, 100.0, 10).unwrap(),
];
for original in temps {
let encoded = bincode::serialize(&original).expect("serialize");
let decoded: Temperature = bincode::deserialize(&encoded).expect("deserialize");
assert_eq!(original.num_steps(), decoded.num_steps());
for i in 0..12 {
assert!((original.beta_at(i) - decoded.beta_at(i)).abs() < 1e-9);
}
}
}
#[test]
fn to_resonator_config_mapping() {
let default_max = amari_holographic::ResonatorConfig::default().max_iterations;
let soft = Temperature::soft().to_resonator_config();
assert_eq!(soft.initial_beta, 1.0);
assert_eq!(soft.final_beta, 1.0);
assert_eq!(soft.max_iterations, default_max);
let hard = Temperature::hard().to_resonator_config();
assert_eq!(hard.initial_beta, f64::MAX);
assert_eq!(hard.final_beta, f64::MAX);
let beta = Temperature::beta(3.0).unwrap().to_resonator_config();
assert_eq!(beta.initial_beta, 3.0);
assert_eq!(beta.final_beta, 3.0);
let annealed = Temperature::annealed(2.0, 50.0, 17)
.unwrap()
.to_resonator_config();
assert_eq!(annealed.initial_beta, 2.0);
assert_eq!(annealed.final_beta, 50.0);
assert_eq!(annealed.max_iterations, 17);
}
}