1use std::collections::{BTreeMap, BTreeSet, VecDeque};
15
16use crate::explore::design_rng::DesignRng;
17use crate::explore::search::evaluation_log::EvaluationLog;
18use crate::explore::search::genome::{ConfigKey, Genome, SearchSpace};
19use crate::explore::search::{
20 Aggregate, Candidate, CandidateOrigin, CandidateTracker, ConfigDraw, Evaluation, GenerationSummary, Goal,
21 Objective, Proposal, RankingEntry, SearchReport, SearchSpecError, Searcher, check_setting, draw_config,
22};
23
24pub const MAX_POPULATION: usize = 1 << 16;
26
27pub const MAX_TOURNAMENT_SIZE: usize = 1 << 16;
29
30#[derive(Debug, Clone, Copy, PartialEq)]
32pub struct GeneticSettings {
33 pub population: usize,
35 pub elite_count: usize,
37 pub tournament_size: usize,
39 pub crossover_rate: f64,
41 pub mutation_rate: f64,
43 pub mutation_scale: f64,
45 pub reevaluate_fraction: f64,
47}
48
49impl Default for GeneticSettings {
50 fn default() -> Self {
51 Self {
52 population: 32,
53 elite_count: 2,
54 tournament_size: 3,
55 crossover_rate: 0.9,
56 mutation_rate: 0.2,
57 mutation_scale: 0.1,
58 reevaluate_fraction: 0.25,
59 }
60 }
61}
62
63impl GeneticSettings {
64 pub fn check(&self) -> Result<(), SearchSpecError> {
72 check_setting(
73 (1..=MAX_POPULATION).contains(&self.population),
74 "genetic.population",
75 self.population,
76 format!("from 1 to {MAX_POPULATION}"),
77 )?;
78 check_setting(
79 self.elite_count < self.population,
80 "genetic.elite_count",
81 self.elite_count,
82 format!("less than the population of {}", self.population),
83 )?;
84 check_setting(
85 (1..=MAX_TOURNAMENT_SIZE).contains(&self.tournament_size),
86 "genetic.tournament_size",
87 self.tournament_size,
88 format!("from 1 to {MAX_TOURNAMENT_SIZE}"),
89 )?;
90 for (key, rate) in [
91 ("genetic.crossover_rate", self.crossover_rate),
92 ("genetic.mutation_rate", self.mutation_rate),
93 ("genetic.reevaluate_fraction", self.reevaluate_fraction),
94 ] {
95 check_setting((0.0..=1.0).contains(&rate), key, rate, "from 0 to 1")?;
96 }
97 check_setting(
98 self.mutation_scale.is_finite() && self.mutation_scale > 0.0,
99 "genetic.mutation_scale",
100 self.mutation_scale,
101 "a positive number",
102 )
103 }
104
105 pub fn reevaluation_count(&self) -> usize {
107 let exact = self.reevaluate_fraction * self.population as f64;
108 ((exact - 1e-9).ceil().max(0.0) as usize).min(self.population)
110 }
111}
112
113#[derive(Debug, Clone)]
115pub struct GeneticAlgorithm {
116 space: SearchSpace,
117 settings: GeneticSettings,
118 rng: DesignRng,
119 tracker: CandidateTracker,
120 log: EvaluationLog,
121 config_candidates: BTreeMap<ConfigKey, u64>,
123 queue: VecDeque<Proposal>,
125 members: Vec<u64>,
127 children: Vec<u64>,
129 revisited: BTreeSet<u64>,
131 generations: Vec<GenerationSummary>,
132}
133
134impl GeneticAlgorithm {
135 pub fn new(
143 space: SearchSpace,
144 settings: GeneticSettings,
145 objective: Objective,
146 max_evaluations: u64,
147 seed: u64,
148 ) -> Result<Self, SearchSpecError> {
149 settings.check()?;
150 let mut rng = DesignRng::new(seed);
151 let mut proposed = BTreeSet::new();
152 let mut queue = VecDeque::with_capacity(settings.population);
153 for _ in 0..settings.population {
154 if let ConfigDraw::New(genome, key) =
155 draw_config(&space, &BTreeMap::new(), &proposed, || space.random_genome(&mut rng))
156 {
157 proposed.insert(key);
158 queue.push_back(Proposal::first_evaluation(genome, CandidateOrigin::Random));
159 }
160 }
161 Ok(Self {
162 space,
163 settings,
164 rng,
165 tracker: CandidateTracker::new(max_evaluations),
166 log: EvaluationLog::new(objective),
167 config_candidates: BTreeMap::new(),
168 queue,
169 members: Vec::new(),
170 children: Vec::new(),
171 revisited: BTreeSet::new(),
172 generations: Vec::new(),
173 })
174 }
175
176 fn finish_generation(&mut self) {
179 let mut members = self.log.rank(&self.members);
180 members.truncate(self.settings.elite_count);
181 members.append(&mut self.children);
182 let mut present: BTreeSet<u64> = members.iter().copied().collect();
183 for candidate_id in std::mem::take(&mut self.revisited) {
184 if present.insert(candidate_id) {
185 members.push(candidate_id);
186 }
187 }
188 let members = self.log.rank(&members);
189 let mut fitness: Vec<f64> = members
190 .iter()
191 .map(|&candidate_id| self.log.objective(candidate_id))
192 .collect();
193 let (best, worst) = (fitness[0], fitness[fitness.len() - 1]);
194 self.generations.push(GenerationSummary {
195 generation: self.generations.len() as u64,
196 best,
197 median: Aggregate::Median.combine(&mut fitness).unwrap_or(worst),
198 worst,
199 });
200 self.members = members;
201 self.breed();
202 }
203
204 fn breed(&mut self) {
209 let reevaluations = self.settings.reevaluation_count().min(self.members.len());
210 let mut reevaluated = BTreeSet::new();
211 for &candidate_id in &self.members[..reevaluations] {
212 let record = self.log.get(candidate_id).expect("every member was evaluated");
213 self.queue.push_back(Proposal::reevaluation(candidate_id, record));
214 reevaluated.insert(candidate_id);
215 }
216 let fitness: Vec<(u64, f64)> = self
217 .members
218 .iter()
219 .map(|&candidate_id| (candidate_id, self.log.objective(candidate_id)))
220 .collect();
221 let mut proposed = BTreeSet::new();
222 for _ in self.settings.elite_count..self.settings.population {
223 let (genome, origin) = self.child_genome(&fitness);
224 let (space, rng, settings) = (&self.space, &mut self.rng, &self.settings);
225 let draw = draw_config(space, &self.config_candidates, &proposed, || {
226 space.mutate(&genome, rng, settings.mutation_rate, settings.mutation_scale)
227 });
228 match draw {
229 ConfigDraw::New(genome, key) => {
230 proposed.insert(key);
231 self.queue.push_back(Proposal::first_evaluation(genome, origin));
232 }
233 ConfigDraw::Known { candidate_id } => {
234 if reevaluated.insert(candidate_id) {
235 let record = self
236 .log
237 .get(candidate_id)
238 .expect("a candidate with a config was evaluated");
239 self.queue.push_back(Proposal::reevaluation(candidate_id, record));
240 }
241 self.revisited.insert(candidate_id);
242 }
243 ConfigDraw::Proposed => {}
244 }
245 }
246 }
247
248 fn child_genome(&mut self, fitness: &[(u64, f64)]) -> (Genome, CandidateOrigin) {
252 let (goal, size) = (self.log.goal(), self.settings.tournament_size);
253 let first_parent_id = tournament(&mut self.rng, fitness, goal, size);
254 let first = self
255 .log
256 .get(first_parent_id)
257 .expect("every member was evaluated")
258 .genome();
259 if self.rng.unit_f64() < self.settings.crossover_rate {
260 let second_parent_id = tournament(&mut self.rng, fitness, goal, size);
261 let second = self
262 .log
263 .get(second_parent_id)
264 .expect("every member was evaluated")
265 .genome();
266 let origin = CandidateOrigin::Crossover {
267 first_parent_id,
268 second_parent_id,
269 };
270 (first.crossover(second, &mut self.rng), origin)
271 } else {
272 let origin = CandidateOrigin::Mutation {
273 parent_id: first_parent_id,
274 };
275 (first.clone(), origin)
276 }
277 }
278}
279
280fn tournament(rng: &mut DesignRng, fitness: &[(u64, f64)], goal: Goal, size: usize) -> u64 {
282 let mut winner = fitness[rng.index(fitness.len() as u64) as usize];
283 for _ in 1..size {
284 let entrant = fitness[rng.index(fitness.len() as u64) as usize];
285 if goal.is_better(entrant.1, winner.1) {
286 winner = entrant;
287 }
288 }
289 winner.0
290}
291
292impl Searcher for GeneticAlgorithm {
293 fn ask(&mut self, max: usize) -> Vec<Candidate> {
294 let count = self.tracker.capacity(max).min(self.queue.len());
295 let asked = self.tracker.issue(self.queue.drain(..count).collect());
296 self.children.extend(
297 asked
298 .iter()
299 .filter(|candidate| candidate.origin.reevaluated_id().is_none())
300 .map(|candidate| candidate.id),
301 );
302 asked
303 }
304
305 fn tell(&mut self, evaluations: &[Evaluation]) {
306 for evaluation in evaluations {
307 if let Some((candidate, batch)) = self.tracker.settle(evaluation.candidate_id) {
308 if candidate.origin.reevaluated_id().is_none() {
309 self.config_candidates
310 .entry(self.space.config_key(&candidate.genome))
311 .or_insert(candidate.id);
312 }
313 self.log.record(&candidate, batch, evaluation);
314 }
315 }
316 let started = !self.children.is_empty() || !self.revisited.is_empty();
317 if self.queue.is_empty() && !self.tracker.is_waiting() && started {
318 self.finish_generation();
319 }
320 }
321
322 fn is_done(&self) -> bool {
323 self.tracker.is_done()
324 }
325
326 fn report(&self) -> SearchReport {
327 SearchReport {
328 ranking: self.log.ranking(),
329 generations: self.generations.clone(),
330 archive: Vec::new(),
331 }
332 }
333
334 fn best(&self) -> Option<RankingEntry> {
335 self.log.best()
336 }
337
338 fn ranking_entry(&self, candidate_id: u64) -> Option<RankingEntry> {
339 self.log.ranking_entry(candidate_id)
340 }
341
342 fn generations(&self) -> &[GenerationSummary] {
343 &self.generations
344 }
345}
346
347#[cfg(test)]
348mod tests {
349 use std::collections::BTreeSet;
350
351 use super::{GeneticAlgorithm, GeneticSettings};
352 use crate::explore::search::tests::support::{drive, level_space, noise, unit_space};
353 use crate::explore::search::{Aggregate, Candidate, CandidateOrigin, Evaluation, Goal, Objective, Searcher as _};
354
355 fn settings() -> GeneticSettings {
356 GeneticSettings {
357 population: 20,
358 elite_count: 2,
359 tournament_size: 3,
360 crossover_rate: 0.9,
361 mutation_rate: 0.3,
362 mutation_scale: 0.15,
363 reevaluate_fraction: 0.25,
364 }
365 }
366
367 fn height(genes: &[f64]) -> f64 {
369 -genes.iter().map(|gene| (gene - 0.6) * (gene - 0.6)).sum::<f64>()
370 }
371
372 #[test]
373 fn the_genetic_algorithm_improves_on_a_noisy_objective() {
374 let objective = Objective {
375 column: "Infected:max".to_owned(),
376 goal: Goal::Maximize,
377 aggregate: Aggregate::Median,
378 };
379 let mut search = GeneticAlgorithm::new(unit_space(3), settings(), objective, 600, 11).expect("valid settings");
380 let noisy = |candidate: &Candidate, replicate| {
381 vec![Some(
382 height(candidate.genome.genes()) + 0.1 * noise(candidate.id, replicate),
383 )]
384 };
385 let asked = drive(&mut search, 16, 3, noisy);
386 let report = search.report();
387 let first = &report.generations[0];
388 let last = report.generations.last().expect("generations finished");
389 assert!(
390 report.generations.len() >= 20,
391 "{} generations",
392 report.generations.len()
393 );
394 assert!(
395 last.median > first.median,
396 "median {} against {} at first",
397 last.median,
398 first.median
399 );
400 assert!(
401 last.best > first.best,
402 "best {} against {} at first",
403 last.best,
404 first.best
405 );
406 let best = report.best().expect("the search scored candidates");
407 let genes = asked[best.candidate_id as usize].genome.genes();
408 assert!(height(genes) > -0.02, "best point {genes:?}");
409 assert!(
410 report.ranking.iter().any(|entry| entry.replicate_count > 3),
411 "some members gathered more replicates than one evaluation gives"
412 );
413 }
414
415 #[test]
416 fn elites_survive_into_the_next_generation() {
417 let objective = Objective {
418 column: "Infected:max".to_owned(),
419 goal: Goal::Minimize,
420 aggregate: Aggregate::Mean,
421 };
422 let settings = GeneticSettings {
423 reevaluate_fraction: 0.0,
424 ..settings()
425 };
426 let mut search = GeneticAlgorithm::new(unit_space(2), settings, objective, 38, 4).expect("valid settings");
427 let sum = |candidate: &Candidate, _| vec![Some(candidate.genome.genes().iter().sum())];
428 drive(&mut search, 7, 1, sum);
429 let report = search.report();
430 assert_eq!(report.generations.len(), 2, "20 random members, then 18 children");
431 assert!(report.generations[1].best <= report.generations[0].best);
432 let elites = search.log.rank(&(0..20).collect::<Vec<u64>>());
433 for elite in &elites[..2] {
434 assert!(search.members.contains(elite), "elite {elite} survives");
435 }
436 assert_eq!(search.members.len(), 20);
437 }
438
439 #[test]
440 fn no_two_first_evaluations_share_a_config() {
441 let objective = Objective {
442 column: "Infected:max".to_owned(),
443 goal: Goal::Maximize,
444 aggregate: Aggregate::Median,
445 };
446 let settings = GeneticSettings {
447 population: 12,
448 ..GeneticSettings::default()
449 };
450 let space = level_space();
452 let mut search = GeneticAlgorithm::new(space.clone(), settings, objective, 150, 3).expect("valid settings");
453 let noisy = |candidate: &Candidate, replicate| {
454 vec![Some(candidate.genome.genes()[0] + 0.1 * noise(candidate.id, replicate))]
455 };
456 let asked = drive(&mut search, 5, 2, noisy);
457 assert_eq!(
458 asked.len(),
459 150,
460 "the search spends its budget once every config is known"
461 );
462 let mut keys = BTreeSet::new();
463 for candidate in asked
464 .iter()
465 .filter(|candidate| candidate.origin.reevaluated_id().is_none())
466 {
467 assert!(
468 keys.insert(space.config_key(&candidate.genome)),
469 "candidate {} repeats the config of an earlier one",
470 candidate.id
471 );
472 }
473 assert_eq!(keys.len(), 15, "the search finds every config");
474 assert!(
475 asked[100..]
476 .iter()
477 .all(|candidate| matches!(candidate.origin, CandidateOrigin::Reevaluation { .. })),
478 "once every config is known, the search re-evaluates"
479 );
480 }
481
482 #[test]
483 fn no_batch_spans_two_generations() {
484 let objective = Objective {
485 column: "Infected:max".to_owned(),
486 goal: Goal::Maximize,
487 aggregate: Aggregate::Mean,
488 };
489 let mut search = GeneticAlgorithm::new(unit_space(2), settings(), objective, 200, 6).expect("valid settings");
491 let mut sizes = Vec::new();
492 while !search.is_done() {
493 let (queued, finished) = (search.queue.len(), search.generations.len());
494 let batch = search.ask(7);
495 let evaluations: Vec<Evaluation> = batch
496 .iter()
497 .map(|candidate| Evaluation {
498 candidate_id: candidate.id,
499 outputs: vec![vec![Some(height(candidate.genome.genes()))]],
500 })
501 .collect();
502 search.tell(&evaluations);
503 let drained = batch.len() == queued;
504 assert_eq!(
505 search.generations.len(),
506 finished + usize::from(drained),
507 "a generation ends with the batch that asks for its last candidate"
508 );
509 sizes.push(batch.len());
510 }
511 assert_eq!(sizes[..7], [7, 7, 6, 7, 7, 7, 2]);
512 }
513
514 #[test]
515 fn reevaluation_counts_round_up() {
516 let count = |reevaluate_fraction, population| {
517 GeneticSettings {
518 population,
519 reevaluate_fraction,
520 ..settings()
521 }
522 .reevaluation_count()
523 };
524 assert_eq!(
525 count(0.1, 30),
526 3,
527 "0.1 * 30 is 3, though the product lands a hair above"
528 );
529 assert_eq!(count(0.25, 10), 3);
530 assert_eq!(count(0.0, 10), 0);
531 assert_eq!(count(1.0, 10), 10);
532 }
533
534 #[test]
535 fn a_population_of_elites_alone_is_refused() {
536 let settings = GeneticSettings {
537 elite_count: 20,
538 ..settings()
539 };
540 let objective = Objective {
541 column: "Infected:max".to_owned(),
542 goal: Goal::Minimize,
543 aggregate: Aggregate::Mean,
544 };
545 assert!(GeneticAlgorithm::new(unit_space(1), settings, objective, 10, 1).is_err());
546 }
547}