use std::cmp::Ordering;
use ndarray::Array1;
use crate::{
genetic::{D12, PopulationSOO},
operators::survival::SurvivalOperator,
random::RandomGenerator,
};
pub struct FitnessSurvival;
impl SurvivalOperator for FitnessSurvival {
type FDim = ndarray::Ix1;
fn operate<ConstrDim>(
&mut self,
population: PopulationSOO<ConstrDim>,
num_survive: usize,
_rng: &mut impl RandomGenerator,
) -> PopulationSOO<ConstrDim>
where
ConstrDim: D12,
{
let pop_size = population.fitness.len();
let mut indices: Vec<usize> = (0..pop_size).collect();
if let Some(violations) = &population.constraint_violation_totals {
indices.sort_by(|&i, &j| {
let ord1 = violations[i]
.partial_cmp(&violations[j])
.unwrap_or(Ordering::Equal);
if ord1 != Ordering::Equal {
ord1
} else {
population.fitness[i]
.partial_cmp(&population.fitness[j])
.unwrap_or(Ordering::Equal)
}
});
} else {
indices.sort_by(|&i, &j| {
population.fitness[i]
.partial_cmp(&population.fitness[j])
.unwrap_or(Ordering::Equal)
});
}
let survive_count = num_survive.min(pop_size);
let selected_indices = &indices[..survive_count];
let mut selected_population = population.selected(selected_indices);
selected_population.set_rank(Array1::from_iter(0..survive_count));
selected_population
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::random::TestDummyRng;
use ndarray::{Array2, array};
struct FakeRandomGenerator {
dummy: TestDummyRng,
}
impl FakeRandomGenerator {
fn new() -> Self {
Self {
dummy: TestDummyRng,
}
}
}
impl RandomGenerator for FakeRandomGenerator {
type R = TestDummyRng;
fn rng(&mut self) -> &mut TestDummyRng {
&mut self.dummy
}
fn shuffle_vec_usize(&mut self, _vector: &mut Vec<usize>) {
}
}
#[test]
fn selects_lowest_two_fitness() {
let genes = Array2::zeros((4, 1));
let fitness = array![0.8, 0.2, 0.5, 0.1];
let pop = PopulationSOO::new_unconstrained(genes, fitness);
let mut rng = FakeRandomGenerator::new();
let mut selector = FitnessSurvival;
let survived = selector.operate(pop, 2, &mut rng);
assert_eq!(survived.fitness.len(), 2);
let expected_fitness = array![0.1, 0.2];
assert_eq!(survived.fitness, expected_fitness);
let expected_ranks = array![0, 1];
assert_eq!(survived.rank.unwrap(), expected_ranks);
}
#[test]
fn selects_lowest_two_constraints_violation() {
let genes = Array2::zeros((4, 1));
let fitness = array![0.8, 0.2, 0.5, 0.1];
let constraints = array![0.0, 0.0, 10.0, 100.0];
let pop = PopulationSOO::new(genes, fitness, constraints);
let mut rng = FakeRandomGenerator::new();
let mut selector = FitnessSurvival;
let survived = selector.operate(pop, 2, &mut rng);
assert_eq!(survived.fitness.len(), 2);
let expected_fitness = array![0.2, 0.8];
assert_eq!(survived.fitness, expected_fitness);
let best = survived.best();
assert_eq!(best.fitness, array![0.2]);
assert_eq!(best.constraint_violation_totals.unwrap(), array![0.0]);
}
#[test]
fn selects_all_when_num_survive_exceeds_population() {
let genes = Array2::zeros((4, 1));
let fitness = array![0.3, 0.1, 0.2];
let pop = PopulationSOO::new_unconstrained(genes, fitness);
let mut rng = FakeRandomGenerator::new();
let mut selector = FitnessSurvival;
let survived = selector.operate(pop, 5, &mut rng);
assert_eq!(survived.fitness.len(), 3);
let expected_fitness = array![0.1, 0.2, 0.3];
assert_eq!(survived.fitness, expected_fitness);
let expected_ranks = array![0, 1, 2];
assert_eq!(survived.rank.unwrap(), expected_ranks);
}
#[test]
fn single_individual_survives_with_rank_zero() {
let genes = Array2::zeros((1, 1));
let fitness = array![0.42];
let pop = PopulationSOO::new_unconstrained(genes, fitness);
let mut rng = FakeRandomGenerator::new();
let mut selector = FitnessSurvival;
let survived = selector.operate(pop, 1, &mut rng);
assert_eq!(survived.fitness.len(), 1);
assert_eq!(survived.fitness[0], 0.42);
let expected_ranks = array![0];
assert_eq!(survived.rank.unwrap(), expected_ranks);
}
}