Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 80 additions & 10 deletions genetic-rs-common/src/builtin/eliminator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -708,20 +708,90 @@ mod speciation {
{
#[cfg(not(feature = "rayon"))]
fn eliminate(&mut self, genomes: Vec<G>) -> Vec<G> {
let mut fitnesses = self.calculate_and_sort(genomes);
self.inner.observer.observe(&fitnesses);
let median_index = (fitnesses.len() as f32) * self.inner.threshold;
fitnesses.truncate(median_index as usize + 1);
fitnesses.into_iter().map(|(g, _)| g).collect()
let population = SpeciatedPopulation::from_genomes(
&genomes,
self.speciation_threshold,
&self.ctx,
);
let mut raw_fitnesses = vec![0.0_f32; genomes.len()];
let mut divided_fitnesses = vec![0.0_f32; genomes.len()];

for species in population.species() {
let len = species.len() as f32;
debug_assert!(len != 0.0);
for &index in species {
let fitness = self.inner.fitness_fn.fitness(&genomes[index]);
raw_fitnesses[index] = fitness;
divided_fitnesses[index] = if fitness < 0.0 {
fitness * len
} else {
fitness / len
};
}
}

// Sort by divided fitness (highest first) for elimination, but expose raw
// fitness values to the observer so it sees the unmodified scores.
let mut pairs: Vec<(G, f32, f32)> = genomes
.into_iter()
.enumerate()
.map(|(i, g)| (g, raw_fitnesses[i], divided_fitnesses[i]))
.collect();
pairs.sort_by(|(_, _, adiv), (_, _, bdiv)| bdiv.partial_cmp(adiv).unwrap());

let median_index = (pairs.len() as f32) * self.inner.threshold;

// Build observer slice with raw fitness values (sorted by divided fitness).
let mut observer_pairs: Vec<(G, f32)> =
pairs.into_iter().map(|(g, raw, _)| (g, raw)).collect();
self.inner.observer.observe(&observer_pairs);

observer_pairs.truncate(median_index as usize + 1);
observer_pairs.into_iter().map(|(g, _)| g).collect()
}

#[cfg(feature = "rayon")]
fn eliminate(&mut self, genomes: Vec<G>) -> Vec<G> {
let mut fitnesses = self.calculate_and_sort(genomes);
self.inner.observer.observe(&fitnesses);
let median_index = (fitnesses.len() as f32) * self.inner.threshold;
fitnesses.truncate(median_index as usize + 1);
fitnesses.into_par_iter().map(|(g, _)| g).collect()
let population = SpeciatedPopulation::from_genomes(
&genomes,
self.speciation_threshold,
&self.ctx,
);
let mut raw_fitnesses = vec![0.0_f32; genomes.len()];
let mut divided_fitnesses = vec![0.0_f32; genomes.len()];

for species in population.species() {
let len = species.len() as f32;
debug_assert!(len != 0.0);
for &index in species {
let fitness = self.inner.fitness_fn.fitness(&genomes[index]);
raw_fitnesses[index] = fitness;
divided_fitnesses[index] = if fitness < 0.0 {
fitness * len
} else {
fitness / len
};
}
}

// Sort by divided fitness (highest first) for elimination, but expose raw
// fitness values to the observer so it sees the unmodified scores.
let mut pairs: Vec<(G, f32, f32)> = genomes
.into_iter()
.enumerate()
.map(|(i, g)| (g, raw_fitnesses[i], divided_fitnesses[i]))
.collect();
pairs.sort_by(|(_, _, adiv), (_, _, bdiv)| bdiv.partial_cmp(adiv).unwrap());

let median_index = (pairs.len() as f32) * self.inner.threshold;

// Build observer slice with raw fitness values (sorted by divided fitness).
let mut observer_pairs: Vec<(G, f32)> =
pairs.into_iter().map(|(g, raw, _)| (g, raw)).collect();
self.inner.observer.observe(&observer_pairs);

observer_pairs.truncate(median_index as usize + 1);
observer_pairs.into_par_iter().map(|(g, _)| g).collect()
}
}
}
Expand Down
49 changes: 49 additions & 0 deletions genetic-rs/tests/speciation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -265,3 +265,52 @@ fn speciation_protects_rare_species() {
"the rare species genome must survive despite lower raw fitness"
);
}

/// The fitness observer on [`SpeciatedFitnessEliminator`] must receive the raw
/// (pre-division) fitness values, not the values after they have been divided by
/// the number of genomes in the species.
///
/// Setup:
/// - 4 genomes of class 0 with val = 1.0 → raw fitness = 1.0, divided = 0.25
/// - 1 genome of class 1 with val = 0.5 → raw fitness = 0.5, divided = 0.5
///
/// If the observer sees raw values, it must observe 1.0 and 0.5 among the scores.
/// If it sees divided values, it would observe 0.25 instead of 1.0 — the test
/// would fail in that case.
#[test]
fn observer_receives_pre_division_fitness() {
use std::sync::{Arc, Mutex};

let observed: Arc<Mutex<Vec<f32>>> = Arc::new(Mutex::new(Vec::new()));
let observed_clone = Arc::clone(&observed);

let observer = move |fitnesses: &[(Genome, f32)]| {
let mut v = observed_clone.lock().unwrap();
v.extend(fitnesses.iter().map(|(_, f)| *f));
};

let mut class0_genomes: Vec<Genome> = (0..4).map(|_| Genome { class: 0, val: 1.0 }).collect();
class0_genomes.push(Genome { class: 1, val: 0.5 });

let mut eliminator = SpeciatedFitnessEliminator::new(fitness, 0.5, 0.5, observer, ());
eliminator.eliminate(class0_genomes);

let scores = observed.lock().unwrap();
// Raw fitness values are 1.0 (×4) and 0.5 (×1).
// Divided values would be 0.25 and 0.5 — we must NOT see 0.25.
assert!(
scores.iter().any(|&f| (f - 1.0_f32).abs() < 1e-6),
"observer must see the raw fitness 1.0, but got: {:?}",
*scores,
);
assert!(
scores.iter().any(|&f| (f - 0.5_f32).abs() < 1e-6),
"observer must see the raw fitness 0.5, but got: {:?}",
*scores,
);
assert!(
!scores.iter().any(|&f| (f - 0.25_f32).abs() < 1e-6),
"observer must NOT see the divided fitness 0.25 (pre-division values expected), but got: {:?}",
*scores,
);
}
Loading