use crate::{ N_INF, factor::{Factor, VarId, VarStore}, gaussian::Gaussian, pdf, }; /// Gaussian observation factor on a diff variable. /// /// Encodes the soft evidence `m_obs ~ N(diff, sigma²)`. The outgoing message /// to `diff` is the constant `N(m_obs, sigma²)`, so this factor converges in a /// single propagation: subsequent calls return a zero delta. #[derive(Debug)] pub struct MarginFactor { pub diff: VarId, pub m_obs: f64, pub sigma: f64, pub(crate) msg: Gaussian, pub(crate) evidence_cached: Option, } impl MarginFactor { #[must_use] pub fn new(diff: VarId, m_obs: f64, sigma: f64) -> Self { debug_assert!(sigma > 0.0, "score sigma must be positive"); Self { diff, m_obs, sigma, msg: N_INF, evidence_cached: None, } } } impl MarginFactor { /// Propagate this factor's message, optionally damping the update in /// natural-parameter space. `alpha = 1.0` matches `Factor::propagate` /// exactly; `alpha < 1.0` writes `α·new_msg + (1−α)·old_msg`. pub(crate) fn propagate_with_alpha(&mut self, vars: &mut VarStore, alpha: f64) -> (f64, f64) { let marginal = vars.get(self.diff); let cavity = marginal / self.msg; if self.evidence_cached.is_none() { self.evidence_cached = Some(cavity_evidence(cavity, self.m_obs, self.sigma)); } let new_msg = Gaussian::from_ms(self.m_obs, self.sigma); let damped = self.msg.damp_natural(new_msg, alpha); let old_msg = self.msg; self.msg = damped; vars.set(self.diff, cavity * damped); old_msg.delta(damped) } } impl Factor for MarginFactor { fn propagate(&mut self, vars: &mut VarStore) -> (f64, f64) { self.propagate_with_alpha(vars, 1.0) } fn log_evidence(&self, _vars: &VarStore) -> f64 { self.evidence_cached.unwrap_or(1.0).ln() } } /// Density of the observed margin under the cavity, clamped to a positive /// floor so a far-out observation cannot underflow to `0.0` and make /// `log_evidence` `-inf`. fn cavity_evidence(cavity: Gaussian, m_obs: f64, sigma: f64) -> f64 { let combined_sigma = (cavity.sigma().powi(2) + sigma.powi(2)).sqrt(); pdf(m_obs, cavity.mu(), combined_sigma).max(f64::MIN_POSITIVE) } #[cfg(test)] mod tests { use super::*; #[test] fn first_propagate_writes_tilted_marginal() { let mut vars = VarStore::new(); let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f = MarginFactor::new(diff, 5.0, 1.0); f.propagate(&mut vars); let result = vars.get(diff); // pi = 1/36 + 1 ≈ 1.027778; tau = 0 + 5 = 5 // mu = 5 / 1.027778 ≈ 4.864865; sigma = 1/sqrt(1.027778) ≈ 0.986394 assert!((result.mu() - 4.864864864864865).abs() < 1e-12); assert!((result.sigma() - 0.986393923832144).abs() < 1e-12); } #[test] fn converges_in_one_step() { let mut vars = VarStore::new(); let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f = MarginFactor::new(diff, 5.0, 1.0); f.propagate(&mut vars); let (dmu, dsig) = f.propagate(&mut vars); assert!( dmu < 1e-12, "expected ~0 delta on second propagate, got {dmu}" ); assert!(dsig < 1e-12); } #[test] fn evidence_cached_on_first_propagate() { let mut vars = VarStore::new(); let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f = MarginFactor::new(diff, 5.0, 1.0); assert!(f.evidence_cached.is_none()); f.propagate(&mut vars); let z = f.evidence_cached.unwrap(); // pdf(5, 0, sqrt(37)) ≈ 0.046783 assert!((z - 0.04678300292616668).abs() < 1e-10); // Subsequent propagations don't change it. f.propagate(&mut vars); assert_eq!(f.evidence_cached.unwrap(), z); } #[test] fn log_evidence_matches_cached_ln() { let mut vars = VarStore::new(); let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f = MarginFactor::new(diff, 5.0, 1.0); f.propagate(&mut vars); let logz = f.log_evidence(&vars); assert!((logz - (-3.062235327364623)).abs() < 1e-10); } #[test] fn propagate_with_alpha_one_matches_undamped_propagate() { let mut vars_a = VarStore::new(); let diff_a = vars_a.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f_a = MarginFactor::new(diff_a, 5.0, 1.0); let delta_a = f_a.propagate(&mut vars_a); let result_a = vars_a.get(diff_a); let mut vars_b = VarStore::new(); let diff_b = vars_b.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f_b = MarginFactor::new(diff_b, 5.0, 1.0); let delta_b = f_b.propagate_with_alpha(&mut vars_b, 1.0); let result_b = vars_b.get(diff_b); assert_eq!(result_a.pi(), result_b.pi()); assert_eq!(result_a.tau(), result_b.tau()); assert_eq!(delta_a, delta_b); assert_eq!(f_a.msg.pi(), f_b.msg.pi()); assert_eq!(f_a.msg.tau(), f_b.msg.tau()); } #[test] fn propagate_with_alpha_half_blends_msg_in_natural_params() { // Run undamped to capture (initial_msg, undamped_new_msg). let mut vars_full = VarStore::new(); let diff_full = vars_full.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f_full = MarginFactor::new(diff_full, 5.0, 1.0); let initial_msg_pi = f_full.msg.pi(); let initial_msg_tau = f_full.msg.tau(); f_full.propagate(&mut vars_full); let undamped_msg_pi = f_full.msg.pi(); let undamped_msg_tau = f_full.msg.tau(); // Run damped at α = 0.5 from the same initial state. let mut vars_half = VarStore::new(); let diff_half = vars_half.alloc(Gaussian::from_ms(0.0, 6.0)); let mut f_half = MarginFactor::new(diff_half, 5.0, 1.0); f_half.propagate_with_alpha(&mut vars_half, 0.5); let expected_pi = 0.5 * undamped_msg_pi + 0.5 * initial_msg_pi; let expected_tau = 0.5 * undamped_msg_tau + 0.5 * initial_msg_tau; assert!((f_half.msg.pi() - expected_pi).abs() < 1e-12); assert!((f_half.msg.tau() - expected_tau).abs() < 1e-12); } }