use rand::Rng;
use rand_distr::{Beta, Distribution, Gamma};
use crate::operators::mutation::{
BitFlipMutation, GaussianMutation, PolynomialMutation, UniformMutation,
};
#[derive(Clone, Debug)]
pub struct BetaPosterior {
pub alpha: f64,
pub beta: f64,
pub alpha0: f64,
pub beta0: f64,
}
impl BetaPosterior {
pub fn uniform() -> Self {
Self::new(1.0, 1.0)
}
pub fn jeffreys() -> Self {
Self::new(0.5, 0.5)
}
pub fn new(alpha: f64, beta: f64) -> Self {
Self {
alpha,
beta,
alpha0: alpha,
beta0: beta,
}
}
pub fn observe_success(&mut self) {
self.alpha += 1.0;
}
pub fn observe_failure(&mut self) {
self.beta += 1.0;
}
pub fn observe(&mut self, success: bool) {
if success {
self.observe_success();
} else {
self.observe_failure();
}
}
pub fn mean(&self) -> f64 {
self.alpha / (self.alpha + self.beta)
}
pub fn mode(&self) -> Option<f64> {
if self.alpha > 1.0 && self.beta > 1.0 {
Some((self.alpha - 1.0) / (self.alpha + self.beta - 2.0))
} else {
None
}
}
pub fn variance(&self) -> f64 {
let sum = self.alpha + self.beta;
(self.alpha * self.beta) / (sum * sum * (sum + 1.0))
}
pub fn std_dev(&self) -> f64 {
self.variance().sqrt()
}
pub fn sample<R: Rng>(&self, rng: &mut R) -> f64 {
Beta::new(self.alpha, self.beta)
.expect("Invalid Beta parameters")
.sample(rng)
}
pub fn credible_interval(&self, probability: f64) -> (f64, f64) {
let mean = self.mean();
let std = self.std_dev();
let z = normal_quantile((1.0 + probability) / 2.0);
let lower = (mean - z * std).max(0.0);
let upper = (mean + z * std).min(1.0);
(lower, upper)
}
pub fn observations(&self) -> f64 {
(self.alpha - self.alpha0) + (self.beta - self.beta0)
}
pub fn decay(&mut self, factor: f64) {
self.alpha = self.alpha0 + factor * (self.alpha - self.alpha0);
self.beta = self.beta0 + factor * (self.beta - self.beta0);
}
}
impl Default for BetaPosterior {
fn default() -> Self {
Self::uniform()
}
}
#[derive(Clone, Debug)]
pub struct GammaPosterior {
pub shape: f64,
pub rate: f64,
pub shape0: f64,
pub rate0: f64,
}
impl GammaPosterior {
pub fn vague() -> Self {
Self::new(1.0, 0.01)
}
pub fn new(shape: f64, rate: f64) -> Self {
Self {
shape,
rate,
shape0: shape,
rate0: rate,
}
}
pub fn observe(&mut self, value: f64) {
self.shape += 1.0;
self.rate += value;
}
pub fn mean(&self) -> f64 {
self.shape / self.rate
}
pub fn posterior_mean_of_mean(&self) -> Option<f64> {
if self.shape > 1.0 {
Some(self.rate / (self.shape - 1.0))
} else {
None
}
}
pub fn mode(&self) -> Option<f64> {
if self.shape >= 1.0 {
Some((self.shape - 1.0) / self.rate)
} else {
None
}
}
pub fn variance(&self) -> f64 {
self.shape / (self.rate * self.rate)
}
pub fn observations(&self) -> f64 {
self.shape - self.shape0
}
pub fn sample<R: Rng>(&self, rng: &mut R) -> f64 {
Gamma::new(self.shape, 1.0 / self.rate)
.expect("Invalid Gamma parameters")
.sample(rng)
}
pub fn decay(&mut self, factor: f64) {
self.shape = self.shape0 + factor * (self.shape - self.shape0);
self.rate = self.rate0 + factor * (self.rate - self.rate0);
}
}
impl Default for GammaPosterior {
fn default() -> Self {
Self::vague()
}
}
#[derive(Clone, Debug, Default)]
pub struct RunningLogMoments {
mean_log: f64,
m2: f64,
n: usize,
}
impl RunningLogMoments {
pub fn new() -> Self {
Self::default()
}
pub fn observe(&mut self, x: f64) {
if x <= 0.0 {
return;
}
let log_x = x.ln();
self.n += 1;
let delta = log_x - self.mean_log;
self.mean_log += delta / self.n as f64;
let delta2 = log_x - self.mean_log;
self.m2 += delta * delta2;
}
pub fn count(&self) -> usize {
self.n
}
pub fn mean_log(&self) -> f64 {
self.mean_log
}
pub fn var_log(&self) -> f64 {
if self.n >= 1 {
self.m2 / self.n as f64
} else {
0.0
}
}
pub fn sample_var_log(&self) -> Option<f64> {
if self.n >= 2 {
Some(self.m2 / (self.n as f64 - 1.0))
} else {
None
}
}
pub fn mean(&self) -> f64 {
(self.mean_log + self.var_log() / 2.0).exp()
}
pub fn mode(&self) -> f64 {
(self.mean_log - self.var_log()).exp()
}
pub fn sample<R: Rng>(&self, rng: &mut R) -> f64 {
use rand_distr::StandardNormal;
let z: f64 = rng.sample(StandardNormal);
(self.mean_log + self.var_log().sqrt() * z).exp()
}
}
pub trait TunableMutation {
fn set_mutation_probability(&mut self, probability: f64);
}
macro_rules! impl_tunable_mutation {
($($ty:ty),+ $(,)?) => {
$(
impl TunableMutation for $ty {
fn set_mutation_probability(&mut self, probability: f64) {
self.mutation_probability = Some(probability.clamp(0.0, 1.0));
}
}
)+
};
}
impl_tunable_mutation!(
PolynomialMutation,
GaussianMutation,
UniformMutation,
BitFlipMutation,
);
#[derive(Clone, Debug)]
pub struct BanditArm {
pub value: f64,
pub posterior: BetaPosterior,
pub selections: u64,
}
#[derive(Clone, Debug)]
pub struct BanditParameter {
pub name: String,
arms: Vec<BanditArm>,
last_selected: Option<usize>,
}
impl BanditParameter {
pub fn new(name: impl Into<String>, values: Vec<f64>) -> Self {
Self::with_prior(name, values, BetaPosterior::uniform())
}
pub fn with_prior(name: impl Into<String>, values: Vec<f64>, prior: BetaPosterior) -> Self {
assert!(
!values.is_empty(),
"BanditParameter requires at least one arm value"
);
let arms = values
.into_iter()
.map(|value| BanditArm {
value,
posterior: prior.clone(),
selections: 0,
})
.collect();
Self {
name: name.into(),
arms,
last_selected: None,
}
}
pub fn select<R: Rng>(&mut self, rng: &mut R) -> f64 {
let mut best_idx = 0;
let mut best_draw = f64::NEG_INFINITY;
for (i, arm) in self.arms.iter().enumerate() {
let draw = arm.posterior.sample(rng);
if draw > best_draw {
best_draw = draw;
best_idx = i;
}
}
self.last_selected = Some(best_idx);
self.arms[best_idx].selections += 1;
self.arms[best_idx].value
}
pub fn observe(&mut self, improved: bool) {
if let Some(idx) = self.last_selected {
self.arms[idx].posterior.observe(improved);
}
}
pub fn arms(&self) -> &[BanditArm] {
&self.arms
}
pub fn values(&self) -> Vec<f64> {
self.arms.iter().map(|a| a.value).collect()
}
pub fn posterior_means(&self) -> Vec<f64> {
self.arms.iter().map(|a| a.posterior.mean()).collect()
}
pub fn selection_counts(&self) -> Vec<u64> {
self.arms.iter().map(|a| a.selections).collect()
}
pub fn selected_value(&self) -> Option<f64> {
self.last_selected.map(|i| self.arms[i].value)
}
pub fn selected_index(&self) -> Option<usize> {
self.last_selected
}
pub fn best_index(&self) -> usize {
self.arms
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| {
a.posterior
.mean()
.partial_cmp(&b.posterior.mean())
.unwrap_or(std::cmp::Ordering::Equal)
})
.map(|(i, _)| i)
.unwrap_or(0)
}
pub fn best_value(&self) -> f64 {
self.arms[self.best_index()].value
}
pub fn total_observations(&self) -> f64 {
self.arms.iter().map(|a| a.posterior.observations()).sum()
}
}
pub const PARAM_MUTATION_RATE: &str = "mutation_rate";
pub const PARAM_CROSSOVER_PROB: &str = "crossover_prob";
#[derive(Clone, Debug)]
pub struct ThompsonConfig {
pub mutation_rate_arms: Vec<f64>,
pub crossover_prob_arms: Vec<f64>,
pub prior: BetaPosterior,
pub record_history: bool,
}
impl Default for ThompsonConfig {
fn default() -> Self {
Self {
mutation_rate_arms: vec![0.01, 0.05, 0.1, 0.2, 0.4],
crossover_prob_arms: vec![0.5, 0.7, 0.9],
prior: BetaPosterior::uniform(),
record_history: false,
}
}
}
impl ThompsonConfig {
pub fn build_tuner(&self) -> ThompsonSamplingTuner {
ThompsonSamplingTuner::from_config(self)
}
}
#[derive(Clone, Debug)]
pub struct TunerSnapshot {
pub generation: usize,
pub parameters: Vec<(String, Option<f64>, Vec<f64>)>,
}
#[derive(Clone, Debug)]
pub struct ThompsonSamplingTuner {
parameters: Vec<BanditParameter>,
record_history: bool,
history: Vec<TunerSnapshot>,
observations: u64,
}
impl ThompsonSamplingTuner {
pub fn new(parameters: Vec<BanditParameter>) -> Self {
Self {
parameters,
record_history: false,
history: Vec::new(),
observations: 0,
}
}
pub fn from_config(cfg: &ThompsonConfig) -> Self {
let mut parameters = Vec::new();
if !cfg.mutation_rate_arms.is_empty() {
parameters.push(BanditParameter::with_prior(
PARAM_MUTATION_RATE,
cfg.mutation_rate_arms.clone(),
cfg.prior.clone(),
));
}
if !cfg.crossover_prob_arms.is_empty() {
parameters.push(BanditParameter::with_prior(
PARAM_CROSSOVER_PROB,
cfg.crossover_prob_arms.clone(),
cfg.prior.clone(),
));
}
Self {
parameters,
record_history: cfg.record_history,
history: Vec::new(),
observations: 0,
}
}
pub fn with_history(mut self, on: bool) -> Self {
self.record_history = on;
self
}
pub fn parameters(&self) -> &[BanditParameter] {
&self.parameters
}
pub fn parameter(&self, name: &str) -> Option<&BanditParameter> {
self.parameters.iter().find(|p| p.name == name)
}
pub fn parameter_mut(&mut self, name: &str) -> Option<&mut BanditParameter> {
self.parameters.iter_mut().find(|p| p.name == name)
}
pub fn is_empty(&self) -> bool {
self.parameters.is_empty()
}
pub fn select_all<R: Rng>(&mut self, rng: &mut R) {
for p in &mut self.parameters {
p.select(rng);
}
}
pub fn selected(&self, name: &str) -> Option<f64> {
self.parameter(name).and_then(|p| p.selected_value())
}
pub fn observe(&mut self, improved: bool) {
for p in &mut self.parameters {
p.observe(improved);
}
self.observations += 1;
}
pub fn total_observations(&self) -> u64 {
self.observations
}
pub fn snapshot(&mut self, generation: usize) {
if !self.record_history {
return;
}
let parameters = self
.parameters
.iter()
.map(|p| (p.name.clone(), p.selected_value(), p.posterior_means()))
.collect();
self.history.push(TunerSnapshot {
generation,
parameters,
});
}
pub fn history(&self) -> &[TunerSnapshot] {
&self.history
}
}
fn normal_quantile(p: f64) -> f64 {
if p <= 0.0 {
return f64::NEG_INFINITY;
}
if p >= 1.0 {
return f64::INFINITY;
}
let t = if p < 0.5 {
(-2.0 * p.ln()).sqrt()
} else {
(-2.0 * (1.0 - p).ln()).sqrt()
};
let c0 = 2.515517;
let c1 = 0.802853;
let c2 = 0.010328;
let d1 = 1.432788;
let d2 = 0.189269;
let d3 = 0.001308;
let q = t - (c0 + c1 * t + c2 * t * t) / (1.0 + d1 * t + d2 * t * t + d3 * t * t * t);
if p < 0.5 {
-q
} else {
q
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::SeedableRng;
#[test]
fn test_beta_posterior_uniform_prior() {
let posterior = BetaPosterior::uniform();
assert!((posterior.mean() - 0.5).abs() < 1e-10);
}
#[test]
fn test_beta_posterior_update() {
let mut posterior = BetaPosterior::uniform();
for _ in 0..7 {
posterior.observe_success();
}
for _ in 0..3 {
posterior.observe_failure();
}
assert!((posterior.mean() - 0.667).abs() < 0.01);
}
#[test]
fn test_beta_posterior_sample() {
let posterior = BetaPosterior::new(5.0, 5.0);
let mut rng = rand::thread_rng();
for _ in 0..100 {
let sample = posterior.sample(&mut rng);
assert!((0.0..=1.0).contains(&sample));
}
}
#[test]
fn test_beta_observations_uses_stored_prior() {
let mut uniform = BetaPosterior::uniform();
for i in 0..10 {
uniform.observe(i % 2 == 0);
}
assert!((uniform.observations() - 10.0).abs() < 1e-12);
let mut jeffreys = BetaPosterior::jeffreys();
for _ in 0..5 {
jeffreys.observe_success();
}
for _ in 0..5 {
jeffreys.observe_failure();
}
assert!((jeffreys.observations() - 10.0).abs() < 1e-12);
let mut informative = BetaPosterior::new(2.0, 2.0);
for _ in 0..4 {
informative.observe_success();
}
assert!((informative.observations() - 4.0).abs() < 1e-12);
}
#[test]
fn test_gamma_rate_and_mean_of_mean() {
let mut posterior = GammaPosterior::new(2.0, 1.0);
for x in [1.0, 2.0, 3.0] {
posterior.observe(x);
}
assert!((posterior.shape - 5.0).abs() < 1e-12);
assert!((posterior.rate - 7.0).abs() < 1e-12);
assert!((posterior.mean() - 5.0 / 7.0).abs() < 1e-12);
assert!((posterior.posterior_mean_of_mean().unwrap() - 1.75).abs() < 1e-12);
assert!((posterior.observations() - 3.0).abs() < 1e-12);
}
#[test]
fn test_gamma_recovers_mean_not_reciprocal() {
let mut posterior = GammaPosterior::vague();
for _ in 0..1000 {
posterior.observe(20.0);
}
assert!((posterior.mean() - 0.05).abs() < 0.005, "rate ~ 1/20");
let mean = posterior.posterior_mean_of_mean().unwrap();
assert!((mean - 20.0).abs() < 0.5, "mean-of-mean ~ 20, got {mean}");
}
#[test]
fn test_running_log_moments_no_prior_contamination() {
let mut moments = RunningLogMoments::new();
moments.observe(std::f64::consts::E); assert!((moments.var_log()).abs() < 1e-12);
assert!((moments.mean_log() - 1.0).abs() < 1e-12);
moments.observe(std::f64::consts::E.powi(3)); assert!((moments.mean_log() - 2.0).abs() < 1e-12);
assert!((moments.var_log() - 1.0).abs() < 1e-12);
assert!((moments.sample_var_log().unwrap() - 2.0).abs() < 1e-12);
}
#[test]
fn test_running_log_moments_mean_original_space() {
let mut moments = RunningLogMoments::new();
for _ in 0..10 {
moments.observe(0.1);
}
assert!((moments.mean() - 0.1).abs() < 1e-12);
}
#[test]
fn test_bandit_parameter_thompson_selects_a_value() {
let mut param = BanditParameter::new(PARAM_MUTATION_RATE, vec![0.01, 0.1, 0.3]);
let mut rng = StdRng::seed_from_u64(1);
let v = param.select(&mut rng);
assert!([0.01, 0.1, 0.3].contains(&v));
param.observe(true);
assert!((param.total_observations() - 1.0).abs() < 1e-12);
}
#[test]
fn test_bandit_concentrates_on_better_arm() {
let mut rng = StdRng::seed_from_u64(20260710);
let mut param = BanditParameter::new(PARAM_MUTATION_RATE, vec![0.01, 0.3]);
let true_p = |v: f64| if v >= 0.3 { 0.55 } else { 0.20 };
let total_rounds = 2000;
let late_start = 1500;
let mut late_good = 0u32;
let mut late_total = 0u32;
for round in 0..total_rounds {
let value = param.select(&mut rng);
let improved = rng.gen::<f64>() < true_p(value);
param.observe(improved);
if round >= late_start {
late_total += 1;
if value >= 0.3 {
late_good += 1;
}
}
}
let frac = late_good as f64 / late_total as f64;
assert!(
frac > 0.70,
"expected >70% of late pulls on the better arm, got {:.2}",
frac
);
assert!((param.best_value() - 0.3).abs() < 1e-12);
}
#[test]
fn test_thompson_tuner_from_config() {
let cfg = ThompsonConfig::default();
let mut tuner = cfg.build_tuner();
assert!(tuner.parameter(PARAM_MUTATION_RATE).is_some());
assert!(tuner.parameter(PARAM_CROSSOVER_PROB).is_some());
let mut rng = StdRng::seed_from_u64(7);
tuner.select_all(&mut rng);
assert!(tuner.selected(PARAM_MUTATION_RATE).is_some());
assert!(tuner.selected(PARAM_CROSSOVER_PROB).is_some());
tuner.observe(true);
tuner.observe(false);
assert_eq!(tuner.total_observations(), 2);
}
#[test]
fn test_thompson_tuner_parameters_are_independent() {
let cfg = ThompsonConfig {
mutation_rate_arms: vec![0.1, 0.3],
crossover_prob_arms: vec![0.5, 0.9],
prior: BetaPosterior::uniform(),
record_history: false,
};
let mut tuner = cfg.build_tuner();
let mut rng = StdRng::seed_from_u64(99);
for _ in 0..50 {
tuner.select_all(&mut rng);
tuner.observe(true);
}
let mr = tuner.parameter(PARAM_MUTATION_RATE).unwrap();
let cx = tuner.parameter(PARAM_CROSSOVER_PROB).unwrap();
assert_eq!(mr.values(), vec![0.1, 0.3]);
assert_eq!(cx.values(), vec![0.5, 0.9]);
}
#[test]
fn test_tunable_mutation_sets_probability() {
let mut m = GaussianMutation::new(0.1);
m.set_mutation_probability(0.25);
assert_eq!(m.mutation_probability, Some(0.25));
m.set_mutation_probability(5.0);
assert_eq!(m.mutation_probability, Some(1.0));
}
#[test]
fn test_credible_interval() {
let posterior = BetaPosterior::new(50.0, 50.0);
let (lower, upper) = posterior.credible_interval(0.95);
assert!(lower < 0.5);
assert!(upper > 0.5);
assert!(lower > 0.0);
assert!(upper < 1.0);
}
}