use std::collections::HashMap;
use fugue::inference::mcmc_utils::DiminishingAdaptation;
use fugue::{
adaptive_mcmc_chain_with_overrides, adaptive_single_site_mh, Address, SiteProposal, Trace,
};
use rand::Rng;
use super::likelihood::GenomeLikelihood;
use super::model::EvolutionModel;
use super::prior::GenomePrior;
pub struct EvolutionChain<P, L>
where
P: GenomePrior,
L: GenomeLikelihood<P::Genome>,
{
model: EvolutionModel<P, L>,
adaptation: DiminishingAdaptation,
overrides: HashMap<Address, SiteProposal>,
}
impl<P, L> EvolutionChain<P, L>
where
P: GenomePrior,
L: GenomeLikelihood<P::Genome>,
{
pub fn new(model: EvolutionModel<P, L>) -> Self {
Self {
model,
adaptation: DiminishingAdaptation::new(0.44, 0.7),
overrides: HashMap::new(),
}
}
pub fn target_rate(mut self, rate: f64) -> Self {
self.adaptation = DiminishingAdaptation::new(rate, 0.7);
self
}
pub fn override_site(mut self, addr: Address, proposal: SiteProposal) -> Self {
self.overrides.insert(addr, proposal);
self
}
pub fn model(&self) -> &EvolutionModel<P, L> {
&self.model
}
pub fn init<R: Rng>(&self, rng: &mut R) -> Trace {
use fugue::runtime::handler::run;
use fugue::runtime::interpreters::PriorHandler;
let (_g, trace) = run(
PriorHandler {
rng,
trace: Trace::default(),
},
(self.model.target_model())(),
);
trace
}
pub fn init_from(&self, genome: &P::Genome) -> Option<Trace> {
let (_g, trace) = self.model.score(genome);
if trace.total_log_weight().is_finite() {
Some(trace)
} else {
None
}
}
pub fn step<R: Rng>(&mut self, rng: &mut R, current: &Trace) -> (P::Genome, Trace) {
adaptive_single_site_mh(
rng,
self.model.target_model(),
current,
&mut self.adaptation,
)
}
pub fn run_chain<R: Rng>(&self, rng: &mut R, n: usize, warmup: usize) -> Vec<(P::Genome, Trace)>
where
P::Genome: Clone,
{
adaptive_mcmc_chain_with_overrides(
rng,
self.model.target_model(),
n,
warmup,
&self.overrides,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::fitness::traits::Fitness;
use crate::genome::bounds::{Bounds, MultiBounds};
use crate::genome::real_vector::RealVector;
use crate::genome::traits::{BinaryGenome, PermutationGenome, RealValuedGenome};
use crate::inference::model::tests::PtrFitness;
use crate::inference::prior::{BitStringPrior, PermutationPrior, UniformBoxPrior};
use rand::rngs::StdRng;
use rand::SeedableRng;
fn linear_x0(g: &RealVector) -> f64 {
g.genes()[0]
}
#[test]
fn test_mh_respects_bounds() {
let prior = UniformBoxPrior::new(MultiBounds::new(vec![Bounds::new(-2.0, 2.0)]));
let model = EvolutionModel::new(prior, PtrFitness(linear_x0)).with_beta(1.0);
let mut chain = EvolutionChain::new(model);
let mut rng = StdRng::seed_from_u64(20260710);
let mut current = chain.init(&mut rng);
let mut samples = Vec::new();
for i in 0..40_000 {
let (g, t) = chain.step(&mut rng, ¤t);
current = t;
let x = g.genes()[0];
assert!((-2.0..=2.0).contains(&x), "MH sample escaped bounds: {}", x);
if i >= 5_000 {
samples.push(x);
}
}
let mean = samples.iter().sum::<f64>() / samples.len() as f64;
let analytic = {
let e2 = 2.0_f64.exp();
let em2 = (-2.0_f64).exp();
(e2 + 3.0 * em2) / (e2 - em2)
};
assert!(
(mean - analytic).abs() < 0.1,
"posterior mean {} deviates from truncated-exponential analytic {}",
mean,
analytic
);
}
#[test]
fn test_bitstring_chain_moves() {
#[derive(Clone, Copy)]
struct OnesCount;
impl Fitness for OnesCount {
type Genome = crate::genome::bit_string::BitString;
type Value = f64;
fn evaluate(&self, g: &Self::Genome) -> f64 {
g.bits().iter().filter(|&&b| b).count() as f64
}
}
let model = EvolutionModel::new(BitStringPrior::uniform(8), OnesCount).with_beta(1.0);
let mut chain = EvolutionChain::new(model);
let mut rng = StdRng::seed_from_u64(11);
let init = chain.init(&mut rng);
let init_bits: Vec<Option<bool>> = (0..8)
.map(|i| init.get_bool(&fugue::addr!("bit", i)))
.collect();
let mut current = init.clone();
let mut moved = false;
for _ in 0..200 {
let (_g, t) = chain.step(&mut rng, ¤t);
current = t;
let bits: Vec<Option<bool>> = (0..8)
.map(|i| current.get_bool(&fugue::addr!("bit", i)))
.collect();
if bits != init_bits {
moved = true;
break;
}
}
assert!(moved, "BitString chain never moved (dead-chain regression)");
}
#[test]
fn test_permutation_chain_moves() {
#[derive(Clone, Copy)]
struct SortedNess;
impl Fitness for SortedNess {
type Genome = crate::genome::permutation::Permutation;
type Value = f64;
fn evaluate(&self, g: &Self::Genome) -> f64 {
g.permutation().windows(2).filter(|w| w[0] < w[1]).count() as f64
}
}
let model = EvolutionModel::new(PermutationPrior::new(5), SortedNess).with_beta(1.0);
let mut chain = EvolutionChain::new(model);
let mut rng = StdRng::seed_from_u64(17);
let init = chain.init(&mut rng);
let read_perm = |t: &Trace| -> Vec<usize> {
(0..5)
.map(|i| t.get_usize(&fugue::addr!("perm", i)).unwrap())
.collect()
};
let init_perm = read_perm(&init);
let mut current = init;
let mut moved = false;
for _ in 0..500 {
let (g, t) = chain.step(&mut rng, ¤t);
current = t;
assert!(
g.is_valid_permutation(),
"chain left the permutation support: {:?}",
g.permutation()
);
if read_perm(¤t) != init_perm {
moved = true;
}
}
assert!(
moved,
"Permutation chain never moved (dead-chain regression)"
);
}
}