1use std::collections::HashMap;
25
26use fugue::inference::mcmc_utils::DiminishingAdaptation;
27use fugue::runtime::handler::run;
28use fugue::runtime::interpreters::{PriorHandler, ScoreGivenTrace};
29use fugue::{
30 adaptive_mcmc_chain_with_overrides, adaptive_single_site_mh_cached, Address, SiteProposal,
31 Trace,
32};
33use rand::Rng;
34
35use super::likelihood::GenomeLikelihood;
36use super::model::EvolutionModel;
37use super::prior::GenomePrior;
38use crate::error::GenomeError;
39
40fn finite_state(trace: Trace) -> Result<Trace, GenomeError> {
41 if trace.total_log_weight().is_finite() {
42 Ok(trace)
43 } else {
44 Err(GenomeError::ConstraintViolation(
45 "genome is outside the prior's support (target log-density is not finite)".to_string(),
46 ))
47 }
48}
49
50pub struct EvolutionChain<P, L>
52where
53 P: GenomePrior,
54 L: GenomeLikelihood<P::Genome>,
55{
56 model: EvolutionModel<P, L>,
57 adaptation: DiminishingAdaptation,
58 overrides: HashMap<Address, SiteProposal>,
59}
60
61impl<P, L> EvolutionChain<P, L>
62where
63 P: GenomePrior,
64 L: GenomeLikelihood<P::Genome>,
65{
66 pub fn new(model: EvolutionModel<P, L>) -> Self {
68 Self {
69 model,
70 adaptation: DiminishingAdaptation::new(0.44, 0.7),
71 overrides: HashMap::new(),
72 }
73 }
74
75 pub fn target_rate(mut self, rate: f64) -> Self {
77 self.adaptation = DiminishingAdaptation::new(rate, 0.7);
78 self
79 }
80
81 pub fn override_site(mut self, addr: Address, proposal: SiteProposal) -> Self {
92 self.overrides.insert(addr, proposal);
93 self
94 }
95
96 pub fn overrides(&self) -> &HashMap<Address, SiteProposal> {
98 &self.overrides
99 }
100
101 pub fn model(&self) -> &EvolutionModel<P, L> {
103 &self.model
104 }
105
106 pub fn init<R: Rng>(&self, rng: &mut R) -> Trace {
109 let (_g, trace) = run(
110 PriorHandler {
111 rng,
112 trace: Trace::default(),
113 },
114 (self.model.target_model())(),
115 );
116 trace
117 }
118
119 pub fn init_from(&self, genome: &P::Genome) -> Option<Trace> {
130 self.try_init_from(genome).ok()
131 }
132
133 pub fn try_init_from(&self, genome: &P::Genome) -> Result<Trace, GenomeError> {
138 let (_g, trace) = self.model.score(genome)?;
139 finite_state(trace)
140 }
141
142 pub fn init_from_with_latents<R: Rng>(
148 &self,
149 rng: &mut R,
150 genome: &P::Genome,
151 ) -> Result<Trace, GenomeError> {
152 let (_g, trace) = self.model.score_with_latents(rng, genome)?;
153 finite_state(trace)
154 }
155
156 pub fn step<R: Rng>(&mut self, rng: &mut R, current: &Trace) -> (P::Genome, Trace) {
180 match self.step_scored(rng, current) {
181 Some((g, t, _log_weight)) => (g, t),
182 None => (self.decode(current), current.clone()),
183 }
184 }
185
186 pub fn step_scored<R: Rng>(
192 &mut self,
193 rng: &mut R,
194 current: &Trace,
195 ) -> Option<(P::Genome, Trace, f64)> {
196 adaptive_single_site_mh_cached(
197 rng,
198 self.model.target_model(),
199 current,
200 &mut self.adaptation,
201 &self.overrides,
202 true,
203 )
204 }
205
206 pub fn decode(&self, state: &Trace) -> P::Genome {
211 let (g, _) = run(
212 ScoreGivenTrace {
213 base: state.clone(),
214 trace: Trace::default(),
215 },
216 self.model.prior().model(),
217 );
218 g
219 }
220
221 pub fn run_chain<R: Rng>(&self, rng: &mut R, n: usize, warmup: usize) -> Vec<(P::Genome, Trace)>
225 where
226 P::Genome: Clone,
227 {
228 adaptive_mcmc_chain_with_overrides(
229 rng,
230 self.model.target_model(),
231 n,
232 warmup,
233 &self.overrides,
234 )
235 }
236}
237
238#[cfg(test)]
239mod tests {
240 use super::*;
241 use crate::fitness::traits::Fitness;
242 use crate::genome::bounds::{Bounds, MultiBounds};
243 use crate::genome::real_vector::RealVector;
244 use crate::genome::traits::{BinaryGenome, PermutationGenome, RealValuedGenome};
245 use crate::inference::model::tests::PtrFitness;
246 use crate::inference::prior::{BitStringPrior, PermutationPrior, UniformBoxPrior};
247 use rand::rngs::StdRng;
248 use rand::SeedableRng;
249
250 fn linear_x0(g: &RealVector) -> f64 {
251 g.genes()[0]
252 }
253
254 #[test]
260 fn test_mh_respects_bounds() {
261 let prior = UniformBoxPrior::new(MultiBounds::new(vec![Bounds::new(-2.0, 2.0)]));
262 let model = EvolutionModel::new(prior, PtrFitness(linear_x0)).with_beta(1.0);
263 let mut chain = EvolutionChain::new(model);
264
265 let mut rng = StdRng::seed_from_u64(20260710);
266 let mut current = chain.init(&mut rng);
267 let mut samples = Vec::new();
268 for i in 0..40_000 {
269 let (g, t) = chain.step(&mut rng, ¤t);
270 current = t;
271 let x = g.genes()[0];
272 assert!((-2.0..=2.0).contains(&x), "MH sample escaped bounds: {}", x);
273 if i >= 5_000 {
274 samples.push(x);
275 }
276 }
277 let mean = samples.iter().sum::<f64>() / samples.len() as f64;
278 let analytic = {
279 let e2 = 2.0_f64.exp();
280 let em2 = (-2.0_f64).exp();
281 (e2 + 3.0 * em2) / (e2 - em2)
282 };
283 assert!(
284 (mean - analytic).abs() < 0.1,
285 "posterior mean {} deviates from truncated-exponential analytic {}",
286 mean,
287 analytic
288 );
289 }
290
291 #[test]
302 fn test_mh_bounded_prior_containing_negatives_mixes_across_zero() {
303 let prior = UniformBoxPrior::new(MultiBounds::new(vec![Bounds::new(-0.5, 0.5)]));
304 let model = EvolutionModel::new(prior, PtrFitness(linear_x0)).with_beta(1.0);
305 let analytic_mean = 0.5 / (0.5f64).tanh() - 1.0;
306 let analytic_p_pos = (0.5f64.exp() - 1.0) / (0.5f64.exp() - (-0.5f64).exp());
307 for seed in [1u64, 2, 3, 20260710] {
308 let mut chain = EvolutionChain::new(model.clone());
309 let mut rng = StdRng::seed_from_u64(seed);
310 let mut current = chain.init(&mut rng);
311 let mut samples = Vec::new();
312 for i in 0..40_000 {
313 let (g, t) = chain.step(&mut rng, ¤t);
314 current = t;
315 let x = g.genes()[0];
316 assert!((-0.5..=0.5).contains(&x), "MH sample escaped bounds: {}", x);
317 if i >= 5_000 {
318 samples.push(x);
319 }
320 }
321 let n = samples.len() as f64;
322 let mean = samples.iter().sum::<f64>() / n;
323 let p_pos = samples.iter().filter(|&&x| x > 0.0).count() as f64 / n;
324 assert!(
325 (mean - analytic_mean).abs() < 0.04,
326 "seed {seed}: posterior mean {mean} deviates from analytic {analytic_mean}"
327 );
328 assert!(
329 (p_pos - analytic_p_pos).abs() < 0.08,
330 "seed {seed}: P(x > 0) = {p_pos} vs analytic {analytic_p_pos} — chain stuck on one sign"
331 );
332 }
333 }
334
335 #[test]
340 fn test_step_honours_override_site() {
341 let prior = || UniformBoxPrior::new(MultiBounds::new(vec![Bounds::new(-2.0, 2.0)]));
342 let start = RealVector::new(vec![0.0]);
343
344 let mut confined = EvolutionChain::new(EvolutionModel::new(prior(), PtrFitness(linear_x0)))
345 .override_site(
346 fugue::addr!("gene", 0),
347 SiteProposal::Reflect {
348 lower: -0.5,
349 upper: 0.5,
350 },
351 );
352 let mut rng = StdRng::seed_from_u64(3);
353 let mut current = confined.init_from(&start).expect("in support");
354 let mut accepted = 0;
355 for _ in 0..5_000 {
356 if let Some((g, t, _)) = confined.step_scored(&mut rng, ¤t) {
357 accepted += 1;
358 current = t;
359 let x = g.genes()[0];
360 assert!(
361 (-0.5..=0.5).contains(&x),
362 "override ignored: reflected chain left [-0.5, 0.5] at {x}"
363 );
364 }
365 }
366 assert!(
367 accepted > 500,
368 "confined chain barely moved ({accepted} acceptances)"
369 );
370
371 let mut free = EvolutionChain::new(EvolutionModel::new(prior(), PtrFitness(linear_x0)));
372 let mut rng = StdRng::seed_from_u64(3);
373 let mut current = free.init_from(&start).expect("in support");
374 let mut escaped = false;
375 for _ in 0..5_000 {
376 let (g, t) = free.step(&mut rng, ¤t);
377 current = t;
378 if g.genes()[0].abs() > 0.5 {
379 escaped = true;
380 break;
381 }
382 }
383 assert!(
384 escaped,
385 "without the override the chain must explore the whole box"
386 );
387 }
388
389 #[test]
394 fn test_step_costs_one_fitness_evaluation() {
395 use std::sync::atomic::{AtomicUsize, Ordering};
396 use std::sync::Arc;
397
398 #[derive(Clone)]
399 struct Counting(Arc<AtomicUsize>);
400 impl Fitness for Counting {
401 type Genome = RealVector;
402 type Value = f64;
403 fn evaluate(&self, g: &RealVector) -> f64 {
404 self.0.fetch_add(1, Ordering::SeqCst);
405 -0.5 * g.genes()[0].powi(2)
406 }
407 }
408
409 let counter = Arc::new(AtomicUsize::new(0));
410 let prior = UniformBoxPrior::new(MultiBounds::new(vec![Bounds::new(-3.0, 3.0)]));
411 let mut chain = EvolutionChain::new(EvolutionModel::new(prior, Counting(counter.clone())));
412 let mut rng = StdRng::seed_from_u64(5);
413 let mut current = chain
414 .init_from(&RealVector::new(vec![0.3]))
415 .expect("in support");
416 assert_eq!(counter.load(Ordering::SeqCst), 1, "init_from scores once");
417
418 let n = 2_000;
419 let mut rejections = 0;
420 for _ in 0..n {
421 let before = counter.load(Ordering::SeqCst);
422 let (g, t) = chain.step(&mut rng, ¤t);
423 assert_eq!(
424 counter.load(Ordering::SeqCst) - before,
425 1,
426 "a step must evaluate the fitness exactly once"
427 );
428 let x = t.get_f64(&fugue::addr!("gene", 0)).unwrap();
429 if x == current.get_f64(&fugue::addr!("gene", 0)).unwrap() {
430 rejections += 1;
431 }
432 assert_eq!(g.genes()[0], x);
433 current = t;
434 }
435 assert!(
436 rejections > 0,
437 "some proposals must be rejected for the test to bite"
438 );
439 assert_eq!(counter.load(Ordering::SeqCst), 1 + n);
440 }
441
442 #[test]
446 fn test_bitstring_chain_moves() {
447 #[derive(Clone, Copy)]
448 struct OnesCount;
449 impl Fitness for OnesCount {
450 type Genome = crate::genome::bit_string::BitString;
451 type Value = f64;
452 fn evaluate(&self, g: &Self::Genome) -> f64 {
453 g.bits().iter().filter(|&&b| b).count() as f64
454 }
455 }
456
457 let model = EvolutionModel::new(BitStringPrior::uniform(8), OnesCount).with_beta(1.0);
458 let mut chain = EvolutionChain::new(model);
459 let mut rng = StdRng::seed_from_u64(11);
460 let init = chain.init(&mut rng);
461 let init_bits: Vec<Option<bool>> = (0..8)
462 .map(|i| init.get_bool(&fugue::addr!("bit", i)))
463 .collect();
464
465 let mut current = init.clone();
466 let mut moved = false;
467 for _ in 0..200 {
468 let (_g, t) = chain.step(&mut rng, ¤t);
469 current = t;
470 let bits: Vec<Option<bool>> = (0..8)
471 .map(|i| current.get_bool(&fugue::addr!("bit", i)))
472 .collect();
473 if bits != init_bits {
474 moved = true;
475 break;
476 }
477 }
478 assert!(moved, "BitString chain never moved (dead-chain regression)");
479 }
480
481 #[test]
485 fn test_permutation_chain_moves() {
486 #[derive(Clone, Copy)]
487 struct SortedNess;
488 impl Fitness for SortedNess {
489 type Genome = crate::genome::permutation::Permutation;
490 type Value = f64;
491 fn evaluate(&self, g: &Self::Genome) -> f64 {
492 g.permutation().windows(2).filter(|w| w[0] < w[1]).count() as f64
494 }
495 }
496
497 let model = EvolutionModel::new(PermutationPrior::new(5), SortedNess).with_beta(1.0);
498 let mut chain = EvolutionChain::new(model);
499 let mut rng = StdRng::seed_from_u64(17);
500 let init = chain.init(&mut rng);
501 let read_perm = |t: &Trace| -> Vec<usize> {
502 (0..5)
503 .map(|i| t.get_usize(&fugue::addr!("perm", i)).unwrap())
504 .collect()
505 };
506 let init_perm = read_perm(&init);
507
508 let mut current = init;
509 let mut moved = false;
510 for _ in 0..500 {
511 let (g, t) = chain.step(&mut rng, ¤t);
512 current = t;
513 assert!(
514 g.is_valid_permutation(),
515 "chain left the permutation support: {:?}",
516 g.permutation()
517 );
518 if read_perm(¤t) != init_perm {
519 moved = true;
520 }
521 }
522 assert!(
523 moved,
524 "Permutation chain never moved (dead-chain regression)"
525 );
526 }
527}