radiate-gp 1.3.1

Extensions for radiate. Genetic Programming implementations for graphs (neural networks) and trees.
Documentation
use crate::collections::GraphChromosome;
use crate::node::{Node, NodeExt};
use radiate_core::genome::*;
use radiate_core::{AlterContext, Crossover, Expr, RateSet, RdRand, random_provider};
use std::cmp::Ordering;
use std::fmt::Debug;

const PARENT_RATE: &str = "crossover.graph.rate.parent";

pub struct GraphCrossover {
    rate: Expr,
    parent_node_rate: Expr,
}

impl GraphCrossover {
    pub fn new(rate: impl Into<Expr>, crossover_parent_node_rate: impl Into<Expr>) -> Self {
        GraphCrossover {
            rate: rate.into(),
            parent_node_rate: crossover_parent_node_rate.into(),
        }
    }
}

impl<T> Crossover<GraphChromosome<T>> for GraphCrossover
where
    T: Clone + PartialEq + Debug,
{
    fn rates(&self) -> RateSet {
        RateSet::new(self.rate.clone()).push(self.parent_node_rate.clone().alias(PARENT_RATE))
    }

    #[inline]
    fn cross(
        &self,
        parent_one: &mut Phenotype<GraphChromosome<T>>,
        parent_two: &mut Phenotype<GraphChromosome<T>>,
        ctx: &mut AlterContext,
    ) -> usize {
        let parent_rate = ctx.internal_rate(0);

        let is_speciated = !parent_one.species().is_empty() && !parent_two.species().is_empty();

        let geno_one = parent_one.genotype_mut();
        let geno_two = parent_two.genotype();

        let num_crosses = random_provider::with_rng(|rand| {
            let chromo_index = rand.range(0..std::cmp::min(geno_one.len(), geno_two.len()));
            let chromo_one = geno_one.get_mut(chromo_index).unwrap();
            let chromo_two = geno_two.get(chromo_index).unwrap();

            if is_speciated {
                crossover_speciated(chromo_one, chromo_two, parent_rate, rand)
            } else {
                crossover_uniform(chromo_one, chromo_two, parent_rate, rand)
            }
        });

        if num_crosses > 0 {
            parent_one.invalidate(ctx.generation());
            return num_crosses;
        }

        num_crosses
    }
}

fn crossover_uniform<T>(
    chromo_one: &mut GraphChromosome<T>,
    chromo_two: &GraphChromosome<T>,
    rate: f32,
    rand: &mut RdRand,
) -> usize
where
    T: Clone + PartialEq,
{
    let mut crosses = 0;
    let min_len = std::cmp::min(chromo_one.len(), chromo_two.len());

    for i in 0..min_len {
        let node_one = chromo_one.get_mut(i);
        let node_two = chromo_two.get(i);

        if let Some((node_one, node_two)) = node_one.zip(node_two) {
            if node_one.arity() != node_two.arity() {
                continue;
            }

            if !rand.bool(rate) {
                continue;
            }

            if node_one.value() != node_two.value() {
                node_one.set_value(node_two.value().clone());
                crosses += 1;
            }
        }
    }

    crosses
}

fn crossover_speciated<T>(
    chromo_one: &mut GraphChromosome<T>,
    chromo_two: &GraphChromosome<T>,
    rate: f32,
    rand: &mut RdRand,
) -> usize
where
    T: Clone + PartialEq + Debug,
{
    let mut crosses = 0;
    let (mut ia, mut ib) = (0, 0);

    while ia < chromo_one.len() && ib < chromo_two.len() {
        let gene_one = chromo_one.get(ia);
        let gene_two = chromo_two.get(ib);

        let Some((gene_one, gene_two)) = gene_one.zip(gene_two) else {
            break;
        };

        match gene_one.innovation().cmp(&gene_two.innovation()) {
            Ordering::Equal => {
                if rand.bool(rate) {
                    let node_one = chromo_one.get_mut(ia);

                    if let Some(node_one) = node_one
                        && node_one.arity() == gene_two.arity()
                        && node_one.value() != gene_two.value()
                    {
                        node_one.set_value(gene_two.value().clone());
                        crosses += 1;
                    }
                }

                ia += 1;
                ib += 1;
            }
            Ordering::Less => ia += 1,
            Ordering::Greater => ib += 1,
        }
    }

    crosses
}