use std::collections::{BTreeSet, HashMap};
use rand::{seq::{IndexedMutRandom, IndexedRandom}, Rng};
use serde::{Serialize, Deserialize};
use rand_distr::{Distribution, Normal};
#[derive(Serialize, Deserialize, Debug)]
pub struct GlobalInnovator {
pub innov: usize,
}
impl GlobalInnovator {
pub fn new() -> Self {
GlobalInnovator { innov: 0 }
}
pub fn next(&mut self) -> usize {
let innov = self.innov;
self.innov += 1;
innov
}
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy)]
pub struct NodeGene {
pub id: usize,
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy)]
pub struct ConnectionGene {
pub in_node: usize,
pub out_node: usize,
pub weight: f64,
pub enabled: bool,
pub innov: usize,
}
impl ConnectionGene {
pub fn get_id(&self) -> (usize, usize) {
(self.in_node, self.out_node)
}
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct Genome {
pub num_inputs: usize,
pub num_outputs: usize,
pub node_genes: Vec<NodeGene>,
pub connection_genes: Vec<ConnectionGene>,
}
const CONNECTION_MUTATION_RATE: f64 = 0.15;
const NODE_MUTATION_RATE: f64 = 0.03; const WEIGHT_MUTATION_RATE: f64 = 0.8;
const PERTUBATION_CHANCE: f64 = 0.9; const PERTUBATION_STD: f64 = 0.1;
const REPLACEMENT_RANGE: f64 = 5.0;
const TOGGLE_MUTATION_RATE: f64 = 0.01;
impl Genome {
pub fn new(num_inputs: usize, num_outputs: usize) -> Self {
let node_genes = (0..(num_inputs + 1 + num_outputs)).into_iter()
.map(|i| NodeGene { id: i})
.collect();
Genome {
num_inputs: num_inputs + 1,
num_outputs,
node_genes,
connection_genes: Vec::new(),
}
}
pub fn crossover(fit_parent: &Genome, unfit_parent: &Genome) -> Genome {
let mut child_connections: Vec<ConnectionGene> = Vec::new();
let mut fitter_map = HashMap::new();
for conn in &fit_parent.connection_genes {
fitter_map.insert(conn.innov, conn);
}
let mut unfit_map = HashMap::new();
for conn in &unfit_parent.connection_genes {
unfit_map.insert(conn.innov, conn);
}
let all_innovs: BTreeSet<usize> = fitter_map.keys().chain(unfit_map.keys()).cloned().collect();
for innov in all_innovs { match (fitter_map.get(&innov), unfit_map.get(&innov)) {
(Some(&a), Some(&b)) => {
child_connections.push(if rand::random() { a.clone() } else { b.clone() });
}
(Some(&a), None) => {
child_connections.push(a.clone());
}
(None, Some(_b)) => {
}
(None, None) => unreachable!(),
}
}
Genome {
num_inputs: fit_parent.num_inputs,
num_outputs: fit_parent.num_outputs,
node_genes: fit_parent.node_genes.clone(), connection_genes: child_connections,
}
}
pub fn mutate(&mut self, innovator: &mut GlobalInnovator, innovations: &mut HashMap<(usize, usize), usize>) {
let mut rng = rand::rng();
self.mutate_weights_and_toggle();
if rng.random::<f64>() < CONNECTION_MUTATION_RATE {
self.add_connection(innovator, innovations);
}
if rng.random::<f64>() < NODE_MUTATION_RATE {
self.add_node(innovator, innovations);
}
}
fn mutate_weights_and_toggle(&mut self) {
let mut rng = rand::rng();
let normal = Normal::new(0.0, PERTUBATION_STD).unwrap();
for connection in &mut self.connection_genes {
if rng.random::<f64>() < TOGGLE_MUTATION_RATE {
connection.enabled = !connection.enabled;
}
if rng.random::<f64>() > WEIGHT_MUTATION_RATE { continue;
}
if rng.random::<f64>() < PERTUBATION_CHANCE {
let pertub_amount = normal.sample(&mut rng);
connection.weight += pertub_amount;
} else {
let new_weight = rng.random_range(-REPLACEMENT_RANGE..REPLACEMENT_RANGE);
connection.weight = new_weight;
}
}
}
fn add_node(&mut self, innovator: &mut GlobalInnovator, innovations: &mut HashMap<(usize, usize), usize>) {
if self.connection_genes.len() < 1 {
return;
}
let mut rng = rand::rng();
let mut collected = self.connection_genes.iter_mut()
.filter(|x| x.enabled)
.collect::<Vec<&mut ConnectionGene>>();
let chosen = collected.choose_mut(&mut rng).unwrap();
let new_id = self.node_genes.last().unwrap().id + 1;
self.node_genes.push(NodeGene { id: new_id});
let innov0;
let innov1;
match innovations.get(&(chosen.in_node, new_id)) {
Some(x) => { innov0 = *x;
},
None => { innov0 = innovator.next();
innovations.insert((chosen.in_node, new_id), innov0);
},
}
match innovations.get(&(new_id, chosen.out_node)) {
Some(x) => { innov1 = *x;
},
None => { innov1 = innovator.next();
innovations.insert((new_id, chosen.out_node), innov1);
},
}
let connection_0 = ConnectionGene {
in_node: chosen.in_node,
out_node: new_id,
weight: chosen.weight, enabled: true,
innov: innov0,
};
let connection_1 = ConnectionGene {
in_node: new_id,
out_node: chosen.out_node,
weight: 1.0, enabled: true,
innov: innov1,
};
chosen.enabled = false;
chosen.weight = 1.0;
self.connection_genes.push(connection_0);
self.connection_genes.push(connection_1);
self.connection_genes.sort_by_key(|c| c.innov);
self.node_genes.sort_by_key(|n| n.id);
}
fn add_connection(&mut self, innovator: &mut GlobalInnovator, innovations: &mut HashMap<(usize, usize), usize>) {
if self.node_genes.len() < 2 {
return;
}
let connected: Vec<(usize, usize)> = self.connection_genes.iter()
.map(|c|
if c.in_node < c.out_node {
(c.in_node, c.out_node)
} else {
(c.out_node, c.in_node)
}
).collect();
let mut candidates = Vec::new();
for (i, a) in self.node_genes.iter().enumerate() {
for b in &self.node_genes[i + 1..] {
let pair = if a.id < b.id { (a.id, b.id) } else { (b.id, a.id) };
if connected.contains(&pair) {
continue;
}
if pair.0 < self.num_inputs && pair.1 < self.num_inputs {
continue;
}
if (pair.0 >= self.num_inputs && pair.0 < self.num_outputs) &&
(pair.1 >= self.num_inputs && pair.1 < self.num_outputs) {
continue;
}
candidates.push(pair);
}
}
if candidates.len() == 0 {
return;
}
let mut rng = rand::rng();
let chosen = candidates.choose(&mut rng).copied().unwrap();
let innov;
match innovations.get(&(chosen.0, chosen.1)) {
Some(x) => { innov = *x;
},
None => { innov = innovator.next();
innovations.insert((chosen.0, chosen.1), innov);
},
}
self.connection_genes.push(ConnectionGene {
in_node: chosen.0,
out_node: chosen.1,
weight: 1.0,
enabled: true,
innov,
});
self.connection_genes.sort_by_key(|c| c.innov);
}
}