Skip to main content

henad_core/explore/search/
genome.rs

1//! Genomes of a search, and the space that decodes them into configs.
2//!
3//! A genome holds one gene per factor of its space, each a fraction from 0 to 1. A continuous gene maps linearly
4//! onto its range. A whole-number gene over `m` values takes the value at position `floor(gene * m)`, and so does a
5//! categorical gene over `m` listed levels.
6
7use std::collections::BTreeSet;
8use std::fmt;
9
10use crate::explore::design::{continuous_level, continuous_value, whole_number_level};
11use crate::explore::design_rng::DesignRng;
12use crate::explore::factor::{Factor, FactorDomain, FactorError, FactorLevel, FactorSpec, FactorTarget};
13use crate::explore::plan::Config;
14use crate::explore::spec::ActionSpec;
15use crate::params::{ParamDescriptor, ParamValue};
16
17/// A point in a search space, one gene per factor, each from 0 to 1.
18#[derive(Debug, Clone, PartialEq)]
19pub struct Genome {
20    genes: Vec<f64>,
21}
22
23impl Genome {
24    /// Genes in factor order, each from 0 to 1.
25    pub fn genes(&self) -> &[f64] {
26        &self.genes
27    }
28
29    /// Returns a child that takes each gene from `self` or `other` with even odds.
30    ///
31    /// # Panics
32    ///
33    /// Panics when the two genomes differ in length.
34    pub fn crossover(&self, other: &Self, rng: &mut DesignRng) -> Self {
35        assert_eq!(self.genes.len(), other.genes.len(), "both parents come from one space");
36        let genes = self
37            .genes
38            .iter()
39            .zip(&other.genes)
40            .map(|(&first, &second)| if rng.unit_f64() < 0.5 { first } else { second })
41            .collect();
42        Self { genes }
43    }
44}
45
46/// Key of the config a genome decodes to, one entry per factor.
47///
48/// Two genomes of one space have equal keys exactly when they decode to the same config.
49#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
50pub struct ConfigKey(Vec<u64>);
51
52/// Factors a search varies, one gene each.
53///
54/// A factor over listed levels is a categorical gene. Any other factor is an ordered gene.
55#[derive(Debug, Clone, PartialEq)]
56pub struct SearchSpace {
57    factors: Vec<Factor>,
58}
59
60impl SearchSpace {
61    /// Returns the space over `factors`.
62    ///
63    /// A level listed twice is kept once, at its first place.
64    ///
65    /// # Errors
66    ///
67    /// Returns [`SearchSpaceError`] for no factors, or a factor with an empty list of levels.
68    pub fn new(mut factors: Vec<Factor>) -> Result<Self, SearchSpaceError> {
69        if factors.is_empty() {
70            return Err(SearchSpaceError::NoFactors);
71        }
72        if let Some(factor_index) = factors
73            .iter()
74            .position(|factor| factor.levels().is_some_and(<[FactorLevel]>::is_empty))
75        {
76            return Err(SearchSpaceError::NoLevels { factor_index });
77        }
78        for factor in &mut factors {
79            if let FactorDomain::Levels(levels) = &mut factor.domain {
80                let mut seen = BTreeSet::new();
81                levels.retain(|level| seen.insert(level_identity(level)));
82            }
83        }
84        Ok(Self { factors })
85    }
86
87    /// Resolves `specs` against `params` and `actions` as a sampled design would, and returns their space.
88    ///
89    /// A range with no step spans every value from its `min` to its `max`. `fixed` lists the parameter values every
90    /// config shares, as `(id, value)` pairs.
91    ///
92    /// # Errors
93    ///
94    /// Returns [`SearchSpaceError`] for no factors, a factor [`FactorSpec::resolve_sampled`] rejects, a target varied
95    /// twice, or a parameter both fixed and varied.
96    pub fn resolve(
97        specs: &[FactorSpec],
98        params: &[ParamDescriptor],
99        actions: &[ActionSpec],
100        fixed: &[(String, String)],
101    ) -> Result<Self, SearchSpaceError> {
102        let mut factors = Vec::with_capacity(specs.len());
103        for (index, spec) in specs.iter().enumerate() {
104            if specs[..index].iter().any(|earlier| earlier.target == spec.target) {
105                return Err(SearchSpaceError::VariedTwice {
106                    target: spec.target.clone(),
107                });
108            }
109            if let FactorTarget::Param(id) = &spec.target
110                && fixed.iter().any(|(fixed_id, _)| fixed_id == id)
111            {
112                return Err(SearchSpaceError::FixedAndVaried { id: id.clone() });
113            }
114            factors.push(
115                spec.resolve_sampled(params, actions)
116                    .map_err(SearchSpaceError::Factor)?,
117            );
118        }
119        Self::new(factors)
120    }
121
122    /// Factors of the space, in gene order.
123    pub fn factors(&self) -> &[Factor] {
124        &self.factors
125    }
126
127    /// Returns a genome with every gene drawn uniformly.
128    pub fn random_genome(&self, rng: &mut DesignRng) -> Genome {
129        Genome {
130            genes: self.factors.iter().map(|_| rng.unit_f64()).collect(),
131        }
132    }
133
134    /// Returns a copy of `genome` in which each gene changes with probability `rate`.
135    ///
136    /// An ordered gene moves by `scale * (r1 + r2 - 1)` for two uniform draws, reflected back into `[0, 1]` at either
137    /// end. A categorical gene is drawn again.
138    ///
139    /// # Panics
140    ///
141    /// Panics when `genome` and the space have different gene counts.
142    pub fn mutate(&self, genome: &Genome, rng: &mut DesignRng, rate: f64, scale: f64) -> Genome {
143        assert_eq!(
144            genome.genes.len(),
145            self.factors.len(),
146            "the genome comes from this space"
147        );
148        let genes = self
149            .factors
150            .iter()
151            .zip(&genome.genes)
152            .map(|(factor, &gene)| {
153                if rng.unit_f64() >= rate {
154                    gene
155                } else if factor.levels().is_some() {
156                    rng.unit_f64()
157                } else {
158                    reflect(gene + scale * rng.triangular())
159                }
160            })
161            .collect();
162        Genome { genes }
163    }
164
165    /// Returns `base` with the level of every gene of `genome` written into it.
166    ///
167    /// # Panics
168    ///
169    /// Panics when `genome` and the space have different gene counts, or a factor's slot is past the end of `base`.
170    pub fn decode(&self, genome: &Genome, base: &Config) -> Config {
171        assert_eq!(
172            genome.genes.len(),
173            self.factors.len(),
174            "the genome comes from this space"
175        );
176        let mut config = base.clone();
177        for (factor, &gene) in self.factors.iter().zip(&genome.genes) {
178            factor.apply(&level(factor, gene), &mut config);
179        }
180        config
181    }
182
183    /// Returns the key of the config `genome` decodes to.
184    ///
185    /// # Panics
186    ///
187    /// Panics when `genome` and the space have different gene counts.
188    pub fn config_key(&self, genome: &Genome) -> ConfigKey {
189        assert_eq!(
190            genome.genes.len(),
191            self.factors.len(),
192            "the genome comes from this space"
193        );
194        let entries = self
195            .factors
196            .iter()
197            .zip(&genome.genes)
198            .map(|(factor, &gene)| match &factor.domain {
199                &FactorDomain::Continuous { min, max } => {
200                    u64::from(continuous_value(min, max, gene.clamp(0.0, 1.0)).to_bits())
201                }
202                &FactorDomain::WholeNumbers { min, max } => level_index(gene, u128::from(max - min) + 1) as u64,
203                FactorDomain::Levels(levels) => level_index(gene, levels.len() as u128) as u64,
204            })
205            .collect();
206        ConfigKey(entries)
207    }
208}
209
210/// Returns the level of `factor` at gene `gene`.
211fn level(factor: &Factor, gene: f64) -> FactorLevel {
212    match &factor.domain {
213        &FactorDomain::Continuous { min, max } => continuous_level(min, max, gene.clamp(0.0, 1.0)),
214        &FactorDomain::WholeNumbers { min, max } => {
215            let offset = level_index(gene, u128::from(max - min) + 1) as u64;
216            whole_number_level(factor.slot, min + offset)
217        }
218        FactorDomain::Levels(levels) => levels[level_index(gene, levels.len() as u128) as usize].clone(),
219    }
220}
221
222/// Returns a key that two levels share when they are equal, with zero and negative zero as one.
223fn level_identity(level: &FactorLevel) -> (u8, u64) {
224    match level {
225        FactorLevel::Param(ParamValue::F32(value)) => (0, u64::from(if *value == 0.0 { 0 } else { value.to_bits() })),
226        FactorLevel::Param(ParamValue::U32(value)) => (1, u64::from(*value)),
227        FactorLevel::Param(ParamValue::Bool(value)) => (2, u64::from(*value)),
228        FactorLevel::Param(ParamValue::Choice(index)) => (3, *index as u64),
229        FactorLevel::Tick(tick) => (4, *tick),
230    }
231}
232
233/// Returns `floor(gene * count)`, kept from 0 to `count - 1`.
234fn level_index(gene: f64, count: u128) -> u128 {
235    // The cast saturates, taking a negative product to 0.
236    ((gene * count as f64) as u128).min(count - 1)
237}
238
239/// Returns `value` folded back into `[0, 1]`, as a reflection at each end would.
240fn reflect(value: f64) -> f64 {
241    let folded = value.rem_euclid(2.0);
242    if folded > 1.0 { 2.0 - folded } else { folded }
243}
244
245/// A search space that cannot be built.
246#[derive(Debug, Clone, PartialEq)]
247pub enum SearchSpaceError {
248    /// A space with no factors.
249    NoFactors,
250    /// Factor `factor_index`, a list with no levels.
251    NoLevels {
252        /// Index of the factor in the space, counting from 0.
253        factor_index: usize,
254    },
255    /// A factor that the search space rejects.
256    Factor(FactorError),
257    /// A target the space varies twice.
258    VariedTwice {
259        /// Parameter or action tick varied twice.
260        target: FactorTarget,
261    },
262    /// Parameter `id`, varied by the search and fixed as well.
263    FixedAndVaried {
264        /// Id of the parameter.
265        id: String,
266    },
267}
268
269impl fmt::Display for SearchSpaceError {
270    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
271        match self {
272            Self::NoFactors => write!(f, "search space has no factors"),
273            Self::NoLevels { factor_index } => write!(f, "factor {factor_index} of the search space has no levels"),
274            Self::Factor(_) => write!(f, "search space"),
275            Self::VariedTwice { target } => write!(f, "search space varies {target} twice"),
276            Self::FixedAndVaried { id } => write!(f, "parameter '{id}' is both fixed and varied by the search"),
277        }
278    }
279}
280
281impl std::error::Error for SearchSpaceError {
282    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
283        match self {
284            Self::Factor(error) => Some(error),
285            Self::NoFactors | Self::NoLevels { .. } | Self::VariedTwice { .. } | Self::FixedAndVaried { .. } => None,
286        }
287    }
288}
289
290#[cfg(test)]
291mod tests {
292    use std::collections::BTreeSet;
293
294    use super::{Genome, SearchSpace, SearchSpaceError, reflect};
295    use crate::explore::design_rng::DesignRng;
296    use crate::explore::factor::{Factor, FactorDomain, FactorLevel, FactorSlot, FactorSpec, FactorTarget, LevelSpec};
297    use crate::explore::plan::Config;
298    use crate::explore::spec::ActionSpec;
299    use crate::explore::value::check_value;
300    use crate::helpers::{choice_param, f32_param, u32_param};
301    use crate::params::{ParamDescriptor, ParamValue};
302
303    const SHAPES: &[&str] = &["ring", "star", "grid"];
304
305    fn params() -> Vec<ParamDescriptor> {
306        vec![
307            f32_param("rate", "Rate", 0.5, 0.0, 1.0, Some(0.01)),
308            u32_param("size", "Size", 10, 1, 100),
309            choice_param("shape", "Shape", SHAPES, 0),
310        ]
311    }
312
313    fn actions() -> Vec<ActionSpec> {
314        vec![ActionSpec::new("outbreak", 50)]
315    }
316
317    fn range(min: f64, max: f64) -> LevelSpec {
318        LevelSpec::Range { min, max, step: None }
319    }
320
321    /// Returns a space over every kind of gene: continuous, whole numbers, ticks and a choice.
322    fn space() -> SearchSpace {
323        let specs = [
324            FactorSpec::param("rate", range(0.05, 0.9)),
325            FactorSpec::param("size", range(3.0, 17.0)),
326            FactorSpec::action("outbreak", range(0.0, 400.0)),
327            FactorSpec::param("shape", LevelSpec::All),
328        ];
329        SearchSpace::resolve(&specs, &params(), &actions(), &[]).expect("every factor resolves")
330    }
331
332    fn base() -> Config {
333        Config {
334            block: 0,
335            params: vec![ParamValue::F32(0.5), ParamValue::U32(10), ParamValue::Choice(0)],
336            action_ticks: vec![50],
337        }
338    }
339
340    fn genome(genes: &[f64]) -> Genome {
341        Genome { genes: genes.to_vec() }
342    }
343
344    #[test]
345    fn a_decoded_point_respects_every_bound() {
346        let space = space();
347        let params = params();
348        let mut rng = DesignRng::new(4);
349        let mut genomes = vec![genome(&[0.0; 4]), genome(&[1.0; 4]), genome(&[0.999_999_999_999; 4])];
350        for _ in 0..1000 {
351            let random = space.random_genome(&mut rng);
352            genomes.push(space.mutate(&random, &mut rng, 1.0, 0.5));
353            genomes.push(random);
354        }
355        for genome in &genomes {
356            let config = space.decode(genome, &base());
357            for (descriptor, value) in params.iter().zip(&config.params) {
358                assert!(
359                    check_value(&descriptor.kind, value).is_ok(),
360                    "{value:?} for {}",
361                    descriptor.id
362                );
363            }
364            let ParamValue::F32(rate) = config.params[0] else {
365                panic!("rate is an f32");
366            };
367            assert!((0.05..=0.9).contains(&rate), "rate {rate}");
368            let ParamValue::U32(size) = config.params[1] else {
369                panic!("size is a u32");
370            };
371            assert!((3..=17).contains(&size), "size {size}");
372            assert!(config.action_ticks[0] <= 400, "tick {}", config.action_ticks[0]);
373        }
374        assert_eq!(
375            space.decode(&genome(&[0.0; 4]), &base()).params[0],
376            ParamValue::F32(0.05)
377        );
378        assert_eq!(
379            space.decode(&genome(&[1.0; 4]), &base()).params[0],
380            ParamValue::F32(0.9)
381        );
382    }
383
384    #[test]
385    fn integer_and_categorical_genes_decode_to_valid_levels() {
386        let space = space();
387        let decode = |gene: f64| space.decode(&genome(&[0.5, gene, gene, gene]), &base());
388        let ends = [decode(0.0), decode(1.0), decode(0.999_999_999_999)];
389        assert_eq!(ends[0].params[1], ParamValue::U32(3));
390        assert_eq!(ends[0].action_ticks[0], 0);
391        assert_eq!(ends[0].params[2], ParamValue::Choice(0));
392        for end in &ends[1..] {
393            assert_eq!(end.params[1], ParamValue::U32(17));
394            assert_eq!(end.action_ticks[0], 400);
395            assert_eq!(end.params[2], ParamValue::Choice(2));
396        }
397        let mut sizes = BTreeSet::new();
398        let mut ticks = BTreeSet::new();
399        let mut shapes = BTreeSet::new();
400        for step in 0..=4000 {
401            let config = decode(f64::from(step) / 4000.0);
402            let ParamValue::U32(size) = config.params[1] else {
403                panic!("size is a u32");
404            };
405            let ParamValue::Choice(shape) = config.params[2] else {
406                panic!("shape is a choice");
407            };
408            sizes.insert(size);
409            ticks.insert(config.action_ticks[0]);
410            shapes.insert(shape);
411        }
412        assert_eq!(sizes, (3..=17).collect(), "every whole number is reachable");
413        assert_eq!(ticks, (0..=400).collect(), "every tick is reachable");
414        assert_eq!(shapes, (0..3).collect(), "every option is reachable");
415    }
416
417    #[test]
418    fn genomes_share_a_config_key_exactly_when_they_share_a_config() {
419        let space = space();
420        let first = genome(&[0.3, 0.50, 0.5, 0.1]);
421        let same = [
422            genome(&[0.3 + 1e-12, 0.52, 0.5001, 0.2]),
423            genome(&[0.3, 0.51, 0.5, 0.3]),
424        ];
425        for other in &same {
426            assert_eq!(space.decode(other, &base()), space.decode(&first, &base()));
427            assert_eq!(space.config_key(other), space.config_key(&first), "{other:?}");
428        }
429        let different = [
430            genome(&[0.31, 0.5, 0.5, 0.1]),
431            genome(&[0.3, 0.6, 0.5, 0.1]),
432            genome(&[0.3, 0.5, 0.6, 0.1]),
433            genome(&[0.3, 0.5, 0.5, 0.9]),
434        ];
435        for other in &different {
436            assert_ne!(space.decode(other, &base()), space.decode(&first, &base()));
437            assert_ne!(space.config_key(other), space.config_key(&first), "{other:?}");
438        }
439    }
440
441    #[test]
442    fn mutation_keeps_every_gene_in_the_unit_interval() {
443        let space = space();
444        let mut rng = DesignRng::new(8);
445        let mut genome = space.random_genome(&mut rng);
446        for _ in 0..10_000 {
447            genome = space.mutate(&genome, &mut rng, 0.7, 0.8);
448            assert!(
449                genome.genes().iter().all(|gene| (0.0..=1.0).contains(gene)),
450                "{genome:?}"
451            );
452        }
453        for (value, folded) in [(-0.25, 0.25), (1.25, 0.75), (0.5, 0.5), (2.5, 0.5), (-1.75, 0.25)] {
454            assert!(
455                (reflect(value) - folded).abs() < 1e-12,
456                "{value} folds to {}",
457                reflect(value)
458            );
459        }
460    }
461
462    #[test]
463    fn a_zero_rate_changes_no_gene() {
464        let space = space();
465        let mut rng = DesignRng::new(2);
466        let genome = space.random_genome(&mut rng);
467        assert_eq!(space.mutate(&genome, &mut rng, 0.0, 0.5), genome);
468    }
469
470    #[test]
471    fn crossover_takes_each_gene_from_a_parent() {
472        let mut rng = DesignRng::new(6);
473        let first = genome(&[0.1; 64]);
474        let second = genome(&[0.9; 64]);
475        let child = first.crossover(&second, &mut rng);
476        assert!(child.genes().iter().all(|&gene| gene == 0.1 || gene == 0.9));
477        assert!(
478            child.genes().contains(&0.1) && child.genes().contains(&0.9),
479            "{child:?}"
480        );
481    }
482
483    #[test]
484    fn a_search_space_refuses_repeated_and_fixed_targets() {
485        let rate = FactorSpec::param("rate", range(0.1, 0.2));
486        assert_eq!(
487            SearchSpace::resolve(&[rate.clone(), rate.clone()], &params(), &actions(), &[]),
488            Err(SearchSpaceError::VariedTwice {
489                target: FactorTarget::Param("rate".to_owned()),
490            })
491        );
492        let fixed = [("rate".to_owned(), "0.3".to_owned())];
493        assert_eq!(
494            SearchSpace::resolve(std::slice::from_ref(&rate), &params(), &actions(), &fixed),
495            Err(SearchSpaceError::FixedAndVaried { id: "rate".to_owned() })
496        );
497        assert_eq!(
498            SearchSpace::resolve(&[], &params(), &actions(), &[]),
499            Err(SearchSpaceError::NoFactors)
500        );
501    }
502
503    #[test]
504    fn listed_levels_are_categorical() {
505        let specs = [FactorSpec::param(
506            "size",
507            LevelSpec::Values(vec!["4".to_owned(), "8".to_owned()]),
508        )];
509        let space = SearchSpace::resolve(&specs, &params(), &actions(), &[]).expect("listed sizes resolve");
510        assert_eq!(
511            space.factors()[0].levels(),
512            Some(
513                &[
514                    FactorLevel::Param(ParamValue::U32(4)),
515                    FactorLevel::Param(ParamValue::U32(8))
516                ][..]
517            )
518        );
519        assert_eq!(space.decode(&genome(&[0.6]), &base()).params[1], ParamValue::U32(8));
520    }
521
522    #[test]
523    fn a_repeated_level_is_kept_once() {
524        // "2" refers to the third shape by its index.
525        let shapes = ["star", "ring", "star", "2", "grid"].map(str::to_owned).to_vec();
526        let specs = [FactorSpec::param("shape", LevelSpec::Values(shapes))];
527        let space = SearchSpace::resolve(&specs, &params(), &actions(), &[]).expect("listed shapes resolve");
528        let choices = [1, 0, 2].map(|index| FactorLevel::Param(ParamValue::Choice(index)));
529        assert_eq!(space.factors()[0].levels(), Some(&choices[..]), "first places kept");
530        for (first, second) in [(0.1, 0.3), (0.4, 0.6), (0.7, 0.9)] {
531            assert_eq!(
532                space.config_key(&genome(&[first])),
533                space.config_key(&genome(&[second]))
534            );
535        }
536        assert_ne!(space.config_key(&genome(&[0.1])), space.config_key(&genome(&[0.9])));
537
538        // Steps below the spacing of f32 values near 0.5 round several levels to one value.
539        let step = LevelSpec::Range {
540            min: 0.5,
541            max: 0.500_000_05,
542            step: Some(1e-8),
543        };
544        let space = SearchSpace::resolve(&[FactorSpec::param("rate", step)], &params(), &actions(), &[])
545            .expect("the rates resolve");
546        let rates = [0.5, f32::from_bits(0.5_f32.to_bits() + 1)].map(|rate| FactorLevel::Param(ParamValue::F32(rate)));
547        assert_eq!(space.factors()[0].levels(), Some(&rates[..]));
548
549        let signed_zeros = [0.0, -0.0, 0.25].map(|rate| FactorLevel::Param(ParamValue::F32(rate)));
550        let space = SearchSpace::new(vec![Factor {
551            slot: FactorSlot::Param(0),
552            domain: FactorDomain::Levels(signed_zeros.to_vec()),
553        }])
554        .expect("three levels");
555        assert_eq!(
556            space.factors()[0].levels().map(<[FactorLevel]>::len),
557            Some(2),
558            "zero is one level"
559        );
560    }
561}