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;
28use fugue::{factor, score_given_trace_reconciled, Model, ModelExt, ReconcileReport, Trace};
29use rand::rngs::StdRng;
30use rand::{Rng, SeedableRng};
31
32use super::likelihood::{FactorFitness, GenomeLikelihood};
33use super::prior::GenomePrior;
34use crate::error::GenomeError;
35use crate::fitness::traits::Fitness;
36
37/// Replay `model` over `base` and require a **complete, exact** assignment:
38/// every site the program visits is in `base` (with the right type) and every
39/// site of `base` is visited. Structural mismatches come back as errors
40/// instead of a panic (EV-N3).
41///
42/// Built on fugue's reconciling scorer rather than the strict one: the strict
43/// handler keeps running the program with `Default::default()` values after a
44/// missing site, and a program whose control flow depends on such a value can
45/// recurse without bound (a grammar reading a missing `#leaf` as "function
46/// node"). The reconciling scorer instead draws a missing site from its prior
47/// — from a fixed-seed generator here, so this function stays deterministic —
48/// and reports it; the draw only ever serves to terminate a replay that is
49/// rejected anyway.
50pub(crate) fn score_complete<A>(base: Trace, model: Model<A>) -> Result<(A, Trace), GenomeError> {
51    let mut rng = StdRng::seed_from_u64(0);
52    let (a, scored, report) = score_given_trace_reconciled(base, &mut rng, model)
53        .map_err(|e| GenomeError::InvalidStructure(e.to_string()))?;
54    check_exact(&report)?;
55    Ok((a, scored))
56}
57
58fn check_exact(report: &ReconcileReport) -> Result<(), GenomeError> {
59    if let Some(addr) = report.fresh_addresses.first() {
60        return Err(GenomeError::MissingAddress(addr.to_string()));
61    }
62    if !report.vanished_addresses.is_empty() {
63        let shown: Vec<String> = report
64            .vanished_addresses
65            .iter()
66            .take(3)
67            .map(|a| a.to_string())
68            .collect();
69        return Err(GenomeError::InvalidStructure(format!(
70            "encoding carries {} site(s) the program never visits: {}{}",
71            report.vanished_addresses.len(),
72            shown.join(", "),
73            if report.vanished_addresses.len() > 3 {
74                ", …"
75            } else {
76                ""
77            }
78        )));
79    }
80    Ok(())
81}
82
83/// A probabilistic model of an evolutionary population: a genome prior
84/// program, an observation program (likelihood), and an inverse temperature
85/// `β`.
86#[derive(Clone)]
87pub struct EvolutionModel<P, L>
88where
89    P: GenomePrior,
90    L: GenomeLikelihood<P::Genome>,
91{
92    prior: P,
93    likelihood: L,
94    beta: f64,
95}
96
97impl<P, F> EvolutionModel<P, FactorFitness<F>>
98where
99    P: GenomePrior,
100    F: Fitness<Genome = P::Genome, Value = f64> + Clone + Send + Sync + 'static,
101{
102    /// Create a model with a black-box scalar fitness entering as
103    /// `factor(β·f(x))` (the classical Gibbs-posterior mode), at `β = 1`.
104    pub fn new(prior: P, fitness: F) -> Self {
105        Self {
106            prior,
107            likelihood: FactorFitness::new(fitness),
108            beta: 1.0,
109        }
110    }
111
112    /// Raw fitness `f(x)` (higher is better).
113    pub fn fitness_value(&self, genome: &P::Genome) -> f64 {
114        self.likelihood.fitness.evaluate(genome)
115    }
116
117    /// The tempered fitness log-factor `β · f(x)`.
118    pub fn log_weight(&self, genome: &P::Genome) -> f64 {
119        self.beta * self.fitness_value(genome)
120    }
121
122    /// Create a fugue trace whose choices equal the genome's encoding under
123    /// this model's prior and whose `total_log_weight()` equals `β · f(x)` —
124    /// the "fitness as likelihood" weighted-trace contract (EV-52): a genuine
125    /// `factor(β·f)` model run through
126    /// [`TraceScoringHandler`](super::effect_handlers::TraceScoringHandler),
127    /// so the mass lands in `log_factors`. Fails with the prior's
128    /// [`GenomePrior::validate`] error (e.g. `DimensionMismatch`) for a genome
129    /// of the wrong shape, before the fitness is evaluated.
130    pub fn to_weighted_trace(&self, genome: &P::Genome) -> Result<Trace, GenomeError> {
131        self.prior.validate(genome)?;
132        let logw = self.log_weight(genome);
133        let base = self.prior.trace_of(genome);
134        let (_r, trace) = run(
135            super::effect_handlers::TraceScoringHandler::new(base),
136            factor(logw),
137        );
138        Ok(trace)
139    }
140}
141
142impl<P, L> EvolutionModel<P, L>
143where
144    P: GenomePrior,
145    L: GenomeLikelihood<P::Genome>,
146{
147    /// Create a model from an arbitrary observation program (`observe`
148    /// statements, latent nuisance parameters, factors), at `β = 1`.
149    pub fn from_likelihood(prior: P, likelihood: L) -> Self {
150        Self {
151            prior,
152            likelihood,
153            beta: 1.0,
154        }
155    }
156
157    /// Set the inverse temperature `β` directly (`β ≥ 0`; negative values
158    /// clamp to `0`, the prior).
159    ///
160    /// # Panics
161    ///
162    /// If `beta` is not finite. `β = ∞` would make the fitness factor
163    /// `∞ · f(x)` — `NaN` wherever `f = 0`, `±∞` elsewhere — a target on
164    /// which no sampler can move (EV-N5). Optimizer mode is
165    /// [`EvolutionSMC::anneal`](super::smc::EvolutionSMC::anneal) with a
166    /// large *finite* `beta_max`.
167    pub fn with_beta(mut self, beta: f64) -> Self {
168        assert!(
169            beta.is_finite(),
170            "EvolutionModel::with_beta: β must be finite (got {beta}); anneal toward a large finite β instead"
171        );
172        self.beta = beta.max(0.0);
173        self
174    }
175
176    /// Set the temperature `T > 0`; equivalent to `β = 1/T`.
177    ///
178    /// # Panics
179    ///
180    /// If `temperature` is not a finite, strictly positive number: `T = 0`
181    /// is `β = ∞`, see [`Self::with_beta`].
182    pub fn with_temperature(self, temperature: f64) -> Self {
183        assert!(
184            temperature.is_finite() && temperature > 0.0,
185            "EvolutionModel::with_temperature: T must be finite and > 0 (got {temperature})"
186        );
187        self.with_beta(1.0 / temperature)
188    }
189
190    /// Current inverse temperature `β`.
191    pub fn beta(&self) -> f64 {
192        self.beta
193    }
194
195    /// Current temperature `T = 1/β`.
196    pub fn temperature(&self) -> f64 {
197        1.0 / self.beta
198    }
199
200    /// The prior program.
201    pub fn prior(&self) -> &P {
202        &self.prior
203    }
204
205    /// The observation program.
206    pub fn likelihood(&self) -> &L {
207        &self.likelihood
208    }
209
210    /// The fixed-β target as a program:
211    /// `log π_β(x) = log p(x) + β·log p(data|x)`. This is the model MH runs
212    /// against.
213    pub fn target_model(&self) -> impl Fn() -> Model<P::Genome> + Clone + '_ {
214        let prior = self.prior.clone();
215        let likelihood = self.likelihood.clone();
216        let beta = self.beta;
217        move || {
218            let likelihood = likelihood.clone();
219            prior
220                .model()
221                .bind(move |g| likelihood.model(&g, beta).map(move |_| g))
222        }
223    }
224
225    /// The **untempered** (`β = 1`) joint program for tempered SMC: fugue's
226    /// `adaptive_smc` supplies β by tempering
227    /// `log_likelihood + log_factors`, applying it exactly once.
228    pub fn smc_model(&self) -> impl Fn() -> Model<P::Genome> + Clone + '_ {
229        let prior = self.prior.clone();
230        let likelihood = self.likelihood.clone();
231        move || {
232            let likelihood = likelihood.clone();
233            prior
234                .model()
235                .bind(move |g| likelihood.model(&g, 1.0).map(move |_| g))
236        }
237    }
238
239    /// Draw a genome from the prior `p(x)` by running the prior program.
240    pub fn sample_prior<R: Rng>(&self, rng: &mut R) -> P::Genome {
241        let (g, _) = run(
242            PriorHandler {
243                rng,
244                trace: Trace::default(),
245            },
246            self.prior.model(),
247        );
248        g
249    }
250
251    /// Score a genome under the fixed-β target by replaying its encoding
252    /// **under this model's prior** ([`GenomePrior::trace_of`]) through the
253    /// target program — so this works for every prior, including generative
254    /// grammars over trees.
255    ///
256    /// The returned trace satisfies `log π_β(g) = trace.total_log_weight()`
257    /// and `log p(g) = trace.log_prior`. A genome outside the prior's support
258    /// scores `log_prior = −∞` (that is a valid score, not an error).
259    ///
260    /// # Errors
261    ///
262    /// Never panics on a structural mismatch (EV-N3). Returns the prior's
263    /// [`GenomePrior::validate`] error for a genome of the wrong shape
264    /// ([`GenomeError::DimensionMismatch`] for the vector priors);
265    /// [`GenomeError::MissingAddress`] when the program visits a site the
266    /// encoding lacks — in particular a **latent nuisance site** of the
267    /// likelihood (`NoiseSpec::Infer`'s `sigma`, a Pareto weight), which is
268    /// not part of `trace_of(g)`: use [`Self::score_with_latents`] to draw it
269    /// from its prior, or the SMC/MH drivers, which sample it; and
270    /// [`GenomeError::InvalidStructure`] when the encoding carries sites the
271    /// program never visits.
272    pub fn score(&self, genome: &P::Genome) -> Result<(P::Genome, Trace), GenomeError> {
273        self.prior.validate(genome)?;
274        score_complete(self.prior.trace_of(genome), (self.target_model())())
275    }
276
277    /// Like [`Self::score`], but a site the program visits that the genome's
278    /// encoding lacks — a latent nuisance parameter of the likelihood — is
279    /// **drawn from its prior** with `rng` and kept in the returned trace,
280    /// which is then a complete, fully scored state of the target (the shape
281    /// [`EvolutionChain::step`](super::mh::EvolutionChain::step) needs).
282    /// Sites of the encoding the program never visits are still an error.
283    pub fn score_with_latents<R: Rng>(
284        &self,
285        rng: &mut R,
286        genome: &P::Genome,
287    ) -> Result<(P::Genome, Trace), GenomeError> {
288        self.prior.validate(genome)?;
289        let (g, scored, report) =
290            score_given_trace_reconciled(self.prior.trace_of(genome), rng, (self.target_model())())
291                .map_err(|e| GenomeError::InvalidStructure(e.to_string()))?;
292        if !report.vanished_addresses.is_empty() {
293            check_exact(&ReconcileReport {
294                fresh_addresses: Vec::new(),
295                vanished_addresses: report.vanished_addresses,
296            })?;
297        }
298        Ok((g, scored))
299    }
300
301    /// Unnormalised log target `log π_β(x) = log p(x) + β·log p(data|x)`;
302    /// `−∞` outside the prior's support. Errors as [`Self::score`].
303    pub fn log_boltzmann_target(&self, genome: &P::Genome) -> Result<f64, GenomeError> {
304        Ok(self.score(genome)?.1.total_log_weight())
305    }
306}
307
308#[cfg(test)]
309pub(crate) mod tests {
310    use super::*;
311    use crate::genome::bounds::MultiBounds;
312    use crate::genome::real_vector::RealVector;
313    use crate::genome::traits::RealValuedGenome;
314    use crate::inference::prior::{GaussianPrior, UniformBoxPrior};
315
316    /// A `Clone`-able fitness wrapping a function pointer.
317    #[derive(Clone, Copy)]
318    pub(crate) struct PtrFitness(pub(crate) fn(&RealVector) -> f64);
319
320    impl Fitness for PtrFitness {
321        type Genome = RealVector;
322        type Value = f64;
323        fn evaluate(&self, genome: &RealVector) -> f64 {
324            (self.0)(genome)
325        }
326    }
327
328    pub(crate) fn quad_origin(g: &RealVector) -> f64 {
329        -0.5 * g.genes().iter().map(|x| x * x).sum::<f64>()
330    }
331
332    #[test]
333    fn test_to_weighted_trace_carries_fitness_mass() {
334        // regression: EV-52 — total_log_weight() must equal β·f(x), not 0.
335        let prior = UniformBoxPrior::new(MultiBounds::symmetric(5.0, 2));
336        let model = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_beta(2.0);
337        let genome = RealVector::new(vec![1.0, 2.0]);
338        let f = model.fitness_value(&genome); // -0.5*(1+4) = -2.5
339        let trace = model
340            .to_weighted_trace(&genome)
341            .expect("matching dimension");
342        assert!((trace.total_log_weight() - 2.0 * f).abs() < 1e-9);
343        assert!((trace.log_factors - 2.0 * f).abs() < 1e-9);
344        assert!(trace.total_log_weight().abs() > 1e-6);
345    }
346
347    #[test]
348    fn test_score_composes_prior_and_factor() {
349        // log π_β = log p + β·f, with each part in its own accumulator.
350        let prior = GaussianPrior::new(0.0, 2.0, 2);
351        let model = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_beta(1.5);
352        let g = RealVector::new(vec![0.5, -1.0]);
353        let (decoded, scored) = model.score(&g).expect("matching dimension");
354        assert_eq!(decoded.genes(), g.genes());
355        assert!((scored.log_factors - 1.5 * quad_origin(&g)).abs() < 1e-12);
356        assert!(scored.log_prior.is_finite());
357        assert!(
358            (scored.total_log_weight() - (scored.log_prior + scored.log_factors)).abs() < 1e-12
359        );
360    }
361
362    #[test]
363    fn test_out_of_bounds_scores_neg_inf() {
364        let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 1));
365        let model = EvolutionModel::new(prior, PtrFitness(quad_origin));
366        let g = RealVector::new(vec![3.0]);
367        assert_eq!(model.log_boltzmann_target(&g), Ok(f64::NEG_INFINITY));
368    }
369
370    /// EV-N3: a `RealVector` shorter or longer than the prior's dimension is
371    /// an error, not a panic (short) or a silent truncation (long) — for
372    /// `score`, `log_boltzmann_target`, `to_weighted_trace` and
373    /// `EvolutionChain::init_from` alike.
374    #[test]
375    fn test_score_rejects_dimension_mismatch() {
376        let prior = GaussianPrior::new(0.0, 2.0, 3);
377        let model = EvolutionModel::new(prior, PtrFitness(quad_origin));
378        let short = RealVector::new(vec![0.5, -1.0]);
379        let long = RealVector::new(vec![0.5, -1.0, 0.0, 2.0]);
380        let mismatch = |expected, actual| GenomeError::DimensionMismatch { expected, actual };
381        assert_eq!(model.score(&short).map(|_| ()), Err(mismatch(3, 2)));
382        assert_eq!(model.score(&long).map(|_| ()), Err(mismatch(3, 4)));
383        assert_eq!(model.log_boltzmann_target(&short), Err(mismatch(3, 2)));
384        assert_eq!(
385            model.to_weighted_trace(&long).map(|_| ()),
386            Err(mismatch(3, 4))
387        );
388        let chain = crate::inference::mh::EvolutionChain::new(model.clone());
389        assert!(chain.init_from(&short).is_none());
390        assert_eq!(chain.try_init_from(&long).map(|_| ()), Err(mismatch(3, 4)));
391        // The right dimension still scores.
392        assert!(model.score(&RealVector::new(vec![0.5, -1.0, 0.0])).is_ok());
393    }
394
395    /// EV-N3: a likelihood with a latent nuisance site (`sigma`) cannot be
396    /// scored from the genome's encoding alone — `score` says which site is
397    /// missing instead of panicking, `init_from` returns `None`, and
398    /// `score_with_latents` / `init_from_with_latents` draw the site from its
399    /// prior and return a complete, fully scored state.
400    #[test]
401    fn test_latent_site_likelihood_is_an_error_not_a_panic() {
402        use crate::inference::likelihood::GenomeLikelihood;
403        use fugue::{addr, sample, Normal, Uniform};
404        use rand::rngs::StdRng;
405        use rand::SeedableRng;
406
407        #[derive(Clone)]
408        struct LatentNoise {
409            y: f64,
410        }
411        impl GenomeLikelihood<RealVector> for LatentNoise {
412            fn model(&self, g: &RealVector, beta: f64) -> Model<()> {
413                let (mu, y) = (g.genes()[0], self.y);
414                sample(addr!("sigma"), Uniform::new(0.1, 2.0).unwrap()).bind(move |sigma| {
415                    crate::inference::likelihood::tempered_observe(
416                        addr!("y"),
417                        Normal::new(mu, sigma).unwrap(),
418                        y,
419                        beta,
420                    )
421                })
422            }
423        }
424
425        let model = EvolutionModel::from_likelihood(
426            GaussianPrior::new(0.0, 2.0, 1),
427            LatentNoise { y: 0.4 },
428        );
429        let g = RealVector::new(vec![0.5]);
430        assert_eq!(
431            model.score(&g).map(|_| ()),
432            Err(GenomeError::MissingAddress("sigma".to_string()))
433        );
434        let chain = crate::inference::mh::EvolutionChain::new(model.clone());
435        assert!(chain.init_from(&g).is_none());
436
437        let mut rng = StdRng::seed_from_u64(1);
438        let (decoded, scored) = model
439            .score_with_latents(&mut rng, &g)
440            .expect("latent drawn");
441        assert_eq!(decoded.genes(), g.genes());
442        let sigma = scored.get_f64(&addr!("sigma")).expect("sigma site present");
443        assert!((0.1..2.0).contains(&sigma));
444        assert!(scored.total_log_weight().is_finite());
445        let analytic = fugue::Distribution::log_prob(&Normal::new(0.5, sigma).unwrap(), &0.4);
446        assert!((scored.log_likelihood - analytic).abs() < 1e-12);
447
448        // …and that state drives the chain.
449        let mut chain = chain;
450        let init = chain
451            .init_from_with_latents(&mut rng, &g)
452            .expect("latent drawn");
453        let mut current = init;
454        for _ in 0..50 {
455            let (_g, t) = chain.step(&mut rng, &current);
456            assert!(t.get_f64(&addr!("sigma")).is_some());
457            current = t;
458        }
459    }
460
461    /// EV-N5: `T = 0` / `β = ∞` used to build `factor(∞·f)` — `NaN` at
462    /// `f = 0` — a chain that never moves. Both are rejected up front.
463    #[test]
464    #[should_panic(expected = "must be finite and > 0")]
465    fn test_with_temperature_zero_is_rejected() {
466        let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 1));
467        let _ = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_temperature(0.0);
468    }
469
470    #[test]
471    #[should_panic(expected = "must be finite")]
472    fn test_with_beta_infinite_is_rejected() {
473        let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 1));
474        let _ = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_beta(f64::INFINITY);
475    }
476
477    #[test]
478    fn test_with_temperature_positive_sets_beta() {
479        let prior = UniformBoxPrior::new(MultiBounds::symmetric(1.0, 1));
480        let m = EvolutionModel::new(prior, PtrFitness(quad_origin)).with_temperature(4.0);
481        assert!((m.beta() - 0.25).abs() < 1e-12);
482        assert!((m.temperature() - 4.0).abs() < 1e-12);
483        // Negative β clamps to the prior, as documented.
484        assert_eq!(m.with_beta(-3.0).beta(), 0.0);
485    }
486
487    #[test]
488    fn test_sample_prior_returns_decoded_genome() {
489        use rand::rngs::StdRng;
490        use rand::SeedableRng;
491        let prior = GaussianPrior::new(0.0, 1.0, 4);
492        let model = EvolutionModel::new(prior, PtrFitness(quad_origin));
493        let mut rng = StdRng::seed_from_u64(3);
494        let g = model.sample_prior(&mut rng);
495        assert_eq!(g.genes().len(), 4);
496    }
497
498    /// An observation-program likelihood: per-datum `observe` statements land
499    /// in `log_likelihood`, decomposed from the prior — structure the scalar
500    /// factor could never expose.
501    #[test]
502    fn test_observation_likelihood_scores_in_log_likelihood() {
503        use crate::inference::likelihood::{tempered_observe, GenomeLikelihood};
504        use fugue::{addr, Normal};
505
506        #[derive(Clone)]
507        struct GaussianData {
508            ys: Vec<f64>,
509            sigma: f64,
510        }
511        impl GenomeLikelihood<RealVector> for GaussianData {
512            fn model(&self, g: &RealVector, beta: f64) -> Model<()> {
513                let mu = g.genes()[0];
514                let sigma = self.sigma;
515                let mut m = fugue::pure(());
516                for (k, &y) in self.ys.iter().enumerate() {
517                    m = m.and_then(move |_| {
518                        tempered_observe(addr!("y", k), Normal::new(mu, sigma).unwrap(), y, beta)
519                    });
520                }
521                m
522            }
523        }
524
525        let prior = GaussianPrior::new(0.0, 2.0, 1);
526        let data = GaussianData {
527            ys: vec![0.4, 0.6, 0.5],
528            sigma: 0.5,
529        };
530        let model = EvolutionModel::from_likelihood(prior, data.clone());
531        let g = RealVector::new(vec![0.5]);
532        let (_, scored) = model.score(&g).expect("no latent sites");
533        let normal = Normal::new(0.5, 0.5).unwrap();
534        let analytic: f64 = data
535            .ys
536            .iter()
537            .map(|y| fugue::Distribution::log_prob(&normal, y))
538            .sum();
539        assert!((scored.log_likelihood - analytic).abs() < 1e-12);
540        assert_eq!(scored.log_factors, 0.0);
541        assert!(scored.log_prior.is_finite());
542    }
543}