1use 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
37pub(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#[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 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 pub fn fitness_value(&self, genome: &P::Genome) -> f64 {
114 self.likelihood.fitness.evaluate(genome)
115 }
116
117 pub fn log_weight(&self, genome: &P::Genome) -> f64 {
119 self.beta * self.fitness_value(genome)
120 }
121
122 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 pub fn from_likelihood(prior: P, likelihood: L) -> Self {
150 Self {
151 prior,
152 likelihood,
153 beta: 1.0,
154 }
155 }
156
157 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 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 pub fn beta(&self) -> f64 {
192 self.beta
193 }
194
195 pub fn temperature(&self) -> f64 {
197 1.0 / self.beta
198 }
199
200 pub fn prior(&self) -> &P {
202 &self.prior
203 }
204
205 pub fn likelihood(&self) -> &L {
207 &self.likelihood
208 }
209
210 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 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 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 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 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 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 #[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 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); 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 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 #[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 assert!(model.score(&RealVector::new(vec![0.5, -1.0, 0.0])).is_ok());
393 }
394
395 #[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 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, ¤t);
456 assert!(t.get_f64(&addr!("sigma")).is_some());
457 current = t;
458 }
459 }
460
461 #[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 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 #[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}