use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ConvergenceStatus {
NotConverged,
Converged(ConvergenceReason),
}
impl ConvergenceStatus {
pub fn is_converged(&self) -> bool {
matches!(self, Self::Converged(_))
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum ConvergenceReason {
FitnessStagnation { generations: usize },
LowDiversity { diversity: u64 }, TargetReached { target: u64 }, MaxGenerations { generations: usize },
MaxEvaluations { evaluations: usize },
RhatConverged { rhat: u64 }, MultipleReasons(Vec<ConvergenceReason>),
Custom(String),
}
impl ConvergenceReason {
pub fn fitness_stagnation(generations: usize) -> Self {
Self::FitnessStagnation { generations }
}
pub fn low_diversity(diversity: f64) -> Self {
Self::LowDiversity {
diversity: diversity.to_bits(),
}
}
pub fn target_reached(target: f64) -> Self {
Self::TargetReached {
target: target.to_bits(),
}
}
pub fn rhat_converged(rhat: f64) -> Self {
Self::RhatConverged {
rhat: rhat.to_bits(),
}
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct ConvergenceConfig {
pub max_generations: Option<usize>,
pub max_evaluations: Option<usize>,
pub target_fitness: Option<f64>,
pub target_tolerance: f64,
pub stagnation_generations: usize,
pub stagnation_threshold: f64,
pub diversity_threshold: f64,
pub rhat_threshold: f64,
pub use_rhat: bool,
}
impl Default for ConvergenceConfig {
fn default() -> Self {
Self {
max_generations: None,
max_evaluations: None,
target_fitness: None,
target_tolerance: 1e-6,
stagnation_generations: 50,
stagnation_threshold: 1e-9,
diversity_threshold: 0.01,
rhat_threshold: 1.1,
use_rhat: false,
}
}
}
impl ConvergenceConfig {
pub fn with_max_generations(generations: usize) -> Self {
Self {
max_generations: Some(generations),
..Default::default()
}
}
pub fn max_generations(mut self, generations: usize) -> Self {
self.max_generations = Some(generations);
self
}
pub fn max_evaluations(mut self, evaluations: usize) -> Self {
self.max_evaluations = Some(evaluations);
self
}
pub fn target_fitness(mut self, target: f64) -> Self {
self.target_fitness = Some(target);
self
}
pub fn target_tolerance(mut self, tolerance: f64) -> Self {
self.target_tolerance = tolerance;
self
}
pub fn stagnation(mut self, generations: usize, threshold: f64) -> Self {
self.stagnation_generations = generations;
self.stagnation_threshold = threshold;
self
}
pub fn diversity_threshold(mut self, threshold: f64) -> Self {
self.diversity_threshold = threshold;
self
}
pub fn with_rhat(mut self, threshold: f64) -> Self {
self.use_rhat = true;
self.rhat_threshold = threshold;
self
}
}
#[derive(Clone, Debug)]
pub struct ConvergenceDetector {
config: ConvergenceConfig,
best_fitness_history: Vec<f64>,
mean_fitness_history: Vec<f64>,
diversity_history: Vec<f64>,
current_generation: usize,
current_evaluations: usize,
best_fitness_overall: f64,
running_best_fitness: f64,
last_improvement_generation: usize,
}
fn two_pass_mean_var(xs: &[f64]) -> (f64, f64) {
let n = xs.len() as f64;
let mean = xs.iter().sum::<f64>() / n;
if xs.len() < 2 {
return (mean, 0.0);
}
let ss: f64 = xs.iter().map(|x| (x - mean).powi(2)).sum();
(mean, ss / (n - 1.0))
}
impl ConvergenceDetector {
pub fn new(config: ConvergenceConfig) -> Self {
Self {
config,
best_fitness_history: Vec::new(),
mean_fitness_history: Vec::new(),
diversity_history: Vec::new(),
current_generation: 0,
current_evaluations: 0,
best_fitness_overall: f64::NEG_INFINITY,
running_best_fitness: f64::NEG_INFINITY,
last_improvement_generation: 0,
}
}
pub fn with_defaults() -> Self {
Self::new(ConvergenceConfig::default())
}
pub fn update(
&mut self,
generation: usize,
evaluations: usize,
best_fitness: f64,
mean_fitness: f64,
diversity: f64,
) {
self.current_generation = generation;
self.current_evaluations = evaluations;
self.best_fitness_history.push(best_fitness);
self.mean_fitness_history.push(mean_fitness);
self.diversity_history.push(diversity);
if best_fitness > self.running_best_fitness {
self.running_best_fitness = best_fitness;
}
if best_fitness > self.best_fitness_overall + self.config.stagnation_threshold {
self.best_fitness_overall = best_fitness;
self.last_improvement_generation = generation;
}
}
pub fn check(&self) -> ConvergenceStatus {
let mut reasons = Vec::new();
if let Some(max_gen) = self.config.max_generations {
if self.current_generation >= max_gen {
reasons.push(ConvergenceReason::MaxGenerations {
generations: self.current_generation,
});
}
}
if let Some(max_eval) = self.config.max_evaluations {
if self.current_evaluations >= max_eval {
reasons.push(ConvergenceReason::MaxEvaluations {
evaluations: self.current_evaluations,
});
}
}
if let Some(target) = self.config.target_fitness {
if !self.best_fitness_history.is_empty() {
let best = self.running_best_fitness;
if (best - target).abs() <= self.config.target_tolerance || best >= target {
reasons.push(ConvergenceReason::target_reached(best));
}
}
}
let generations_since_improvement =
self.current_generation - self.last_improvement_generation;
if generations_since_improvement >= self.config.stagnation_generations {
reasons.push(ConvergenceReason::fitness_stagnation(
generations_since_improvement,
));
}
if let Some(&diversity) = self.diversity_history.last() {
if diversity < self.config.diversity_threshold {
reasons.push(ConvergenceReason::low_diversity(diversity));
}
}
if self.config.use_rhat && self.mean_fitness_history.len() >= 10 {
let rhat = self.compute_rhat();
if rhat < self.config.rhat_threshold {
reasons.push(ConvergenceReason::rhat_converged(rhat));
}
}
match reasons.len() {
0 => ConvergenceStatus::NotConverged,
1 => ConvergenceStatus::Converged(reasons.pop().unwrap()),
_ => ConvergenceStatus::Converged(ConvergenceReason::MultipleReasons(reasons)),
}
}
fn compute_rhat(&self) -> f64 {
let n = self.mean_fitness_history.len();
let l = n / 2;
if l < 5 {
return f64::INFINITY; }
let l_f = l as f64;
let (mean1, var1) = two_pass_mean_var(&self.mean_fitness_history[0..l]);
let (mean2, var2) = two_pass_mean_var(&self.mean_fitness_history[l..2 * l]);
let m = 2.0;
let grand_mean = (mean1 + mean2) / m;
let b = l_f / (m - 1.0) * ((mean1 - grand_mean).powi(2) + (mean2 - grand_mean).powi(2));
let w = (var1 + var2) / m;
if w <= 0.0 {
return 1.0; }
let var_plus = ((l_f - 1.0) / l_f) * w + b / l_f;
(var_plus / w).sqrt()
}
pub fn best_fitness(&self) -> f64 {
self.running_best_fitness
}
pub fn generations_without_improvement(&self) -> usize {
self.current_generation - self.last_improvement_generation
}
pub fn current_diversity(&self) -> Option<f64> {
self.diversity_history.last().copied()
}
pub fn fitness_history(&self) -> &[f64] {
&self.best_fitness_history
}
pub fn diversity_history(&self) -> &[f64] {
&self.diversity_history
}
pub fn reset(&mut self) {
self.best_fitness_history.clear();
self.mean_fitness_history.clear();
self.diversity_history.clear();
self.current_generation = 0;
self.current_evaluations = 0;
self.best_fitness_overall = f64::NEG_INFINITY;
self.running_best_fitness = f64::NEG_INFINITY;
self.last_improvement_generation = 0;
}
}
pub fn evolutionary_rhat(runs: &[Vec<f64>]) -> f64 {
if runs.is_empty() || runs[0].is_empty() {
return f64::INFINITY;
}
let m = runs.len() as f64;
let n_len = runs.iter().map(|r| r.len()).min().unwrap_or(0);
let n = n_len as f64;
if n < 2.0 || m < 2.0 {
return f64::INFINITY;
}
let chain_means: Vec<f64> = runs
.iter()
.map(|r| r[..n_len].iter().sum::<f64>() / n)
.collect();
let grand_mean = chain_means.iter().sum::<f64>() / m;
let b = n / (m - 1.0)
* chain_means
.iter()
.map(|cm| (cm - grand_mean).powi(2))
.sum::<f64>();
let w: f64 = runs
.iter()
.map(|r| {
let chain = &r[..n_len];
let mean = chain.iter().sum::<f64>() / n;
chain.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / (n - 1.0)
})
.sum::<f64>()
/ m;
if w == 0.0 {
return 1.0; }
let var_plus = ((n - 1.0) / n) * w + b / n;
(var_plus / w).sqrt()
}
pub fn evolutionary_ess(weights: &[f64]) -> f64 {
if weights.is_empty() {
return 0.0;
}
let sum: f64 = weights.iter().sum();
if sum == 0.0 {
return weights.len() as f64;
}
let normalized: Vec<f64> = weights.iter().map(|w| w / sum).collect();
let sum_sq: f64 = normalized.iter().map(|w| w * w).sum();
if sum_sq == 0.0 {
weights.len() as f64
} else {
1.0 / sum_sq
}
}
pub fn evolutionary_ess_log(log_weights: &[f64]) -> f64 {
if log_weights.is_empty() {
return 0.0;
}
let max_log = log_weights
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
if max_log.is_infinite() {
return log_weights.len() as f64;
}
let weights: Vec<f64> = log_weights.iter().map(|lw| (lw - max_log).exp()).collect();
evolutionary_ess(&weights)
}
pub fn detect_stagnation(fitness_history: &[f64], threshold: f64) -> usize {
if fitness_history.len() < 2 {
return 0;
}
let best = fitness_history
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
let mut stagnant_count: usize = 0;
for &fitness in fitness_history.iter().rev() {
if (fitness - best).abs() <= threshold {
stagnant_count += 1;
} else {
break;
}
}
stagnant_count.saturating_sub(1) }
pub fn fitness_convergence(fitness_values: &[f64]) -> f64 {
if fitness_values.len() < 2 {
return 1.0;
}
let mean = fitness_values.iter().sum::<f64>() / fitness_values.len() as f64;
if mean.abs() < f64::EPSILON {
return 1.0;
}
let variance = fitness_values
.iter()
.map(|f| (f - mean).powi(2))
.sum::<f64>()
/ (fitness_values.len() - 1) as f64;
let std = variance.sqrt();
let cv = std / mean.abs();
(-cv).exp()
}
#[derive(Clone, Debug)]
pub struct TerminationCriteria {
criteria: Vec<TerminationCriterion>,
require_all: bool,
}
#[derive(Clone, Debug)]
pub enum TerminationCriterion {
MaxGenerations(usize),
MaxEvaluations(usize),
TargetFitness(f64, f64), Stagnation(usize, f64), DiversityThreshold(f64),
TimeLimit(f64),
Custom(String), }
impl TerminationCriteria {
pub fn new() -> Self {
Self {
criteria: Vec::new(),
require_all: false,
}
}
pub fn require_all() -> Self {
Self {
criteria: Vec::new(),
require_all: true,
}
}
pub fn add(mut self, criterion: TerminationCriterion) -> Self {
self.criteria.push(criterion);
self
}
pub fn max_generations(self, generations: usize) -> Self {
self.add(TerminationCriterion::MaxGenerations(generations))
}
pub fn max_evaluations(self, evaluations: usize) -> Self {
self.add(TerminationCriterion::MaxEvaluations(evaluations))
}
pub fn target_fitness(self, target: f64, tolerance: f64) -> Self {
self.add(TerminationCriterion::TargetFitness(target, tolerance))
}
pub fn stagnation(self, generations: usize, threshold: f64) -> Self {
self.add(TerminationCriterion::Stagnation(generations, threshold))
}
pub fn diversity_threshold(self, threshold: f64) -> Self {
self.add(TerminationCriterion::DiversityThreshold(threshold))
}
pub fn time_limit(self, seconds: f64) -> Self {
self.add(TerminationCriterion::TimeLimit(seconds))
}
pub fn should_terminate(
&self,
generation: usize,
evaluations: usize,
best_fitness: f64,
diversity: f64,
fitness_history: &[f64],
elapsed_seconds: f64,
) -> Option<ConvergenceReason> {
let mut satisfied = Vec::new();
for criterion in &self.criteria {
let met = match criterion {
TerminationCriterion::MaxGenerations(max) => generation >= *max,
TerminationCriterion::MaxEvaluations(max) => evaluations >= *max,
TerminationCriterion::TargetFitness(target, tolerance) => {
(best_fitness - target).abs() <= *tolerance || best_fitness >= *target
}
TerminationCriterion::Stagnation(gens, threshold) => {
detect_stagnation(fitness_history, *threshold) >= *gens
}
TerminationCriterion::DiversityThreshold(thresh) => diversity < *thresh,
TerminationCriterion::TimeLimit(limit) => elapsed_seconds >= *limit,
TerminationCriterion::Custom(_) => false, };
if met {
satisfied.push(criterion.to_reason(
generation,
evaluations,
best_fitness,
diversity,
));
}
}
if satisfied.is_empty() {
return None;
}
if self.require_all && satisfied.len() < self.criteria.len() {
return None;
}
if satisfied.len() == 1 {
Some(satisfied.pop().unwrap())
} else {
Some(ConvergenceReason::MultipleReasons(satisfied))
}
}
pub fn criteria(&self) -> &[TerminationCriterion] {
&self.criteria
}
}
impl Default for TerminationCriteria {
fn default() -> Self {
Self::new()
}
}
impl TerminationCriterion {
fn to_reason(
&self,
generation: usize,
evaluations: usize,
best_fitness: f64,
diversity: f64,
) -> ConvergenceReason {
match self {
Self::MaxGenerations(_) => ConvergenceReason::MaxGenerations {
generations: generation,
},
Self::MaxEvaluations(_) => ConvergenceReason::MaxEvaluations { evaluations },
Self::TargetFitness(_, _) => ConvergenceReason::target_reached(best_fitness),
Self::Stagnation(gens, _) => ConvergenceReason::fitness_stagnation(*gens),
Self::DiversityThreshold(_) => ConvergenceReason::low_diversity(diversity),
Self::TimeLimit(t) => ConvergenceReason::Custom(format!("Time limit of {t}s reached")),
Self::Custom(desc) => ConvergenceReason::Custom(desc.clone()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convergence_detector_basic() {
let config = ConvergenceConfig::with_max_generations(100);
let mut detector = ConvergenceDetector::new(config);
for i in 0..50 {
detector.update(i, i * 10, i as f64, i as f64 * 0.5, 0.5);
}
let status = detector.check();
assert!(!status.is_converged());
}
#[test]
fn test_convergence_detector_max_generations() {
let config = ConvergenceConfig::with_max_generations(50);
let mut detector = ConvergenceDetector::new(config);
for i in 0..60 {
detector.update(i, i * 10, i as f64, i as f64 * 0.5, 0.5);
}
let status = detector.check();
assert!(status.is_converged());
if let ConvergenceStatus::Converged(reason) = status {
assert!(matches!(reason, ConvergenceReason::MaxGenerations { .. }));
}
}
#[test]
fn test_convergence_detector_target_fitness() {
let config = ConvergenceConfig::default()
.target_fitness(100.0)
.target_tolerance(1.0);
let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 99.5, 50.0, 0.5);
let status = detector.check();
assert!(status.is_converged());
}
#[test]
fn test_convergence_detector_stagnation() {
let config = ConvergenceConfig::default().stagnation(10, 1e-9);
let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 50.0, 50.0, 0.5);
for i in 1..20 {
detector.update(i, i * 10, 50.0, 50.0, 0.5);
}
let status = detector.check();
assert!(status.is_converged());
if let ConvergenceStatus::Converged(reason) = status {
assert!(matches!(
reason,
ConvergenceReason::FitnessStagnation { .. }
));
}
}
#[test]
fn test_convergence_detector_low_diversity() {
let config = ConvergenceConfig::default().diversity_threshold(0.1);
let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 50.0, 50.0, 0.05);
let status = detector.check();
assert!(status.is_converged());
if let ConvergenceStatus::Converged(reason) = status {
assert!(matches!(reason, ConvergenceReason::LowDiversity { .. }));
}
}
#[test]
fn test_evolutionary_rhat() {
let chain1 = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let chain2 = vec![1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5, 10.5];
let rhat = evolutionary_rhat(&[chain1, chain2]);
assert!(rhat < 1.2, "R-hat was {}, expected < 1.2", rhat);
}
#[test]
fn test_evolutionary_rhat_divergent() {
let chain1 = vec![1.0, 2.0, 1.5, 2.5, 1.2, 2.8, 1.8, 2.2, 1.3, 2.7];
let chain2 = vec![
100.0, 101.0, 100.5, 101.5, 100.2, 101.8, 100.8, 101.2, 100.3, 101.7,
];
let rhat = evolutionary_rhat(&[chain1, chain2]);
assert!(rhat > 1.5, "R-hat was {}, expected > 1.5", rhat);
}
#[test]
fn test_evolutionary_ess() {
let weights = vec![1.0, 1.0, 1.0, 1.0];
let ess = evolutionary_ess(&weights);
assert!((ess - 4.0).abs() < 0.01);
}
#[test]
fn test_evolutionary_ess_unequal() {
let weights = vec![1.0, 0.0, 0.0, 0.0];
let ess = evolutionary_ess(&weights);
assert!((ess - 1.0).abs() < 0.01);
}
#[test]
fn test_evolutionary_ess_log() {
let log_weights = vec![0.0, 0.0, 0.0, 0.0];
let ess = evolutionary_ess_log(&log_weights);
assert!((ess - 4.0).abs() < 0.01);
}
#[test]
fn test_detect_stagnation() {
let history = vec![10.0, 20.0, 30.0, 30.0, 30.0, 30.0];
let stagnant = detect_stagnation(&history, 1e-9);
assert_eq!(stagnant, 3); }
#[test]
fn test_detect_stagnation_improving() {
let history = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let stagnant = detect_stagnation(&history, 1e-9);
assert_eq!(stagnant, 0);
}
#[test]
fn test_fitness_convergence() {
let fitness = vec![50.0, 50.0, 50.0, 50.0];
let conv = fitness_convergence(&fitness);
assert!((conv - 1.0).abs() < 0.01);
}
#[test]
fn test_fitness_convergence_diverse() {
let fitness = vec![0.0, 100.0, 0.0, 100.0];
let conv = fitness_convergence(&fitness);
assert!(conv < 0.5);
}
#[test]
fn test_termination_criteria_max_gen() {
let criteria = TerminationCriteria::new().max_generations(100);
let result = criteria.should_terminate(50, 500, 10.0, 0.5, &[], 10.0);
assert!(result.is_none());
let result = criteria.should_terminate(100, 1000, 10.0, 0.5, &[], 20.0);
assert!(result.is_some());
}
#[test]
fn test_termination_criteria_target() {
let criteria = TerminationCriteria::new().target_fitness(100.0, 1.0);
let result = criteria.should_terminate(10, 100, 50.0, 0.5, &[], 5.0);
assert!(result.is_none());
let result = criteria.should_terminate(10, 100, 99.5, 0.5, &[], 5.0);
assert!(result.is_some());
}
#[test]
fn test_termination_criteria_multiple() {
let criteria = TerminationCriteria::new()
.max_generations(100)
.target_fitness(100.0, 1.0);
let result = criteria.should_terminate(10, 100, 50.0, 0.5, &[], 5.0);
assert!(result.is_none());
let result = criteria.should_terminate(10, 100, 100.0, 0.5, &[], 5.0);
assert!(result.is_some());
let result = criteria.should_terminate(100, 1000, 50.0, 0.5, &[], 50.0);
assert!(result.is_some());
}
#[test]
fn test_termination_criteria_require_all() {
let criteria = TerminationCriteria::require_all()
.max_generations(100)
.stagnation(10, 1e-9);
let short_flat = [50.0; 6];
let result = criteria.should_terminate(100, 1000, 50.0, 0.5, &short_flat, 50.0);
assert!(result.is_none());
let long_flat = [50.0; 11];
let result = criteria.should_terminate(100, 1000, 50.0, 0.5, &long_flat, 50.0);
assert!(result.is_some());
}
#[test]
fn test_convergence_config_builder() {
let config = ConvergenceConfig::with_max_generations(500)
.max_evaluations(10000)
.target_fitness(1.0)
.target_tolerance(0.01)
.stagnation(100, 1e-6)
.diversity_threshold(0.05)
.with_rhat(1.05);
assert_eq!(config.max_generations, Some(500));
assert_eq!(config.max_evaluations, Some(10000));
assert_eq!(config.target_fitness, Some(1.0));
assert_eq!(config.target_tolerance, 0.01);
assert_eq!(config.stagnation_generations, 100);
assert_eq!(config.stagnation_threshold, 1e-6);
assert_eq!(config.diversity_threshold, 0.05);
assert!(config.use_rhat);
assert_eq!(config.rhat_threshold, 1.05);
}
#[test]
fn test_convergence_detector_reset() {
let config = ConvergenceConfig::default();
let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 50.0, 50.0, 0.5);
detector.update(1, 20, 60.0, 55.0, 0.4);
assert_eq!(detector.fitness_history().len(), 2);
assert_eq!(detector.best_fitness(), 60.0);
detector.reset();
assert!(detector.fitness_history().is_empty());
assert_eq!(detector.best_fitness(), f64::NEG_INFINITY);
}
#[test]
fn test_convergence_status_is_converged() {
let not_converged = ConvergenceStatus::NotConverged;
assert!(!not_converged.is_converged());
let converged =
ConvergenceStatus::Converged(ConvergenceReason::MaxGenerations { generations: 100 });
assert!(converged.is_converged());
}
#[test]
fn test_evolutionary_rhat_truncates_unequal_chains() {
let chain1 = vec![1.0, 2.0, 3.0, 4.0];
let chain2 = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let rhat = evolutionary_rhat(&[chain1, chain2]);
let expected = 0.75_f64.sqrt();
assert!(
(rhat - expected).abs() < 1e-9,
"R-hat was {rhat}, expected {expected}"
);
assert!(
rhat < 0.9,
"R-hat {rhat} still shows the unequal-length bug"
);
}
#[test]
fn test_target_fitness_uses_running_best() {
let config = ConvergenceConfig::default()
.target_fitness(100.0)
.target_tolerance(1e-6)
.stagnation(10_000, 1e-9); let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 100.0, 50.0, 1.0); detector.update(1, 20, 50.0, 50.0, 1.0);
assert_eq!(detector.best_fitness(), 100.0);
let status = detector.check();
assert!(
status.is_converged(),
"a target reached earlier must remain converged"
);
if let ConvergenceStatus::Converged(reason) = status {
assert!(matches!(reason, ConvergenceReason::TargetReached { .. }));
}
}
#[test]
fn test_target_fitness_survives_stagnation_throttle() {
let config = ConvergenceConfig::default()
.target_fitness(100.0)
.target_tolerance(1e-6)
.stagnation(50, 0.01);
let mut detector = ConvergenceDetector::new(config);
detector.update(0, 10, 99.995, 50.0, 1.0);
detector.update(1, 20, 100.0, 50.0, 1.0);
assert_eq!(
detector.best_fitness(),
100.0,
"the running best must reflect the true max, not the throttled value"
);
let status = detector.check();
assert!(
status.is_converged(),
"target reached-and-held must be reported even when \
stagnation_threshold > target_tolerance"
);
assert!(
matches!(
status,
ConvergenceStatus::Converged(ConvergenceReason::TargetReached { .. })
| ConvergenceStatus::Converged(ConvergenceReason::MultipleReasons(_))
),
"convergence reason must include TargetReached, got {status:?}"
);
for g in 2..10 {
detector.update(g, 10 * (g + 1), 100.0, 50.0, 1.0);
let s = detector.check();
let has_target = match &s {
ConvergenceStatus::Converged(ConvergenceReason::TargetReached { .. }) => true,
ConvergenceStatus::Converged(ConvergenceReason::MultipleReasons(rs)) => rs
.iter()
.any(|r| matches!(r, ConvergenceReason::TargetReached { .. })),
_ => false,
};
assert!(
has_target,
"target must stay reported while held, gen {g}: {s:?}"
);
}
}
#[test]
fn test_stagnation_threshold_is_wired() {
let history = [10.0, 9.5, 9.5, 9.5, 9.5];
let loose = TerminationCriteria::new().stagnation(3, 1.0);
let tight = TerminationCriteria::new().stagnation(3, 0.1);
assert!(
loose
.should_terminate(0, 0, 9.5, 1.0, &history, 0.0)
.is_some(),
"loose threshold should read the flat tail as stagnant"
);
assert!(
tight
.should_terminate(0, 0, 9.5, 1.0, &history, 0.0)
.is_none(),
"tight threshold should not read a 0.5 drop as stagnant"
);
}
#[test]
fn test_compute_rhat_matches_naive_recompute() {
let mut detector = ConvergenceDetector::with_defaults();
let values: Vec<f64> = (0..40)
.map(|i| {
let x = i as f64;
(x * 0.37).sin() * 3.0 + (x * 0.11).cos() * 1.5 + x * 0.05
})
.collect();
for (i, &v) in values.iter().enumerate() {
detector.update(i, i * 10, v, v, 0.5);
if detector.mean_fitness_history.len() >= 10 {
let n = detector.mean_fitness_history.len();
let half = n / 2;
let chain1 = detector.mean_fitness_history[..half].to_vec();
let chain2 = detector.mean_fitness_history[half..].to_vec();
let naive = evolutionary_rhat(&[chain1, chain2]);
let incremental = detector.compute_rhat();
assert!(
(naive - incremental).abs() < 1e-9,
"at n={n}: incremental R-hat {incremental} != naive {naive}"
);
}
}
}
#[test]
fn test_compute_rhat_stable_under_large_offset() {
let offset = 1e6;
let chain1: Vec<f64> = (0..20)
.map(|i| offset + (i as f64 * 0.7).sin() * 1.3)
.collect();
let chain2: Vec<f64> = (0..20)
.map(|i| offset + 4.0 + (i as f64 * 0.9 + 0.5).cos() * 1.1)
.collect();
let mut detector = ConvergenceDetector::with_defaults();
for (i, &v) in chain1.iter().chain(chain2.iter()).enumerate() {
detector.update(i, i, v, v, 0.5);
}
let got = detector.compute_rhat();
let tp = |xs: &[f64]| -> (f64, f64) {
let n = xs.len() as f64;
let mean = xs.iter().sum::<f64>() / n;
let ss: f64 = xs.iter().map(|x| (x - mean).powi(2)).sum();
(mean, ss / (n - 1.0))
};
let (mean1, var1) = tp(&chain1);
let (mean2, var2) = tp(&chain2);
let l = chain1.len() as f64;
let m = 2.0;
let grand = (mean1 + mean2) / m;
let b = l / (m - 1.0) * ((mean1 - grand).powi(2) + (mean2 - grand).powi(2));
let w = (var1 + var2) / m;
assert!(w > 0.0, "true within-chain variance must be positive");
let var_plus = ((l - 1.0) / l) * w + b / l;
let reference = (var_plus / w).sqrt();
assert!(
(got - reference).abs() < 1e-9,
"compute_rhat {got} != directly computed two-pass R-hat {reference}"
);
assert!(
(got - 1.0).abs() > 1e-6,
"compute_rhat collapsed to the spurious 1.0 (got {got})"
);
assert!(got.is_finite(), "compute_rhat must be finite, got {got}");
let via_public = evolutionary_rhat(&[chain1, chain2]);
assert!(
(got - via_public).abs() < 1e-9,
"compute_rhat {got} != evolutionary_rhat {via_public}"
);
}
}