1use 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#[derive(Debug, Clone, PartialEq)]
19pub struct Genome {
20 genes: Vec<f64>,
21}
22
23impl Genome {
24 pub fn genes(&self) -> &[f64] {
26 &self.genes
27 }
28
29 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#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
50pub struct ConfigKey(Vec<u64>);
51
52#[derive(Debug, Clone, PartialEq)]
56pub struct SearchSpace {
57 factors: Vec<Factor>,
58}
59
60impl SearchSpace {
61 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 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 pub fn factors(&self) -> &[Factor] {
124 &self.factors
125 }
126
127 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 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 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 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
210fn 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
222fn 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
233fn level_index(gene: f64, count: u128) -> u128 {
235 ((gene * count as f64) as u128).min(count - 1)
237}
238
239fn 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#[derive(Debug, Clone, PartialEq)]
247pub enum SearchSpaceError {
248 NoFactors,
250 NoLevels {
252 factor_index: usize,
254 },
255 Factor(FactorError),
257 VariedTwice {
259 target: FactorTarget,
261 },
262 FixedAndVaried {
264 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 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, ¶ms(), &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()], ¶ms(), &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), ¶ms(), &actions(), &fixed),
495 Err(SearchSpaceError::FixedAndVaried { id: "rate".to_owned() })
496 );
497 assert_eq!(
498 SearchSpace::resolve(&[], ¶ms(), &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, ¶ms(), &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 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, ¶ms(), &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 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)], ¶ms(), &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}