use alloc::collections::VecDeque;
use alloc::vec::Vec;
use crate::statistics::{OnlineStats, StatsSnapshot};
use super::{kl_divergence_gaussian, CalibrationSnapshot, Posterior};
pub struct AdaptiveState {
pub baseline_samples: Vec<u64>,
pub sample_samples: Vec<u64>,
pub previous_posterior: Option<Posterior>,
pub recent_kl_divergences: VecDeque<f64>,
pub batch_count: usize,
baseline_stats: OnlineStats,
sample_stats: OnlineStats,
ns_per_tick: Option<f64>,
}
impl AdaptiveState {
pub fn new() -> Self {
Self {
baseline_samples: Vec::new(),
sample_samples: Vec::new(),
previous_posterior: None,
recent_kl_divergences: VecDeque::with_capacity(5),
batch_count: 0,
baseline_stats: OnlineStats::new(),
sample_stats: OnlineStats::new(),
ns_per_tick: None,
}
}
pub fn with_capacity(expected_samples: usize) -> Self {
Self {
baseline_samples: Vec::with_capacity(expected_samples),
sample_samples: Vec::with_capacity(expected_samples),
previous_posterior: None,
recent_kl_divergences: VecDeque::with_capacity(5),
batch_count: 0,
baseline_stats: OnlineStats::new(),
sample_stats: OnlineStats::new(),
ns_per_tick: None,
}
}
pub fn n_total(&self) -> usize {
self.baseline_samples.len()
}
pub fn add_batch(&mut self, baseline: Vec<u64>, sample: Vec<u64>) {
debug_assert_eq!(
baseline.len(),
sample.len(),
"Baseline and sample batch sizes must match"
);
self.baseline_samples.extend(baseline);
self.sample_samples.extend(sample);
self.batch_count += 1;
}
pub fn add_batch_with_conversion(
&mut self,
baseline: Vec<u64>,
sample: Vec<u64>,
ns_per_tick: f64,
) {
debug_assert_eq!(
baseline.len(),
sample.len(),
"Baseline and sample batch sizes must match"
);
self.ns_per_tick = Some(ns_per_tick);
for &t in &baseline {
self.baseline_stats.update(t as f64 * ns_per_tick);
}
for &t in &sample {
self.sample_stats.update(t as f64 * ns_per_tick);
}
self.baseline_samples.extend(baseline);
self.sample_samples.extend(sample);
self.batch_count += 1;
}
pub fn update_kl(&mut self, kl: f64) {
self.recent_kl_divergences.push_back(kl);
if self.recent_kl_divergences.len() > 5 {
self.recent_kl_divergences.pop_front();
}
}
pub fn recent_kl_sum(&self) -> f64 {
self.recent_kl_divergences.iter().sum()
}
pub fn has_kl_history(&self) -> bool {
self.recent_kl_divergences.len() >= 5
}
pub fn update_posterior(&mut self, new_posterior: Posterior) -> f64 {
let kl = if let Some(ref prev) = self.previous_posterior {
kl_divergence_gaussian(&new_posterior, prev)
} else {
0.0
};
self.previous_posterior = Some(new_posterior);
if kl.is_finite() {
self.update_kl(kl);
}
kl
}
pub fn current_posterior(&self) -> Option<&Posterior> {
self.previous_posterior.as_ref()
}
pub fn baseline_ns(&self, ns_per_tick: f64) -> Vec<f64> {
self.baseline_samples
.iter()
.map(|&t| t as f64 * ns_per_tick)
.collect()
}
pub fn sample_ns(&self, ns_per_tick: f64) -> Vec<f64> {
self.sample_samples
.iter()
.map(|&t| t as f64 * ns_per_tick)
.collect()
}
pub fn baseline_stats(&self) -> Option<StatsSnapshot> {
if self.baseline_stats.count() < 2 {
return None;
}
Some(self.baseline_stats.finalize())
}
pub fn sample_stats(&self) -> Option<StatsSnapshot> {
if self.sample_stats.count() < 2 {
return None;
}
Some(self.sample_stats.finalize())
}
pub fn get_stats_snapshot(&self) -> Option<CalibrationSnapshot> {
let baseline = self.baseline_stats()?;
let sample = self.sample_stats()?;
Some(CalibrationSnapshot::new(baseline, sample))
}
pub fn has_stats_tracking(&self) -> bool {
self.ns_per_tick.is_some() && self.baseline_stats.count() > 0
}
pub fn reset(&mut self) {
self.baseline_samples.clear();
self.sample_samples.clear();
self.previous_posterior = None;
self.recent_kl_divergences.clear();
self.batch_count = 0;
self.baseline_stats = OnlineStats::new();
self.sample_stats = OnlineStats::new();
self.ns_per_tick = None;
}
}
impl Default for AdaptiveState {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{Matrix9, Vector9};
fn make_test_posterior(leak_prob: f64, n: usize) -> Posterior {
Posterior::new(
Vector9::zeros(),
Matrix9::identity(),
Vec::new(), leak_prob,
100.0, n,
)
}
#[test]
fn test_adaptive_state_new() {
let state = AdaptiveState::new();
assert_eq!(state.n_total(), 0);
assert_eq!(state.batch_count, 0);
assert!(state.previous_posterior.is_none());
assert!(!state.has_kl_history());
}
#[test]
fn test_add_batch() {
let mut state = AdaptiveState::new();
state.add_batch(vec![100, 101, 102], vec![200, 201, 202]);
assert_eq!(state.n_total(), 3);
assert_eq!(state.batch_count, 1);
assert_eq!(state.baseline_samples, vec![100, 101, 102]);
assert_eq!(state.sample_samples, vec![200, 201, 202]);
}
#[test]
fn test_kl_history() {
let mut state = AdaptiveState::new();
for i in 0..5 {
state.update_kl(0.1 * (i + 1) as f64);
}
assert!(state.has_kl_history());
assert!((state.recent_kl_sum() - 1.5).abs() < 1e-10);
state.update_kl(1.0);
assert!((state.recent_kl_sum() - 2.4).abs() < 1e-10); }
#[test]
fn test_posterior_update() {
let mut state = AdaptiveState::new();
let posterior1 = make_test_posterior(0.75, 1000);
let kl1 = state.update_posterior(posterior1.clone());
assert_eq!(kl1, 0.0);
assert!(state.current_posterior().is_some());
let posterior2 = make_test_posterior(0.80, 2000);
let kl2 = state.update_posterior(posterior2);
assert!(kl2 >= 0.0);
}
#[test]
fn test_add_batch_with_conversion() {
let mut state = AdaptiveState::new();
state.add_batch_with_conversion(vec![100, 110, 120], vec![200, 210, 220], 2.0);
assert_eq!(state.n_total(), 3);
assert_eq!(state.batch_count, 1);
assert!(state.has_stats_tracking());
assert_eq!(state.baseline_samples, vec![100, 110, 120]);
assert_eq!(state.sample_samples, vec![200, 210, 220]);
}
#[test]
fn test_online_stats_tracking() {
let mut state = AdaptiveState::new();
let baseline: Vec<u64> = (0..100).map(|i| 1000 + (i % 10)).collect();
let sample: Vec<u64> = (0..100).map(|i| 1100 + (i % 10)).collect();
state.add_batch_with_conversion(baseline, sample, 1.0);
let baseline_stats = state.baseline_stats().expect("Should have baseline stats");
assert_eq!(baseline_stats.count, 100);
assert!(
(baseline_stats.mean - 1004.5).abs() < 1.0,
"Baseline mean {} should be near 1004.5",
baseline_stats.mean
);
let sample_stats = state.sample_stats().expect("Should have sample stats");
assert_eq!(sample_stats.count, 100);
assert!(
(sample_stats.mean - 1104.5).abs() < 1.0,
"Sample mean {} should be near 1104.5",
sample_stats.mean
);
}
#[test]
fn test_reset() {
let mut state = AdaptiveState::new();
state.add_batch_with_conversion(vec![100, 110], vec![200, 210], 1.0);
state.update_kl(0.5);
let posterior = make_test_posterior(0.75, 100);
state.update_posterior(posterior);
assert!(state.n_total() > 0);
state.reset();
assert_eq!(state.n_total(), 0);
assert_eq!(state.batch_count, 0);
assert!(state.previous_posterior.is_none());
assert!(!state.has_kl_history());
assert!(!state.has_stats_tracking());
}
#[test]
fn test_stats_not_tracked_without_conversion() {
let mut state = AdaptiveState::new();
state.add_batch(vec![100, 110, 120], vec![200, 210, 220]);
assert!(!state.has_stats_tracking());
assert!(state.baseline_stats().is_none());
assert!(state.sample_stats().is_none());
}
}