use std::fmt;
use rand::Rng;
use rand::SeedableRng;
use rand::rngs::StdRng;
use rand::seq::SliceRandom;
use crate::rules::{RewriteRule, Rule};
use crate::state::BinaryGraphState;
pub trait Schedule: fmt::Debug + Send + Sync {
fn name(&self) -> &str;
fn timing(&self) -> &str;
fn selection(&self) -> &str;
fn step(
&self,
state: &BinaryGraphState,
rules: &[RewriteRule],
rng: &mut impl Rng,
) -> BinaryGraphState;
}
#[derive(Debug, Clone)]
pub struct AllVerticesSchedule;
impl AllVerticesSchedule {
pub fn new() -> Self {
Self
}
}
impl Default for AllVerticesSchedule {
fn default() -> Self {
Self::new()
}
}
impl Schedule for AllVerticesSchedule {
fn name(&self) -> &str {
"all_vertices"
}
fn timing(&self) -> &str {
"asynchronous"
}
fn selection(&self) -> &str {
"exhaustive"
}
fn step(
&self,
state: &BinaryGraphState,
rules: &[RewriteRule],
rng: &mut impl Rng,
) -> BinaryGraphState {
let mut current = state.clone();
let n = current.n_vertices();
let mut vertices: Vec<usize> = (0..n).collect();
vertices.shuffle(rng);
for &vertex in &vertices {
let mut rule_indices: Vec<usize> = (0..rules.len()).collect();
rule_indices.shuffle(rng);
for &ri in &rule_indices {
if let Some(info) = rules[ri].matches(¤t, vertex) {
current = rules[ri].apply(¤t, &info, rng);
break; }
}
}
current
}
}
pub static DEFAULT_SCHEDULE: AllVerticesSchedule = AllVerticesSchedule;
pub fn generate_trajectory<O>(
initial_state: &BinaryGraphState,
rules: &[RewriteRule],
steps: usize,
window_size: usize,
schedule: &impl Schedule,
obs_fn: &dyn Fn(&[BinaryGraphState]) -> O,
seed: u64,
) -> Vec<O> {
let mut rng = StdRng::seed_from_u64(seed);
let mut state_history = Vec::with_capacity(window_size);
state_history.push(initial_state.clone());
let mut observations = Vec::with_capacity(steps + 1);
observations.push(obs_fn(&state_history));
let mut current = initial_state.clone();
for _ in 0..steps {
current = schedule.step(¤t, rules, &mut rng);
state_history.push(current.clone());
if state_history.len() > window_size {
state_history.remove(0);
}
observations.push(obs_fn(&state_history));
}
observations
}
#[allow(clippy::needless_range_loop)]
#[allow(clippy::too_many_arguments)]
pub fn generate_ensemble<O: Clone>(
initial_states: &[BinaryGraphState],
rules: &[RewriteRule],
steps: usize,
n_ensemble: usize,
window_size: usize,
schedule: &impl Schedule,
obs_fn: &dyn Fn(&[BinaryGraphState]) -> O,
base_seed: u64,
) -> Vec<Vec<O>> {
assert!(n_ensemble >= 2, "Ensemble size must be at least 2");
assert!(
initial_states.len() >= n_ensemble,
"Need at least {} initial states, got {}",
n_ensemble,
initial_states.len()
);
let mut trajectories = Vec::with_capacity(n_ensemble);
for i in 0..n_ensemble {
let seed = base_seed + i as u64 * 137;
let traj = generate_trajectory(
&initial_states[i],
rules,
steps,
window_size,
schedule,
obs_fn,
seed,
);
trajectories.push(traj);
}
trajectories
}
pub fn test_boolean_function(
rules: &[RewriteRule],
truth_table: &[((u8, u8), u8)],
steps: usize,
n_trials: usize,
schedule: &impl Schedule,
) -> bool {
use ndarray::{arr1, arr2};
for &((a, b), expected) in truth_table {
let mut results = Vec::new();
for trial in 0..n_trials {
let adj = arr2(&[[0, 0, 1], [0, 0, 1], [0, 0, 0]]);
let labels = arr1(&[a as i8, b as i8, 0i8]);
let mut state = BinaryGraphState::new(3, adj.view(), labels.view())
.expect("Boolean test state creation");
let mut rng = StdRng::seed_from_u64((trial * 100) as u64);
for _ in 0..steps {
state = schedule.step(&state, rules, &mut rng);
}
results.push(state.label(2));
}
let ones = results.iter().filter(|&&x| x == 1).count();
let observed = if ones > results.len() / 2 { 1 } else { 0 };
if observed != expected {
return false;
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{rules::create_structured_rules, state::State};
use ndarray::{arr1, arr2};
fn make_test_state() -> BinaryGraphState {
let adj = arr2(&[[0, 1, 0], [1, 0, 1], [0, 0, 0]]);
let labels = arr1(&[1, 0, 1]);
BinaryGraphState::new(3, adj.view(), labels.view()).unwrap()
}
#[test]
fn test_all_vertices_schedule_applies_rules() {
let state = make_test_state();
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let mut rng = StdRng::seed_from_u64(42);
let new_state = schedule.step(&state, &rules, &mut rng);
assert_eq!(new_state.n_vertices(), 3);
for i in 0..3 {
assert!(new_state.label(i) <= 1);
}
}
#[test]
fn test_trajectory_deterministic() {
let state = make_test_state();
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let obs_fn = |s: &[BinaryGraphState]| s[0].canonical_encoding();
let traj1 = generate_trajectory(&state, &rules, 10, 1, &schedule, &obs_fn, 42);
let traj2 = generate_trajectory(&state, &rules, 10, 1, &schedule, &obs_fn, 42);
assert_eq!(traj1.len(), traj2.len());
for (o1, o2) in traj1.iter().zip(traj2.iter()) {
assert_eq!(o1, o2);
}
}
#[test]
fn test_trajectory_different_seeds_diverge() {
let state = make_test_state();
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let obs_fn = |s: &[BinaryGraphState]| s[0].canonical_encoding();
let traj1 = generate_trajectory(&state, &rules, 20, 1, &schedule, &obs_fn, 42);
let traj2 = generate_trajectory(&state, &rules, 20, 1, &schedule, &obs_fn, 99);
let all_same = traj1.iter().zip(traj2.iter()).all(|(a, b)| a == b);
assert!(
!all_same,
"Different seeds should produce different trajectories"
);
}
#[test]
fn test_ensemble_size() {
let state = make_test_state();
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let obs_fn = |s: &[BinaryGraphState]| s[0].canonical_encoding();
let initial_states = vec![state.clone(), state.clone(), state.clone()];
let ensemble = generate_ensemble(&initial_states, &rules, 10, 3, 1, &schedule, &obs_fn, 42);
assert_eq!(ensemble.len(), 3);
for traj in &ensemble {
assert_eq!(traj.len(), 11); }
}
#[test]
#[should_panic(expected = "Ensemble size must be at least 2")]
fn test_ensemble_too_small() {
let state = make_test_state();
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let obs_fn = |s: &[BinaryGraphState]| s[0].canonical_encoding();
let initial_states = vec![state.clone()];
generate_ensemble(&initial_states, &rules, 10, 1, 1, &schedule, &obs_fn, 42);
}
#[test]
fn test_nand_boolean_function() {
let rules = create_structured_rules();
let schedule = AllVerticesSchedule::new();
let truth_table = vec![((0, 0), 1), ((0, 1), 1), ((1, 0), 1), ((1, 1), 0)];
let nand_rules: Vec<RewriteRule> = rules
.into_iter()
.filter(|r| r.name() == "NAND" || r.name() == "IDENTITY")
.collect();
let result = test_boolean_function(&nand_rules, &truth_table, 10, 5, &schedule);
let _ = result;
}
}