use rand::RngExt;
use rand::SeedableRng;
use rand::rngs::StdRng;
use crate::dynamics::{DEFAULT_SCHEDULE, generate_ensemble};
use crate::metrics::{compute_memory, compute_persistence, compute_storage};
use crate::rules::RewriteRule;
use crate::rules::Rule;
use crate::rules::create_destructive_rules;
use crate::state::BinaryGraphState;
#[allow(clippy::too_many_arguments)]
pub fn generate_null_trajectories<O: Clone>(
n_universes: usize,
n_vertices: usize,
n_ensemble: usize,
steps: usize,
window_size: usize,
max_rules_per_subset: usize,
obs_fn: &dyn Fn(&[BinaryGraphState]) -> O,
seed: u64,
) -> Vec<Vec<Vec<O>>> {
let mut rng = StdRng::seed_from_u64(seed);
let destructive_pool = create_destructive_rules();
let scramblers: Vec<&RewriteRule> = destructive_pool
.iter()
.filter(|r| r.name().starts_with("DESTROY_SCRAMBLE_ALL"))
.collect();
let state_pool: Vec<BinaryGraphState> = (0..n_ensemble * 2)
.map(|_| BinaryGraphState::random(n_vertices, &mut rng))
.collect();
let mut null_ensembles = Vec::with_capacity(n_universes);
for i in 0..n_universes {
let size = rng.random_range(1..=max_rules_per_subset);
let mut rules = Vec::with_capacity(size);
let scrambler_idx = rng.random_range(0..scramblers.len());
rules.push(scramblers[scrambler_idx].clone());
for _ in 1..size {
let idx = rng.random_range(0..destructive_pool.len());
rules.push(destructive_pool[idx].clone());
}
for j in (1..rules.len()).rev() {
let k = rng.random_range(0..=j);
rules.swap(j, k);
}
let initial_states: Vec<BinaryGraphState> =
state_pool.iter().take(n_ensemble).cloned().collect();
let ensemble = generate_ensemble(
&initial_states,
&rules,
steps,
n_ensemble,
window_size,
&DEFAULT_SCHEDULE,
obs_fn,
seed + i as u64 * 1000,
);
null_ensembles.push(ensemble);
}
null_ensembles
}
#[derive(Debug, Clone)]
pub struct CalibrationResult {
pub persistence_threshold: f64,
pub storage_threshold: f64,
pub memory_threshold: f64,
pub null_persistence: NullStats,
pub null_storage: NullStats,
pub null_memory: NullStats,
pub percentile: f64,
}
#[derive(Debug, Clone)]
pub struct NullStats {
pub mean: f64,
pub std: f64,
pub scores: Vec<f64>,
}
impl NullStats {
pub fn empirical_p(&self, observed_value: f64) -> f64 {
let count = self.scores.iter().filter(|&&s| s >= observed_value).count();
if self.scores.is_empty() {
1.0
} else {
count as f64 / self.scores.len() as f64
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn calibrate_thresholds<O: Eq + std::hash::Hash + Clone>(
null_ensembles: &[Vec<Vec<O>>],
percentile: f64,
floor_persistence: f64,
floor_storage: f64,
floor_memory: f64,
max_delta: usize,
n_shuffles: usize,
seed: u64,
) -> CalibrationResult {
let mut persistence_scores = Vec::with_capacity(null_ensembles.len());
let mut storage_scores = Vec::with_capacity(null_ensembles.len());
let mut memory_scores = Vec::with_capacity(null_ensembles.len());
for ensemble in null_ensembles {
persistence_scores.push(compute_persistence(ensemble, 1, n_shuffles, seed));
storage_scores.push(compute_storage(ensemble, max_delta, n_shuffles, seed));
memory_scores.push(compute_memory(ensemble, max_delta, n_shuffles, seed));
}
let persistence_threshold =
percentile_value(&persistence_scores, percentile).max(floor_persistence);
let storage_threshold = percentile_value(&storage_scores, percentile).max(floor_storage);
let memory_threshold = percentile_value(&memory_scores, percentile).max(floor_memory);
CalibrationResult {
persistence_threshold,
storage_threshold,
memory_threshold,
null_persistence: NullStats::from_scores(persistence_scores),
null_storage: NullStats::from_scores(storage_scores),
null_memory: NullStats::from_scores(memory_scores),
percentile,
}
}
fn percentile_value(scores: &[f64], percentile: f64) -> f64 {
if scores.is_empty() {
return 0.0;
}
if scores.len() == 1 {
return scores[0];
}
let mut sorted: Vec<f64> = scores.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let k = (percentile / 100.0) * (sorted.len() - 1) as f64;
let lo = k.floor() as usize;
let hi = k.ceil() as usize;
if lo == hi {
sorted[lo]
} else {
let frac = k - lo as f64;
sorted[lo] * (1.0 - frac) + sorted[hi] * frac
}
}
impl NullStats {
fn from_scores(scores: Vec<f64>) -> Self {
let n = scores.len() as f64;
let mean = if n > 0.0 {
scores.iter().sum::<f64>() / n
} else {
0.0
};
let std = if n > 1.0 {
let variance = scores.iter().map(|s| (s - mean).powi(2)).sum::<f64>() / (n - 1.0);
variance.sqrt()
} else {
0.0
};
Self { mean, std, scores }
}
}
#[allow(clippy::too_many_arguments)]
pub fn calibrate<O: Eq + std::hash::Hash + Clone>(
n_null_universes: usize,
n_vertices: usize,
n_ensemble: usize,
steps: usize,
window_size: usize,
max_rules_per_subset: usize,
obs_fn: &dyn Fn(&[BinaryGraphState]) -> O,
percentile: f64,
max_delta: usize,
n_shuffles: usize,
seed: u64,
) -> CalibrationResult {
let null_ensembles = generate_null_trajectories(
n_null_universes,
n_vertices,
n_ensemble,
steps,
window_size,
max_rules_per_subset,
obs_fn,
seed,
);
calibrate_thresholds(
&null_ensembles,
percentile,
0.01, 0.01, 0.01, max_delta,
n_shuffles,
seed,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::observation::observe_windowed;
#[test]
fn test_generate_null_trajectories() {
let nulls = generate_null_trajectories(10, 3, 5, 20, 1, 3, &observe_windowed, 42);
assert_eq!(nulls.len(), 10);
for ensemble in &nulls {
assert_eq!(ensemble.len(), 5);
for traj in ensemble {
assert_eq!(traj.len(), 21);
}
}
}
#[test]
fn test_calibrate_thresholds_produces_valid_result() {
let nulls = generate_null_trajectories(20, 3, 5, 20, 1, 3, &observe_windowed, 42);
let result = calibrate_thresholds(&nulls, 95.0, 0.01, 0.01, 0.01, 10, 5, 42);
assert!(result.persistence_threshold >= 0.01);
assert!(result.storage_threshold >= 0.01);
assert!(result.memory_threshold >= 0.01);
assert_eq!(result.percentile, 95.0);
assert!(!result.null_storage.scores.is_empty());
}
#[test]
fn test_null_stats_empirical_p() {
let stats = NullStats {
mean: 0.1,
std: 0.05,
scores: vec![0.05, 0.08, 0.10, 0.12, 0.15],
};
assert!((stats.empirical_p(0.12) - 0.4).abs() < 0.01);
assert_eq!(stats.empirical_p(1.0), 0.0);
assert_eq!(stats.empirical_p(0.0), 1.0);
}
#[test]
fn test_percentile_value() {
let scores = vec![0.0, 0.2, 0.4, 0.6, 0.8, 1.0];
assert!((percentile_value(&scores, 50.0) - 0.5).abs() < 0.01);
assert!((percentile_value(&scores, 95.0) - 0.95).abs() < 0.01);
assert!((percentile_value(&scores, 0.0) - 0.0).abs() < 0.01);
}
#[test]
fn test_calibrate_convenience() {
let result = calibrate(10, 3, 5, 20, 1, 3, &observe_windowed, 95.0, 10, 5, 42);
assert!(result.storage_threshold >= 0.01);
}
}