Closes the last open item in #25. `cargo clippy -W missing_errors_doc -W missing_panics_doc -W must_use_candidate -W doc_markdown` went from 56 warnings to zero. The 13 hand-written sections name the actual variants each function returns rather than gesturing at "an error". Establishing that meant reading the error paths — `Game::ranked` alone returns four distinct variants, and `record_draw` can hit TieWithoutDrawProbability where `record_winner` provably cannot, since a two-team decisive outcome has nothing to tie. Documenting those as interchangeable would have been worse than leaving them undocumented, because a reader would trust it. Two existing doc comments already described panics in prose but not under a `# Panics` heading, so neither rustdoc nor clippy surfaced them: `Outcome::winner` and `EventBuilder::weights`. Both now carry the heading, and `Outcome::winner` gained the note that it ties every loser, so `n >= 3` needs a positive p_draw — the crate's easiest error to hit by accident. The 43 mechanical fixes (31 `#[must_use]` on pure accessors, 11 missing backticks) were applied with `cargo clippy --fix`. `#[must_use]` on Gaussian's arithmetic and on `posteriors()` matters: discarding those results is always a bug, and until now nothing said so. Also documented why `[profile.release] debug = true` exists — cargo-flamegraph needs the symbols, and library profile settings are ignored downstream, so it reads as an oversight without the note. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01T5SYDExxL4vZgvunrcNSMc
183 lines
6.2 KiB
Rust
183 lines
6.2 KiB
Rust
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<f64>,
|
||
}
|
||
|
||
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);
|
||
}
|
||
}
|