use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use fugue::{factor, observe, pure, Address, Distribution, Model, SampleType};
use crate::fitness::traits::Fitness;
pub trait GenomeLikelihood<G>: Clone + Send + Sync + 'static {
fn model(&self, genome: &G, beta: f64) -> Model<()>;
}
pub fn tempered_observe<T: SampleType>(
addr: Address,
dist: impl Distribution<T> + 'static,
value: T,
beta: f64,
) -> Model<()> {
if beta == 1.0 {
observe(addr, dist, value)
} else {
factor(beta * dist.log_prob(&value))
}
}
#[derive(Clone, Debug)]
pub struct FactorFitness<F> {
pub fitness: F,
}
impl<F> FactorFitness<F> {
pub fn new(fitness: F) -> Self {
Self { fitness }
}
}
impl<G, F> GenomeLikelihood<G> for FactorFitness<F>
where
F: Fitness<Genome = G, Value = f64> + Clone + Send + Sync + 'static,
G: 'static,
{
fn model(&self, genome: &G, beta: f64) -> Model<()> {
factor(beta * self.fitness.evaluate(genome))
}
}
#[derive(Clone)]
pub struct MemoizedFitness<F> {
inner: F,
cache: Arc<Mutex<HashMap<Vec<u8>, f64>>>,
}
impl<F> MemoizedFitness<F> {
pub fn new(fitness: F) -> Self {
Self {
inner: fitness,
cache: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn cache_len(&self) -> usize {
self.cache.lock().map(|c| c.len()).unwrap_or(0)
}
}
impl<F> Fitness for MemoizedFitness<F>
where
F: Fitness<Value = f64>,
F::Genome: serde::Serialize,
{
type Genome = F::Genome;
type Value = f64;
fn evaluate(&self, genome: &Self::Genome) -> f64 {
let key = match bincode::serialize(genome) {
Ok(k) => k,
Err(_) => return self.inner.evaluate(genome), };
if let Ok(cache) = self.cache.lock() {
if let Some(&v) = cache.get(&key) {
return v;
}
}
let v = self.inner.evaluate(genome);
if let Ok(mut cache) = self.cache.lock() {
cache.insert(key, v);
}
v
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct NoLikelihood;
impl<G: 'static> GenomeLikelihood<G> for NoLikelihood {
fn model(&self, _genome: &G, _beta: f64) -> Model<()> {
pure(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::genome::real_vector::RealVector;
use crate::genome::traits::RealValuedGenome;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn test_memoized_fitness_evaluates_once_per_genome() {
static CALLS: AtomicUsize = AtomicUsize::new(0);
#[derive(Clone)]
struct Counting;
impl Fitness for Counting {
type Genome = RealVector;
type Value = f64;
fn evaluate(&self, g: &RealVector) -> f64 {
CALLS.fetch_add(1, Ordering::SeqCst);
-g.genes().iter().map(|x| x * x).sum::<f64>()
}
}
let memo = MemoizedFitness::new(Counting);
let a = RealVector::new(vec![1.0, 2.0]);
let b = RealVector::new(vec![3.0, 4.0]);
let fa = memo.evaluate(&a);
for _ in 0..10 {
assert_eq!(memo.evaluate(&a), fa);
}
memo.evaluate(&b);
memo.evaluate(&b);
assert_eq!(
CALLS.load(Ordering::SeqCst),
2,
"each genome evaluated once"
);
assert_eq!(memo.cache_len(), 2);
let clone = memo.clone();
clone.evaluate(&a);
assert_eq!(CALLS.load(Ordering::SeqCst), 2);
}
#[test]
fn test_tempered_observe_matches_observe_at_beta_one() {
use fugue::runtime::handler::run;
use fugue::runtime::interpreters::PriorHandler;
use fugue::{addr, Normal, Trace};
use rand::rngs::StdRng;
use rand::SeedableRng;
let mut rng = StdRng::seed_from_u64(1);
let dist = Normal::new(0.0, 1.0).unwrap();
let (_, t1) = run(
PriorHandler {
rng: &mut rng,
trace: Trace::default(),
},
tempered_observe(addr!("y"), dist, 0.7, 1.0),
);
let (_, t2) = run(
PriorHandler {
rng: &mut rng,
trace: Trace::default(),
},
tempered_observe(addr!("y"), dist, 0.7, 0.5),
);
let lp = fugue::Distribution::log_prob(&dist, &0.7);
assert!((t1.log_likelihood - lp).abs() < 1e-12);
assert_eq!(t1.log_factors, 0.0);
assert!((t2.log_factors - 0.5 * lp).abs() < 1e-12);
assert_eq!(t2.log_likelihood, 0.0);
}
}