Skip to main content

henad_core/explore/search/
genetic.rs

1//! Generational genetic algorithm for a noisy objective.
2//!
3//! Generation 0 is drawn at random. Each later generation re-evaluates the best members of the previous generation,
4//! keeps its elites and fills the rest with children. The children are bred as the generation starts, from the
5//! fitness the members had before its re-evaluations. The elites are chosen once the re-evaluations are told.
6//!
7//! A child comes from a tournament winner, crossed with a second winner at the crossover rate, then mutated gene by
8//! gene. Fitness is the objective over every replicate a member has, re-evaluations included.
9//!
10//! No two first evaluations share a config. A child is mutated again when an earlier candidate has the same config.
11//! When no draw finds a new config, the generation re-evaluates that candidate, and the candidate joins it in the
12//! child's place.
13
14use 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
24/// Maximum number of members in a generation.
25pub const MAX_POPULATION: usize = 1 << 16;
26
27/// Maximum number of members that one tournament can draw.
28pub const MAX_TOURNAMENT_SIZE: usize = 1 << 16;
29
30/// Settings of a genetic algorithm.
31#[derive(Debug, Clone, Copy, PartialEq)]
32pub struct GeneticSettings {
33    /// Number of members in each generation.
34    pub population: usize,
35    /// Number of best members carried unchanged into the next generation.
36    pub elite_count: usize,
37    /// Number of members drawn for each tournament, the best of whom becomes a parent.
38    pub tournament_size: usize,
39    /// Probability that a child has a second parent.
40    pub crossover_rate: f64,
41    /// Probability that each gene of a child changes.
42    pub mutation_rate: f64,
43    /// Largest step of a changed gene, as a fraction of its range.
44    pub mutation_scale: f64,
45    /// Share of the population re-evaluated each generation, rounded up, best members first.
46    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    /// Checks every setting against its range.
65    ///
66    /// # Errors
67    ///
68    /// Returns [`SearchSpecError::Setting`] for a population outside 1 to [`MAX_POPULATION`], as many elites as
69    /// members, a tournament size outside 1 to [`MAX_TOURNAMENT_SIZE`], a rate or fraction outside `[0, 1]`, or a
70    /// mutation scale that is not a positive number.
71    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    /// Returns the number of members re-evaluated each generation.
106    pub fn reevaluation_count(&self) -> usize {
107        let exact = self.reevaluate_fraction * self.population as f64;
108        // A product such as 0.1 * 30 comes out a hair above the whole number it represents, and would round up past it.
109        ((exact - 1e-9).ceil().max(0.0) as usize).min(self.population)
110    }
111}
112
113/// A generational genetic algorithm.
114#[derive(Debug, Clone)]
115pub struct GeneticAlgorithm {
116    space: SearchSpace,
117    settings: GeneticSettings,
118    rng: DesignRng,
119    tracker: CandidateTracker,
120    log: EvaluationLog,
121    /// First candidate of each config, by config key.
122    config_candidates: BTreeMap<ConfigKey, u64>,
123    /// Candidates of the current generation not yet requested.
124    queue: VecDeque<Proposal>,
125    /// Members of the last finished generation.
126    members: Vec<u64>,
127    /// Children of the current generation requested so far.
128    children: Vec<u64>,
129    /// Earlier candidates that join the current generation instead of a child with the same config.
130    revisited: BTreeSet<u64>,
131    generations: Vec<GenerationSummary>,
132}
133
134impl GeneticAlgorithm {
135    /// Returns a genetic algorithm of `max_evaluations` evaluations over `space`, drawing from `seed`.
136    ///
137    /// Note that generation 0 is smaller than the population when the space has fewer configs.
138    ///
139    /// # Errors
140    ///
141    /// Returns [`SearchSpecError`] when [`GeneticSettings::check`] rejects `settings`.
142    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    /// Ends the current generation, whose members are the elites of the last generation, the new children and the
177    /// candidates revisited instead of a child.
178    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    /// Queues the next generation, re-evaluations of the best members first and the children after them.
205    ///
206    /// A child whose draws all produce a known config re-evaluates the candidate the first draw matched in its place,
207    /// or is left out when the match is another child of the generation.
208    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    /// Returns the genome of a child before mutation, and its origin.
249    ///
250    /// The first parent wins a tournament. At the crossover rate, the winner of a second tournament crosses with it.
251    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
280/// Returns the best of `size` members drawn from `fitness` with replacement, the earliest drawn among equals.
281fn 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    /// Returns the height of a dome peaking at 0.6 in every gene.
368    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        // Five whole numbers by three options: 15 configs, fewer than the budget evaluates.
451        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        // 20 members, then 5 re-evaluations and 18 children a generation, in batches of at most 7.
490        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}