1use std::collections::HashSet;
7
8use fugue::{Address, ChoiceValue, Trace};
9use rand::Rng;
10use rand_distr::{Distribution, Normal};
11
12use crate::error::GenomeError;
13use crate::genome::trace_genome::TraceGenome;
14
15pub trait MutationSelector: Send + Sync {
17 fn select_sites<R: Rng>(&self, trace: &Trace, rng: &mut R) -> HashSet<Address>;
19}
20
21#[derive(Clone, Debug)]
25pub struct UniformMutationSelector {
26 pub mutation_probability: f64,
28}
29
30impl UniformMutationSelector {
31 pub fn new(probability: f64) -> Self {
33 Self {
34 mutation_probability: probability.clamp(0.0, 1.0),
35 }
36 }
37
38 pub fn one_over_n(n: usize) -> Self {
40 Self::new(1.0 / n as f64)
41 }
42}
43
44impl MutationSelector for UniformMutationSelector {
45 fn select_sites<R: Rng>(&self, trace: &Trace, rng: &mut R) -> HashSet<Address> {
46 trace
47 .choices
48 .keys()
49 .filter(|_| rng.gen::<f64>() < self.mutation_probability)
50 .cloned()
51 .collect()
52 }
53}
54
55#[derive(Clone, Debug, Default)]
59pub struct SingleSiteMutationSelector;
60
61impl SingleSiteMutationSelector {
62 pub fn new() -> Self {
64 Self
65 }
66}
67
68impl MutationSelector for SingleSiteMutationSelector {
69 fn select_sites<R: Rng>(&self, trace: &Trace, rng: &mut R) -> HashSet<Address> {
70 let addresses: Vec<_> = trace.choices.keys().collect();
71 if addresses.is_empty() {
72 return HashSet::new();
73 }
74
75 let idx = rng.gen_range(0..addresses.len());
76 let mut sites = HashSet::new();
77 sites.insert(addresses[idx].clone());
78 sites
79 }
80}
81
82#[derive(Clone, Debug)]
86pub struct MultiSiteMutationSelector {
87 pub num_sites: usize,
89}
90
91impl MultiSiteMutationSelector {
92 pub fn new(num_sites: usize) -> Self {
94 Self { num_sites }
95 }
96}
97
98impl MutationSelector for MultiSiteMutationSelector {
99 fn select_sites<R: Rng>(&self, trace: &Trace, rng: &mut R) -> HashSet<Address> {
100 let addresses: Vec<_> = trace.choices.keys().collect();
101 if addresses.is_empty() {
102 return HashSet::new();
103 }
104
105 let k = self.num_sites.min(addresses.len());
106 let mut selected = HashSet::new();
107 let mut indices: Vec<usize> = (0..addresses.len()).collect();
108
109 for i in 0..k {
111 let j = rng.gen_range(i..addresses.len());
112 indices.swap(i, j);
113 selected.insert(addresses[indices[i]].clone());
114 }
115
116 selected
117 }
118}
119
120pub trait CrossoverMask: Send + Sync {
122 fn from_parent1(&self, addr: &Address) -> bool;
124}
125
126#[derive(Clone, Debug)]
130pub struct UniformCrossoverMask {
131 pub bias: f64,
133 selected: HashSet<Address>,
135}
136
137impl UniformCrossoverMask {
138 pub fn new<R: Rng>(bias: f64, trace: &Trace, rng: &mut R) -> Self {
140 let selected = trace
141 .choices
142 .keys()
143 .filter(|_| rng.gen::<f64>() < bias)
144 .cloned()
145 .collect();
146
147 Self { bias, selected }
148 }
149
150 pub fn balanced<R: Rng>(trace: &Trace, rng: &mut R) -> Self {
152 Self::new(0.5, trace, rng)
153 }
154}
155
156impl CrossoverMask for UniformCrossoverMask {
157 fn from_parent1(&self, addr: &Address) -> bool {
158 self.selected.contains(addr)
159 }
160}
161
162#[derive(Clone, Debug)]
167pub struct SinglePointCrossoverMask {
168 parent1_addresses: HashSet<Address>,
170}
171
172impl SinglePointCrossoverMask {
173 pub fn new<R: Rng>(trace: &Trace, rng: &mut R) -> Self {
175 let addresses: Vec<_> = trace.choices.keys().cloned().collect();
176 if addresses.is_empty() {
177 return Self {
178 parent1_addresses: HashSet::new(),
179 };
180 }
181
182 let crossover_point = rng.gen_range(0..=addresses.len());
183 let parent1_addresses: HashSet<_> = addresses.into_iter().take(crossover_point).collect();
184
185 Self { parent1_addresses }
186 }
187}
188
189impl CrossoverMask for SinglePointCrossoverMask {
190 fn from_parent1(&self, addr: &Address) -> bool {
191 self.parent1_addresses.contains(addr)
192 }
193}
194
195#[derive(Clone, Debug)]
200pub struct TwoPointCrossoverMask {
201 parent1_addresses: HashSet<Address>,
203}
204
205impl TwoPointCrossoverMask {
206 pub fn new<R: Rng>(trace: &Trace, rng: &mut R) -> Self {
208 let addresses: Vec<_> = trace.choices.keys().cloned().collect();
209 if addresses.is_empty() {
210 return Self {
211 parent1_addresses: HashSet::new(),
212 };
213 }
214
215 let mut point1 = rng.gen_range(0..addresses.len());
216 let mut point2 = rng.gen_range(0..addresses.len());
217 if point1 > point2 {
218 std::mem::swap(&mut point1, &mut point2);
219 }
220
221 let parent1_addresses: HashSet<_> = addresses
223 .iter()
224 .enumerate()
225 .filter(|(i, _)| *i < point1 || *i >= point2)
226 .map(|(_, addr)| addr.clone())
227 .collect();
228
229 Self { parent1_addresses }
230 }
231}
232
233impl CrossoverMask for TwoPointCrossoverMask {
234 fn from_parent1(&self, addr: &Address) -> bool {
235 self.parent1_addresses.contains(addr)
236 }
237}
238
239pub fn mutate_trace<G, S, R>(
243 genome: &G,
244 selector: &S,
245 mutation_fn: impl Fn(&Address, &ChoiceValue, &mut R) -> ChoiceValue,
246 rng: &mut R,
247) -> Result<G, GenomeError>
248where
249 G: TraceGenome,
250 S: MutationSelector,
251 R: Rng,
252{
253 let trace = genome.to_trace();
254 let mutation_sites = selector.select_sites(&trace, rng);
255
256 let mut new_trace = Trace::default();
257
258 for (addr, choice) in &trace.choices {
259 let new_value = if mutation_sites.contains(addr) {
260 mutation_fn(addr, &choice.value, rng)
261 } else {
262 choice.value.clone()
263 };
264 new_trace.insert_choice(addr.clone(), new_value, choice.logp);
265 }
266
267 G::from_trace(&new_trace)
268}
269
270pub fn crossover_traces<G, M, R>(
274 parent1: &G,
275 parent2: &G,
276 mask: &M,
277 _rng: &mut R,
278) -> Result<(G, G), GenomeError>
279where
280 G: TraceGenome,
281 M: CrossoverMask,
282 R: Rng,
283{
284 let trace1 = parent1.to_trace();
285 let trace2 = parent2.to_trace();
286
287 let mut child1_trace = Trace::default();
288 let mut child2_trace = Trace::default();
289
290 let all_addresses: HashSet<Address> = trace1
292 .choices
293 .keys()
294 .chain(trace2.choices.keys())
295 .cloned()
296 .collect();
297
298 for addr in all_addresses {
299 let (val_for_child1, val_for_child2) = if mask.from_parent1(&addr) {
300 (
302 trace1
303 .choices
304 .get(&addr)
305 .map(|c| c.value.clone())
306 .unwrap_or(ChoiceValue::F64(0.0)),
307 trace2
308 .choices
309 .get(&addr)
310 .map(|c| c.value.clone())
311 .unwrap_or(ChoiceValue::F64(0.0)),
312 )
313 } else {
314 (
316 trace2
317 .choices
318 .get(&addr)
319 .map(|c| c.value.clone())
320 .unwrap_or(ChoiceValue::F64(0.0)),
321 trace1
322 .choices
323 .get(&addr)
324 .map(|c| c.value.clone())
325 .unwrap_or(ChoiceValue::F64(0.0)),
326 )
327 };
328
329 child1_trace.insert_choice(addr.clone(), val_for_child1, 0.0);
330 child2_trace.insert_choice(addr, val_for_child2, 0.0);
331 }
332
333 let child1 = G::from_trace(&child1_trace)?;
334 let child2 = G::from_trace(&child2_trace)?;
335
336 Ok((child1, child2))
337}
338
339pub fn gaussian_mutation<R: Rng>(
344 sigma: f64,
345) -> impl Fn(&Address, &ChoiceValue, &mut R) -> ChoiceValue {
346 let normal = Normal::new(0.0, sigma.max(0.0)).expect("sigma must be finite and non-negative");
347 move |_addr, value, rng| {
348 if let ChoiceValue::F64(v) = value {
349 ChoiceValue::F64(v + normal.sample(rng))
350 } else {
351 value.clone()
352 }
353 }
354}
355
356pub fn bit_flip_mutation<R: Rng>() -> impl Fn(&Address, &ChoiceValue, &mut R) -> ChoiceValue {
358 move |_addr, value, _rng| {
359 if let ChoiceValue::Bool(b) = value {
360 ChoiceValue::Bool(!b)
361 } else {
362 value.clone()
363 }
364 }
365}
366
367pub fn bounded_mutation<R: Rng>(
372 sigma: f64,
373 lower: f64,
374 upper: f64,
375) -> impl Fn(&Address, &ChoiceValue, &mut R) -> ChoiceValue {
376 let normal = Normal::new(0.0, sigma.max(0.0)).expect("sigma must be finite and non-negative");
377 move |_addr, value, rng| {
378 if let ChoiceValue::F64(v) = value {
379 let mutated = (v + normal.sample(rng)).clamp(lower, upper);
380 ChoiceValue::F64(mutated)
381 } else {
382 value.clone()
383 }
384 }
385}
386
387#[cfg(test)]
388mod tests {
389 use super::*;
390 use crate::genome::real_vector::RealVector;
391 use crate::genome::traits::{EvolutionaryGenome, RealValuedGenome};
392
393 #[test]
394 fn test_uniform_mutation_selector() {
395 let mut rng = rand::thread_rng();
396 let genome = RealVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
397 let trace = genome.to_trace();
398
399 let selector = UniformMutationSelector::new(0.9);
401 let sites = selector.select_sites(&trace, &mut rng);
402 assert!(sites.len() <= 5);
404
405 let selector_low = UniformMutationSelector::new(0.1);
407 let _sites_low = selector_low.select_sites(&trace, &mut rng);
408 }
409
410 #[test]
411 fn test_single_site_mutation_selector() {
412 let mut rng = rand::thread_rng();
413 let genome = RealVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
414 let trace = genome.to_trace();
415
416 let selector = SingleSiteMutationSelector::new();
417 let sites = selector.select_sites(&trace, &mut rng);
418
419 assert_eq!(sites.len(), 1);
420 }
421
422 #[test]
423 fn test_multi_site_mutation_selector() {
424 let mut rng = rand::thread_rng();
425 let genome = RealVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
426 let trace = genome.to_trace();
427
428 let selector = MultiSiteMutationSelector::new(3);
429 let sites = selector.select_sites(&trace, &mut rng);
430
431 assert_eq!(sites.len(), 3);
432 }
433
434 #[test]
435 fn test_mutate_trace() {
436 let mut rng = rand::thread_rng();
437 let genome = RealVector::new(vec![1.0, 2.0, 3.0]);
438
439 let selector = UniformMutationSelector::new(1.0); let mutation_fn = gaussian_mutation(0.1);
441
442 let mutated = mutate_trace(&genome, &selector, mutation_fn, &mut rng).unwrap();
443
444 assert_eq!(mutated.dimension(), genome.dimension());
446 }
448
449 #[test]
450 fn test_gaussian_mutation_achieves_sigma() {
451 use rand::rngs::StdRng;
454 use rand::SeedableRng;
455
456 let sigma = 1.0;
457 let mutate = gaussian_mutation::<StdRng>(sigma);
458 let mut rng = StdRng::seed_from_u64(2026);
459 let addr = fugue::addr!("gene", 0);
460 let base = ChoiceValue::F64(0.0);
461
462 let deltas: Vec<f64> = (0..50_000)
463 .map(|_| match mutate(&addr, &base, &mut rng) {
464 ChoiceValue::F64(v) => v,
465 _ => unreachable!(),
466 })
467 .collect();
468
469 let mean = deltas.iter().sum::<f64>() / deltas.len() as f64;
470 let std =
471 (deltas.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / deltas.len() as f64).sqrt();
472 assert!(mean.abs() < 0.03, "mean should be ~0, got {}", mean);
473 assert!(
474 (std - sigma).abs() < 0.03,
475 "empirical std {} should match sigma {} (old kernel gave ~0.816)",
476 std,
477 sigma
478 );
479 }
480
481 #[test]
482 fn test_bounded_mutation_respects_bounds() {
483 use rand::rngs::StdRng;
486 use rand::SeedableRng;
487
488 let mutate = bounded_mutation::<StdRng>(1.0, -0.5, 0.5);
489 let mut rng = StdRng::seed_from_u64(5);
490 let addr = fugue::addr!("gene", 0);
491 let base = ChoiceValue::F64(0.0);
492 for _ in 0..10_000 {
493 if let ChoiceValue::F64(v) = mutate(&addr, &base, &mut rng) {
494 assert!((-0.5..=0.5).contains(&v));
495 }
496 }
497 }
498
499 #[test]
500 fn test_crossover_traces() {
501 let mut rng = rand::thread_rng();
502 let parent1 = RealVector::new(vec![1.0, 2.0, 3.0]);
503 let parent2 = RealVector::new(vec![4.0, 5.0, 6.0]);
504
505 let trace1 = parent1.to_trace();
506 let mask = UniformCrossoverMask::balanced(&trace1, &mut rng);
507
508 let (child1, child2) = crossover_traces(&parent1, &parent2, &mask, &mut rng).unwrap();
509
510 assert_eq!(child1.dimension(), 3);
511 assert_eq!(child2.dimension(), 3);
512
513 for i in 0..3 {
515 let c1_val = child1.genes()[i];
516 let c2_val = child2.genes()[i];
517 let p1_val = parent1.genes()[i];
518 let p2_val = parent2.genes()[i];
519
520 assert!(
521 (c1_val - p1_val).abs() < 1e-10 || (c1_val - p2_val).abs() < 1e-10,
522 "Child1 value {} not from either parent ({} or {})",
523 c1_val,
524 p1_val,
525 p2_val
526 );
527 assert!(
528 (c2_val - p1_val).abs() < 1e-10 || (c2_val - p2_val).abs() < 1e-10,
529 "Child2 value {} not from either parent ({} or {})",
530 c2_val,
531 p1_val,
532 p2_val
533 );
534 }
535 }
536
537 #[test]
538 fn test_single_point_crossover_mask() {
539 let mut rng = rand::thread_rng();
540 let genome = RealVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
541 let trace = genome.to_trace();
542
543 for _ in 0..10 {
544 let mask = SinglePointCrossoverMask::new(&trace, &mut rng);
545
546 let addresses: Vec<_> = trace.choices.keys().collect();
548 let mut found_split = false;
549
550 for i in 1..addresses.len() {
551 let prev_from_p1 = mask.from_parent1(addresses[i - 1]);
552 let curr_from_p1 = mask.from_parent1(addresses[i]);
553
554 if prev_from_p1 && !curr_from_p1 {
555 assert!(!found_split, "Multiple splits found");
556 found_split = true;
557 }
558 }
559 }
560 }
561}