Skip to main content

fugue_evo/inference/
model.rs

1//! The Boltzmann posterior over genomes, as a fugue program
2//!
3//! Fix a prior program `p(x)` (a [`GenomePrior`]) and an observation program
4//! `p(data | x)` (a [`GenomeLikelihood`] — `observe` statements, latent
5//! nuisance parameters, and/or black-box `factor`s). The target
6//!
7//! ```text
8//!     π_β(x) ∝ p(x) · p(data | x)^β
9//! ```
10//!
11//! *is literally a fugue model*: `prior.model().bind(|g|
12//! likelihood.model(&g, β).map(|_| g))`. Every density this layer needs —
13//! prior mass, tempered joint, MH acceptance — is obtained by running or
14//! replaying that program; there is no hand-written density code.
15//!
16//! For the classical black-box case, `EvolutionModel::new(prior, fitness)`
17//! wraps a scalar [`Fitness`] in [`FactorFitness`] (`factor(β·f(x))` — the
18//! Gibbs / generalized-Bayes posterior); `EvolutionModel::from_likelihood`
19//! accepts any observation program.
20//!
21//! Two builders exist because MH wants a fixed-β target while tempered SMC
22//! must receive the β = 1 program (fugue's `adaptive_smc` supplies β by
23//! tempering `log_likelihood + log_factors`; baking β in as well would
24//! double-count it).
25
26use fugue::runtime::handler::run;
27use fugue::runtime::interpreters::{PriorHandler, ScoreGivenTrace};
28use fugue::{factor, Model, ModelExt, Trace};
29use rand::Rng;
30
31use super::likelihood::{FactorFitness, GenomeLikelihood};
32use super::prior::GenomePrior;
33use crate::fitness::traits::Fitness;
34
35/// A probabilistic model of an evolutionary population: a genome prior
36/// program, an observation program (likelihood), and an inverse temperature
37/// `β`.
38#[derive(Clone)]
39pub struct EvolutionModel<P, L>
40where
41    P: GenomePrior,
42    L: GenomeLikelihood<P::Genome>,
43{
44    prior: P,
45    likelihood: L,
46    beta: f64,
47}
48
49impl<P, F> EvolutionModel<P, FactorFitness<F>>
50where
51    P: GenomePrior,
52    F: Fitness<Genome = P::Genome, Value = f64> + Clone + Send + Sync + 'static,
53{
54    /// Create a model with a black-box scalar fitness entering as
55    /// `factor(β·f(x))` (the classical Gibbs-posterior mode), at `β = 1`.
56    pub fn new(prior: P, fitness: F) -> Self {
57        Self {
58            prior,
59            likelihood: FactorFitness::new(fitness),
60            beta: 1.0,
61        }
62    }
63
64    /// Raw fitness `f(x)` (higher is better).
65    pub fn fitness_value(&self, genome: &P::Genome) -> f64 {
66        self.likelihood.fitness.evaluate(genome)
67    }
68
69    /// The tempered fitness log-factor `β · f(x)`.
70    pub fn log_weight(&self, genome: &P::Genome) -> f64 {
71        self.beta * self.fitness_value(genome)
72    }
73
74    /// Create a fugue trace whose choices equal the genome's encoding under
75    /// this model's prior and whose `total_log_weight()` equals `β · f(x)` —
76    /// the "fitness as likelihood" weighted-trace contract (EV-52): a genuine
77    /// `factor(β·f)` model run through
78    /// [`TraceScoringHandler`](super::effect_handlers::TraceScoringHandler),
79    /// so the mass lands in `log_factors`.
80    pub fn to_weighted_trace(&self, genome: &P::Genome) -> Trace {
81        let logw = self.log_weight(genome);
82        let base = self.prior.trace_of(genome);
83        let (_r, trace) = run(
84            super::effect_handlers::TraceScoringHandler::new(base),
85            factor(logw),
86        );
87        trace
88    }
89}
90
91impl<P, L> EvolutionModel<P, L>
92where
93    P: GenomePrior,
94    L: GenomeLikelihood<P::Genome>,
95{
96    /// Create a model from an arbitrary observation program (`observe`
97    /// statements, latent nuisance parameters, factors), at `β = 1`.
98    pub fn from_likelihood(prior: P, likelihood: L) -> Self {
99        Self {
100            prior,
101            likelihood,
102            beta: 1.0,
103        }
104    }
105
106    /// Set the inverse temperature `β` directly (`β ≥ 0`).
107    pub fn with_beta(mut self, beta: f64) -> Self {
108        self.beta = beta.max(0.0);
109        self
110    }
111
112    /// Set the temperature `T`; equivalent to `β = 1/T`.
113    pub fn with_temperature(mut self, temperature: f64) -> Self {
114        self.beta = if temperature > 0.0 {
115            1.0 / temperature
116        } else {
117            f64::INFINITY
118        };
119        self
120    }
121
122    /// Current inverse temperature `β`.
123    pub fn beta(&self) -> f64 {
124        self.beta
125    }
126
127    /// Current temperature `T = 1/β`.
128    pub fn temperature(&self) -> f64 {
129        1.0 / self.beta
130    }
131
132    /// The prior program.
133    pub fn prior(&self) -> &P {
134        &self.prior
135    }
136
137    /// The observation program.
138    pub fn likelihood(&self) -> &L {
139        &self.likelihood
140    }
141
142    /// The fixed-β target as a program:
143    /// `log π_β(x) = log p(x) + β·log p(data|x)`. This is the model MH runs
144    /// against.
145    pub fn target_model(&self) -> impl Fn() -> Model<P::Genome> + Clone + '_ {
146        let prior = self.prior.clone();
147        let likelihood = self.likelihood.clone();
148        let beta = self.beta;
149        move || {
150            let likelihood = likelihood.clone();
151            prior
152                .model()
153                .bind(move |g| likelihood.model(&g, beta).map(move |_| g))
154        }
155    }
156
157    /// The **untempered** (`β = 1`) joint program for tempered SMC: fugue's
158    /// `adaptive_smc` supplies β by tempering
159    /// `log_likelihood + log_factors`, applying it exactly once.
160    pub fn smc_model(&self) -> impl Fn() -> Model<P::Genome> + Clone + '_ {
161        let prior = self.prior.clone();
162        let likelihood = self.likelihood.clone();
163        move || {
164            let likelihood = likelihood.clone();
165            prior
166                .model()
167                .bind(move |g| likelihood.model(&g, 1.0).map(move |_| g))
168        }
169    }
170
171    /// Draw a genome from the prior `p(x)` by running the prior program.
172    pub fn sample_prior<R: Rng>(&self, rng: &mut R) -> P::Genome {
173        let (g, _) = run(
174            PriorHandler {
175                rng,
176                trace: Trace::default(),
177            },
178            self.prior.model(),
179        );
180        g
181    }
182
183    /// Score a genome under the fixed-β target by replaying its encoding
184    /// **under this model's prior** ([`GenomePrior::trace_of`]) through the
185    /// target program — so this works for every prior, including generative
186    /// grammars over trees.
187    ///
188    /// The returned trace satisfies `log π_β(g) = trace.total_log_weight()`
189    /// and `log p(g) = trace.log_prior`. A genome outside the prior's support
190    /// scores `log_prior = −∞`.
191    ///
192    /// Note: if the likelihood contains latent nuisance sites, they are not
193    /// part of `trace_of(g)` and would abort a strict replay — score via the
194    /// SMC/MH drivers in that case (which sample them), or marginalize them
195    /// externally.
196    pub fn score(&self, genome: &P::Genome) -> (P::Genome, Trace) {
197        run(
198            ScoreGivenTrace {
199                base: self.prior.trace_of(genome),
200                trace: Trace::default(),
201            },
202            (self.target_model())(),
203        )
204    }
205
206    /// Unnormalised log target `log π_β(x) = log p(x) + β·log p(data|x)`.
207    pub fn log_boltzmann_target(&self, genome: &P::Genome) -> f64 {
208        self.score(genome).1.total_log_weight()
209    }
210}
211
212#[cfg(test)]
213pub(crate) mod tests {
214    use super::*;
215    use crate::genome::bounds::MultiBounds;
216    use crate::genome::real_vector::RealVector;
217    use crate::genome::traits::RealValuedGenome;
218    use crate::inference::prior::{GaussianPrior, UniformBoxPrior};
219
220    /// A `Clone`-able fitness wrapping a function pointer.
221    #[derive(Clone, Copy)]
222    pub(crate) struct PtrFitness(pub(crate) fn(&RealVector) -> f64);
223
224    impl Fitness for PtrFitness {
225        type Genome = RealVector;
226        type Value = f64;
227        fn evaluate(&self, genome: &RealVector) -> f64 {
228            (self.0)(genome)
229        }
230    }
231
232    pub(crate) fn quad_origin(g: &RealVector) -> f64 {
233        -0.5 * g.genes().iter().map(|x| x * x).sum::<f64>()
234    }
235
236    #[test]
237    fn test_to_weighted_trace_carries_fitness_mass() {
238        // regression: EV-52 — total_log_weight() must equal β·f(x), not 0.
239        let prior = UniformBoxPrior::new(MultiBounds::symmetric(5.0, 2));
240        let model = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_beta(2.0);
241        let genome = RealVector::new(vec![1.0, 2.0]);
242        let f = model.fitness_value(&genome); // -0.5*(1+4) = -2.5
243        let trace = model.to_weighted_trace(&genome);
244        assert!((trace.total_log_weight() - 2.0 * f).abs() < 1e-9);
245        assert!((trace.log_factors - 2.0 * f).abs() < 1e-9);
246        assert!(trace.total_log_weight().abs() > 1e-6);
247    }
248
249    #[test]
250    fn test_score_composes_prior_and_factor() {
251        // log π_β = log p + β·f, with each part in its own accumulator.
252        let prior = GaussianPrior::new(0.0, 2.0, 2);
253        let model = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_beta(1.5);
254        let g = RealVector::new(vec![0.5, -1.0]);
255        let (decoded, scored) = model.score(&g);
256        assert_eq!(decoded.genes(), g.genes());
257        assert!((scored.log_factors - 1.5 * quad_origin(&g)).abs() < 1e-12);
258        assert!(scored.log_prior.is_finite());
259        assert!(
260            (scored.total_log_weight() - (scored.log_prior + scored.log_factors)).abs() < 1e-12
261        );
262    }
263
264    #[test]
265    fn test_out_of_bounds_scores_neg_inf() {
266        let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 1));
267        let model = EvolutionModel::new(prior, PtrFitness(quad_origin));
268        let g = RealVector::new(vec![3.0]);
269        assert_eq!(model.log_boltzmann_target(&g), f64::NEG_INFINITY);
270    }
271
272    #[test]
273    fn test_sample_prior_returns_decoded_genome() {
274        use rand::rngs::StdRng;
275        use rand::SeedableRng;
276        let prior = GaussianPrior::new(0.0, 1.0, 4);
277        let model = EvolutionModel::new(prior, PtrFitness(quad_origin));
278        let mut rng = StdRng::seed_from_u64(3);
279        let g = model.sample_prior(&mut rng);
280        assert_eq!(g.genes().len(), 4);
281    }
282
283    /// An observation-program likelihood: per-datum `observe` statements land
284    /// in `log_likelihood`, decomposed from the prior — structure the scalar
285    /// factor could never expose.
286    #[test]
287    fn test_observation_likelihood_scores_in_log_likelihood() {
288        use crate::inference::likelihood::{tempered_observe, GenomeLikelihood};
289        use fugue::{addr, Normal};
290
291        #[derive(Clone)]
292        struct GaussianData {
293            ys: Vec<f64>,
294            sigma: f64,
295        }
296        impl GenomeLikelihood<RealVector> for GaussianData {
297            fn model(&self, g: &RealVector, beta: f64) -> Model<()> {
298                let mu = g.genes()[0];
299                let sigma = self.sigma;
300                let mut m = fugue::pure(());
301                for (k, &y) in self.ys.iter().enumerate() {
302                    m = m.and_then(move |_| {
303                        tempered_observe(addr!("y", k), Normal::new(mu, sigma).unwrap(), y, beta)
304                    });
305                }
306                m
307            }
308        }
309
310        let prior = GaussianPrior::new(0.0, 2.0, 1);
311        let data = GaussianData {
312            ys: vec![0.4, 0.6, 0.5],
313            sigma: 0.5,
314        };
315        let model = EvolutionModel::from_likelihood(prior, data.clone());
316        let g = RealVector::new(vec![0.5]);
317        let (_, scored) = model.score(&g);
318        let normal = Normal::new(0.5, 0.5).unwrap();
319        let analytic: f64 = data
320            .ys
321            .iter()
322            .map(|y| fugue::Distribution::log_prob(&normal, y))
323            .sum();
324        assert!((scored.log_likelihood - analytic).abs() < 1e-12);
325        assert_eq!(scored.log_factors, 0.0);
326        assert!(scored.log_prior.is_finite());
327    }
328}