Skip to main content

fugue_evo/inference/
trace_operators.rs

1//! Trace-based genetic operators
2//!
3//! These operators work directly on Fugue traces, enabling probabilistic
4//! interpretations of mutation and crossover.
5
6use 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
15/// Trait for selecting which addresses to mutate
16pub trait MutationSelector: Send + Sync {
17    /// Select addresses that should be mutated
18    fn select_sites<R: Rng>(&self, trace: &Trace, rng: &mut R) -> HashSet<Address>;
19}
20
21/// Uniform random mutation selector
22///
23/// Each address has an independent probability of being selected for mutation.
24#[derive(Clone, Debug)]
25pub struct UniformMutationSelector {
26    /// Probability of mutating each address
27    pub mutation_probability: f64,
28}
29
30impl UniformMutationSelector {
31    /// Create a new uniform mutation selector
32    pub fn new(probability: f64) -> Self {
33        Self {
34            mutation_probability: probability.clamp(0.0, 1.0),
35        }
36    }
37
38    /// Default 1/n mutation probability
39    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/// Single-site mutation selector
56///
57/// Selects exactly one random address for mutation.
58#[derive(Clone, Debug, Default)]
59pub struct SingleSiteMutationSelector;
60
61impl SingleSiteMutationSelector {
62    /// Create a new single-site selector
63    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/// Multi-site mutation selector
83///
84/// Selects exactly k random addresses for mutation.
85#[derive(Clone, Debug)]
86pub struct MultiSiteMutationSelector {
87    /// Number of sites to mutate
88    pub num_sites: usize,
89}
90
91impl MultiSiteMutationSelector {
92    /// Create a new multi-site selector
93    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        // Fisher-Yates partial shuffle
110        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
120/// Trait for determining which parent contributes to each address during crossover
121pub trait CrossoverMask: Send + Sync {
122    /// Returns true if parent1's value should be used at the given address
123    fn from_parent1(&self, addr: &Address) -> bool;
124}
125
126/// Uniform crossover mask
127///
128/// Each address independently chosen from either parent.
129#[derive(Clone, Debug)]
130pub struct UniformCrossoverMask {
131    /// Probability of choosing parent1's value
132    pub bias: f64,
133    /// Set of addresses that should come from parent1
134    selected: HashSet<Address>,
135}
136
137impl UniformCrossoverMask {
138    /// Create a new uniform crossover mask
139    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    /// Create with 50/50 probability
151    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/// Single-point crossover mask
163///
164/// All addresses before the crossover point come from parent1,
165/// all after come from parent2.
166#[derive(Clone, Debug)]
167pub struct SinglePointCrossoverMask {
168    /// Addresses from parent1 (before crossover point)
169    parent1_addresses: HashSet<Address>,
170}
171
172impl SinglePointCrossoverMask {
173    /// Create a new single-point crossover mask
174    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/// Two-point crossover mask
196///
197/// Addresses between the two points come from parent2,
198/// outside comes from parent1.
199#[derive(Clone, Debug)]
200pub struct TwoPointCrossoverMask {
201    /// Addresses from parent1 (outside crossover segment)
202    parent1_addresses: HashSet<Address>,
203}
204
205impl TwoPointCrossoverMask {
206    /// Create a new two-point crossover mask
207    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        // Parent1 gets addresses outside [point1, point2)
222        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
239/// Trace-based mutation operator
240///
241/// Mutates a genome by selectively resampling addresses in its trace representation.
242pub 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
270/// Trace-based crossover operator
271///
272/// Creates offspring by merging parent traces according to a crossover mask.
273pub 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    // Collect all addresses from both parents
291    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            // Child1 gets parent1, child2 gets parent2
301            (
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            // Child1 gets parent2, child2 gets parent1
315            (
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
339/// Gaussian mutation function for f64 values.
340///
341/// Adds a true `Normal(0, sigma)` perturbation (via `rand_distr::Normal`), so
342/// the mutation's standard deviation is exactly `sigma`.
343pub 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
356/// Bit flip mutation function for boolean values
357pub 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
367/// Bounded mutation function that respects bounds.
368///
369/// Adds a true `Normal(0, sigma)` perturbation and clamps the result to
370/// `[lower, upper]`.
371pub 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        // High probability should select many sites
400        let selector = UniformMutationSelector::new(0.9);
401        let sites = selector.select_sites(&trace, &mut rng);
402        // Most likely selects multiple sites (probabilistic, so not deterministic)
403        assert!(sites.len() <= 5);
404
405        // Low probability should select few sites
406        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); // Mutate all
440        let mutation_fn = gaussian_mutation(0.1);
441
442        let mutated = mutate_trace(&genome, &selector, mutation_fn, &mut rng).unwrap();
443
444        // Should have same dimension
445        assert_eq!(mutated.dimension(), genome.dimension());
446        // Values should have changed (with high probability)
447    }
448
449    #[test]
450    fn test_gaussian_mutation_achieves_sigma() {
451        // regression: EV-54 — the perturbation std must equal sigma, not
452        // 0.816·sigma (the old uniform kernel scaled by sqrt(2)).
453        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        // regression: EV-54 — bounded_mutation also uses a true Gaussian kernel
484        // and never escapes [lower, upper].
485        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        // Children should have values from either parent
514        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            // Check that it's a valid partition (addresses are either all before or all after)
547            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}