use fugue::{addr, plate, sample, Bernoulli, Categorical, Model, ModelExt, Normal, Uniform};
use crate::genome::bit_string::BitString;
use crate::genome::bounds::MultiBounds;
use crate::genome::permutation::Permutation;
use crate::genome::real_vector::RealVector;
use crate::genome::trace_genome::TraceGenome;
use crate::genome::traits::{BinaryGenome, PermutationGenome, RealValuedGenome};
pub trait GenomePrior: Clone + Send + Sync + 'static {
type Genome: TraceGenome;
fn model(&self) -> Model<Self::Genome>;
}
#[derive(Clone, Debug)]
pub struct UniformBoxPrior {
bounds: MultiBounds,
}
impl UniformBoxPrior {
pub fn new(bounds: MultiBounds) -> Self {
Self { bounds }
}
pub fn bounds(&self) -> &MultiBounds {
&self.bounds
}
}
impl GenomePrior for UniformBoxPrior {
type Genome = RealVector;
fn model(&self) -> Model<RealVector> {
let bounds = self.bounds.clone();
plate!(i in 0..bounds.dimension().max(1) => {
let (lo, hi) = match bounds.get(i) {
Some(b) if b.max > b.min => (b.min, b.max),
Some(b) => (b.min - 1e-9, b.min + 1e-9),
None => (-1.0, 1.0),
};
sample(addr!("gene", i), Uniform::new(lo, hi).expect("valid uniform prior bounds"))
})
.map(|genes| RealVector::from_genes(genes).expect("plate produced genes"))
}
}
#[derive(Clone, Debug)]
pub struct GaussianPrior {
mean: f64,
std: f64,
dim: usize,
}
impl GaussianPrior {
pub fn new(mean: f64, std: f64, dim: usize) -> Self {
assert!(std > 0.0, "Gaussian prior std must be > 0");
Self { mean, std, dim }
}
}
impl GenomePrior for GaussianPrior {
type Genome = RealVector;
fn model(&self) -> Model<RealVector> {
let (mean, std, dim) = (self.mean, self.std, self.dim.max(1));
plate!(i in 0..dim => {
sample(addr!("gene", i), Normal::new(mean, std).expect("valid Gaussian prior"))
})
.map(|genes| RealVector::from_genes(genes).expect("plate produced genes"))
}
}
#[derive(Clone, Debug)]
pub struct BitStringPrior {
p_one: f64,
len: usize,
}
impl BitStringPrior {
pub fn new(p_one: f64, len: usize) -> Self {
assert!(
p_one > 0.0 && p_one < 1.0,
"BitStringPrior p_one must be in (0, 1)"
);
Self { p_one, len }
}
pub fn uniform(len: usize) -> Self {
Self::new(0.5, len)
}
}
impl GenomePrior for BitStringPrior {
type Genome = BitString;
fn model(&self) -> Model<BitString> {
let (p, len) = (self.p_one, self.len.max(1));
plate!(i in 0..len => {
sample(addr!("bit", i), Bernoulli::new(p).expect("valid Bernoulli prior"))
})
.map(|bits| BitString::from_bits(bits).expect("plate produced bits"))
}
}
#[derive(Clone, Debug)]
pub struct PermutationPrior {
n: usize,
}
impl PermutationPrior {
pub fn new(n: usize) -> Self {
assert!(n > 0, "PermutationPrior needs n > 0");
Self { n }
}
}
impl GenomePrior for PermutationPrior {
type Genome = Permutation;
fn model(&self) -> Model<Permutation> {
let n = self.n;
fn rank_model(n: usize, i: usize, ranks: Vec<usize>) -> Model<Vec<usize>> {
if i == n {
return fugue::pure(ranks);
}
let k = n - i;
let probs = vec![1.0 / k as f64; k];
sample(
addr!("perm", i),
Categorical::new(probs).expect("valid categorical prior"),
)
.bind(move |r| {
let mut ranks = ranks;
ranks.push(r);
rank_model(n, i + 1, ranks)
})
}
rank_model(n, 0, Vec::with_capacity(n)).map(move |ranks| {
let mut available: Vec<usize> = (0..n).collect();
let perm: Vec<usize> = ranks.into_iter().map(|r| available.remove(r)).collect();
Permutation::from_permutation(perm).expect("Lehmer decode produced a permutation")
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use fugue::runtime::handler::run;
use fugue::runtime::interpreters::{PriorHandler, ScoreGivenTrace};
use fugue::Trace;
use rand::rngs::StdRng;
use rand::SeedableRng;
#[test]
fn test_gaussian_prior_draws_match_moments() {
let prior = GaussianPrior::new(0.0, 2.0, 1);
let mut rng = StdRng::seed_from_u64(99);
let xs: Vec<f64> = (0..5000)
.map(|_| {
let (g, _) = run(
PriorHandler {
rng: &mut rng,
trace: Trace::default(),
},
prior.model(),
);
g.genes()[0]
})
.collect();
let mean = xs.iter().sum::<f64>() / xs.len() as f64;
let var = xs.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / xs.len() as f64;
assert!(mean.abs() < 0.2, "prior mean {}", mean);
assert!((var.sqrt() - 2.0).abs() < 0.2, "prior std {}", var.sqrt());
}
#[test]
fn test_prior_model_log_prior_matches_analytic() {
let prior = GaussianPrior::new(1.0, 2.0, 3);
let g = RealVector::new(vec![0.5, 1.5, -2.0]);
let (_, scored) = run(
ScoreGivenTrace {
base: g.to_trace(),
trace: Trace::default(),
},
prior.model(),
);
let normal = Normal::new(1.0, 2.0).unwrap();
let analytic: f64 = g
.genes()
.iter()
.map(|x| fugue::Distribution::log_prob(&normal, x))
.sum();
assert!((scored.log_prior - analytic).abs() < 1e-12);
}
#[test]
fn test_uniform_prior_out_of_box_scores_neg_inf() {
let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 2));
let g = RealVector::new(vec![0.5, 5.0]); let (_, scored) = run(
ScoreGivenTrace {
base: g.to_trace(),
trace: Trace::default(),
},
prior.model(),
);
assert_eq!(scored.log_prior, f64::NEG_INFINITY);
}
#[test]
fn test_permutation_prior_generates_valid_permutations() {
let prior = PermutationPrior::new(6);
let mut rng = StdRng::seed_from_u64(7);
for _ in 0..50 {
let (p, trace) = run(
PriorHandler {
rng: &mut rng,
trace: Trace::default(),
},
prior.model(),
);
assert!(p.is_valid_permutation());
let expected = -(720.0f64).ln(); assert!((trace.log_prior - expected).abs() < 1e-9);
let canonical = p.to_trace();
for (addr, choice) in &canonical.choices {
assert_eq!(trace.choices[addr].value, choice.value);
}
}
}
#[test]
fn test_bitstring_prior_scores_canonical_trace() {
let prior = BitStringPrior::uniform(4);
let g = BitString::from_bits(vec![true, false, true, true]).unwrap();
let (decoded, scored) = run(
ScoreGivenTrace {
base: g.to_trace(),
trace: Trace::default(),
},
prior.model(),
);
assert_eq!(decoded.bits(), g.bits());
assert!((scored.log_prior - 4.0 * (0.5f64).ln()).abs() < 1e-12);
}
}