refactor!: close the remaining API gaps from #21
Three unrelated small defects, all requiring signature changes: - `Game::one_v_one` hardcoded `GameOptions::default()`, so a 1v1 could never set `p_draw` or convergence options — and a drawn 1v1 was therefore unreachable through it, since the default `p_draw` is zero. It now takes `&GameOptions` like every other constructor. - `Observer::on_batch_processed` was declared on the trait and never called from anywhere: implementors wired up a callback that could not fire. It is now called after each slice sweep, and renamed `on_slice_processed` to match the vocabulary the codebase adopted in T2 — the unit of work is a `TimeSlice`, not a batch. A slice is swept once travelling backward and once forward, so a multi-slice history fires it twice per slice per iteration; the doc comment says so. - `pub mod factors` sat beside `pub(crate) mod factor`, two module paths differing by one character with only one of them importable. The public facade is now `graph`. Tests cover each as a behaviour rather than a compile check: a drawn 1v1 succeeds only when p_draw is supplied, and the observer tests fail if any callback stops firing. BREAKING CHANGE: `Game::one_v_one` takes a fourth `&GameOptions` argument; `Observer::on_batch_processed` is renamed `on_slice_processed`; the `factors` module is renamed `graph`. Closes #21 Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_011hcFjNDmHXZF8URGLku5zZ
This commit is contained in:
+11
-9
@@ -532,18 +532,20 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Convenience wrapper over [`Game::ranked`] for two single-player teams.
|
||||||
|
///
|
||||||
/// # Errors
|
/// # Errors
|
||||||
///
|
///
|
||||||
/// Delegates to [`Game::ranked`] with default options, so it returns the
|
/// Delegates to [`Game::ranked`], so it returns the same errors — in
|
||||||
/// same errors — in practice `WrongOutcomeKind` for a non-ranked outcome,
|
/// practice `WrongOutcomeKind` for a non-ranked outcome, or
|
||||||
/// or `TieWithoutDrawProbability` for a draw, since the default `p_draw`
|
/// `TieWithoutDrawProbability` for a draw when `options.p_draw` is zero.
|
||||||
/// applies rather than one you chose.
|
|
||||||
pub fn one_v_one(
|
pub fn one_v_one(
|
||||||
a: &Rating<T, D>,
|
a: &Rating<T, D>,
|
||||||
b: &Rating<T, D>,
|
b: &Rating<T, D>,
|
||||||
outcome: crate::Outcome,
|
outcome: crate::Outcome,
|
||||||
|
options: &GameOptions,
|
||||||
) -> Result<(Gaussian, Gaussian), crate::InferenceError> {
|
) -> Result<(Gaussian, Gaussian), crate::InferenceError> {
|
||||||
let game = Self::ranked(&[&[*a], &[*b]], outcome, &GameOptions::default())?;
|
let game = Self::ranked(&[&[*a], &[*b]], outcome, options)?;
|
||||||
let post = game.posteriors();
|
let post = game.posteriors();
|
||||||
Ok((post[0][0], post[1][0]))
|
Ok((post[0][0], post[1][0]))
|
||||||
}
|
}
|
||||||
@@ -563,11 +565,11 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[doc(hidden)]
|
#[doc(hidden)]
|
||||||
pub fn custom<S: crate::factors::Schedule>(
|
pub fn custom<S: crate::graph::Schedule>(
|
||||||
factors: &mut [crate::factors::BuiltinFactor],
|
factors: &mut [crate::graph::BuiltinFactor],
|
||||||
vars: &mut crate::factors::VarStore,
|
vars: &mut crate::graph::VarStore,
|
||||||
schedule: &S,
|
schedule: &S,
|
||||||
) -> crate::factors::ScheduleReport {
|
) -> crate::graph::ScheduleReport {
|
||||||
schedule.run(factors, vars)
|
schedule.run(factors, vars)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,9 @@
|
|||||||
//! Factor-graph public API.
|
//! Factor-graph public API.
|
||||||
//!
|
//!
|
||||||
|
//! Named `graph` rather than `factors` because the private implementation
|
||||||
|
//! module beside it is `factor`: two module paths differing by one character,
|
||||||
|
//! one public and one not, was a standing invitation to import the wrong one.
|
||||||
|
//!
|
||||||
//! The factor types, `VarStore` and the `Schedule` trait are public so custom
|
//! The factor types, `VarStore` and the `Schedule` trait are public so custom
|
||||||
//! schedules can be written against them.
|
//! schedules can be written against them.
|
||||||
//!
|
//!
|
||||||
@@ -261,6 +261,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
|
|||||||
let old = self.time_slices[j].posteriors();
|
let old = self.time_slices[j].posteriors();
|
||||||
|
|
||||||
self.time_slices[j].new_backward_info(&self.agents);
|
self.time_slices[j].new_backward_info(&self.agents);
|
||||||
|
self.observer.on_slice_processed(
|
||||||
|
&self.time_slices[j].time,
|
||||||
|
j,
|
||||||
|
self.time_slices[j].events.len(),
|
||||||
|
);
|
||||||
|
|
||||||
let new = self.time_slices[j].posteriors();
|
let new = self.time_slices[j].posteriors();
|
||||||
|
|
||||||
@@ -280,6 +285,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
|
|||||||
let old = self.time_slices[j].posteriors();
|
let old = self.time_slices[j].posteriors();
|
||||||
|
|
||||||
self.time_slices[j].new_forward_info(&self.agents);
|
self.time_slices[j].new_forward_info(&self.agents);
|
||||||
|
self.observer.on_slice_processed(
|
||||||
|
&self.time_slices[j].time,
|
||||||
|
j,
|
||||||
|
self.time_slices[j].events.len(),
|
||||||
|
);
|
||||||
|
|
||||||
let new = self.time_slices[j].posteriors();
|
let new = self.time_slices[j].posteriors();
|
||||||
|
|
||||||
@@ -292,6 +302,11 @@ impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O
|
|||||||
let old = self.time_slices[0].posteriors();
|
let old = self.time_slices[0].posteriors();
|
||||||
|
|
||||||
self.time_slices[0].iteration(0, &self.agents);
|
self.time_slices[0].iteration(0, &self.agents);
|
||||||
|
self.observer.on_slice_processed(
|
||||||
|
&self.time_slices[0].time,
|
||||||
|
0,
|
||||||
|
self.time_slices[0].events.len(),
|
||||||
|
);
|
||||||
|
|
||||||
let new = self.time_slices[0].posteriors();
|
let new = self.time_slices[0].posteriors();
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -118,9 +118,9 @@ mod error;
|
|||||||
mod event;
|
mod event;
|
||||||
mod event_builder;
|
mod event_builder;
|
||||||
pub(crate) mod factor;
|
pub(crate) mod factor;
|
||||||
pub mod factors;
|
|
||||||
mod game;
|
mod game;
|
||||||
pub mod gaussian;
|
pub mod gaussian;
|
||||||
|
pub mod graph;
|
||||||
mod history;
|
mod history;
|
||||||
mod key_table;
|
mod key_table;
|
||||||
mod matrix;
|
mod matrix;
|
||||||
|
|||||||
+8
-2
@@ -14,8 +14,13 @@ pub trait Observer<T: Time>: Send + Sync {
|
|||||||
/// Called after each convergence iteration across the whole history.
|
/// Called after each convergence iteration across the whole history.
|
||||||
fn on_iteration_end(&self, _iter: usize, _max_step: (f64, f64)) {}
|
fn on_iteration_end(&self, _iter: usize, _max_step: (f64, f64)) {}
|
||||||
|
|
||||||
/// Called after each time slice is processed within an iteration.
|
/// Called after each time slice is swept within an iteration.
|
||||||
fn on_batch_processed(&self, _time: &T, _slice_idx: usize, _n_events: usize) {}
|
///
|
||||||
|
/// A convergence iteration sweeps every slice twice — once travelling
|
||||||
|
/// backward through the history and once forward — so a multi-slice
|
||||||
|
/// history fires this twice per slice per iteration. A single-slice
|
||||||
|
/// history is swept once and fires once.
|
||||||
|
fn on_slice_processed(&self, _time: &T, _slice_idx: usize, _n_events: usize) {}
|
||||||
|
|
||||||
/// Called once when convergence completes (or max iters is reached).
|
/// Called once when convergence completes (or max iters is reached).
|
||||||
fn on_converged(&self, _iters: usize, _final_step: (f64, f64), _converged: bool) {}
|
fn on_converged(&self, _iters: usize, _final_step: (f64, f64), _converged: bool) {}
|
||||||
@@ -35,6 +40,7 @@ mod tests {
|
|||||||
fn null_observer_compiles_for_i64() {
|
fn null_observer_compiles_for_i64() {
|
||||||
let o = NullObserver;
|
let o = NullObserver;
|
||||||
<NullObserver as Observer<i64>>::on_iteration_end(&o, 1, (0.0, 0.0));
|
<NullObserver as Observer<i64>>::on_iteration_end(&o, 1, (0.0, 0.0));
|
||||||
|
<NullObserver as Observer<i64>>::on_slice_processed(&o, &7, 0, 3);
|
||||||
<NullObserver as Observer<i64>>::on_converged(&o, 5, (1e-6, 1e-6), true);
|
<NullObserver as Observer<i64>>::on_converged(&o, 5, (1e-6, 1e-6), true);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -19,7 +19,8 @@ fn ts_rating(mu: f64, sigma: f64, beta: f64, gamma: f64) -> R {
|
|||||||
fn game_1v1_golden_matches_historical() {
|
fn game_1v1_golden_matches_historical() {
|
||||||
let a = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
|
let a = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
|
||||||
let b = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
|
let b = ts_rating(25.0, 25.0 / 3.0, 25.0 / 6.0, 25.0 / 300.0);
|
||||||
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2)).unwrap();
|
let (a_post, b_post) =
|
||||||
|
Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &GameOptions::default()).unwrap();
|
||||||
// Historical golden from pre-T2 test_1vs1 (team 0 wins):
|
// Historical golden from pre-T2 test_1vs1 (team 0 wins):
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
a_post,
|
a_post,
|
||||||
|
|||||||
+44
-1
@@ -32,7 +32,8 @@ fn game_ranked_1v1_golden() {
|
|||||||
fn game_one_v_one_shortcut() {
|
fn game_one_v_one_shortcut() {
|
||||||
let a = default_rating();
|
let a = default_rating();
|
||||||
let b = default_rating();
|
let b = default_rating();
|
||||||
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2)).unwrap();
|
let (a_post, b_post) =
|
||||||
|
Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &GameOptions::default()).unwrap();
|
||||||
assert!(a_post.mu() > 25.0);
|
assert!(a_post.mu() > 25.0);
|
||||||
assert!(b_post.mu() < 25.0);
|
assert!(b_post.mu() < 25.0);
|
||||||
}
|
}
|
||||||
@@ -95,3 +96,45 @@ fn game_log_evidence_is_finite() {
|
|||||||
assert!(g.log_evidence().is_finite());
|
assert!(g.log_evidence().is_finite());
|
||||||
assert!(g.log_evidence() < 0.0);
|
assert!(g.log_evidence() < 0.0);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// `one_v_one` used to hardcode `GameOptions::default()`, so a 1v1 could
|
||||||
|
/// never set `p_draw` and a drawn 1v1 was unreachable through it.
|
||||||
|
#[test]
|
||||||
|
fn one_v_one_honours_the_draw_probability_it_is_given() {
|
||||||
|
let a = default_rating();
|
||||||
|
let b = default_rating();
|
||||||
|
|
||||||
|
// Default options still reject a draw, because the default p_draw is zero.
|
||||||
|
let err = Game::<i64, _>::one_v_one(&a, &b, Outcome::draw(2), &GameOptions::default())
|
||||||
|
.expect_err("a draw needs a positive p_draw");
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::TieWithoutDrawProbability { .. }
|
||||||
|
));
|
||||||
|
|
||||||
|
// With a draw probability supplied it succeeds — which was impossible
|
||||||
|
// before the signature took options.
|
||||||
|
let options = GameOptions {
|
||||||
|
p_draw: 0.25,
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
let (a_post, b_post) = Game::<i64, _>::one_v_one(&a, &b, Outcome::draw(2), &options)
|
||||||
|
.expect("a draw is representable once p_draw is positive");
|
||||||
|
|
||||||
|
// A symmetric draw leaves the means alone and sharpens both sides.
|
||||||
|
assert!((a_post.mu() - b_post.mu()).abs() < 1e-9);
|
||||||
|
assert!(a_post.sigma() < 25.0 / 3.0);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Convergence options reach the 1v1 path too, not just `p_draw`.
|
||||||
|
#[test]
|
||||||
|
fn one_v_one_honours_convergence_options() {
|
||||||
|
let a = default_rating();
|
||||||
|
let b = default_rating();
|
||||||
|
let options = GameOptions {
|
||||||
|
convergence: ConvergenceOptions::default(),
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
let (a_post, _) = Game::<i64, _>::one_v_one(&a, &b, Outcome::winner(0, 2), &options).unwrap();
|
||||||
|
assert!(a_post.mu() > 25.0);
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,105 @@
|
|||||||
|
//! `Observer` callbacks must actually fire.
|
||||||
|
//!
|
||||||
|
//! `on_slice_processed` (formerly `on_batch_processed`) was declared on the
|
||||||
|
//! trait and never called from anywhere, so implementors wired up a callback
|
||||||
|
//! that could not run. These tests exist so that cannot silently recur.
|
||||||
|
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
use trueskill_tt::{History, Observer};
|
||||||
|
|
||||||
|
/// `History` takes its observer by value and never hands it back, so a test
|
||||||
|
/// that wants to read what was recorded shares the storage rather than the
|
||||||
|
/// observer: the handles are cloned, the buffers are not.
|
||||||
|
#[derive(Clone, Default)]
|
||||||
|
struct Recorder {
|
||||||
|
iterations: Arc<Mutex<Vec<usize>>>,
|
||||||
|
slices: Arc<Mutex<Vec<(i64, usize, usize)>>>,
|
||||||
|
converged: Arc<Mutex<Vec<(usize, bool)>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Observer<i64> for Recorder {
|
||||||
|
fn on_iteration_end(&self, iter: usize, _max_step: (f64, f64)) {
|
||||||
|
self.iterations.lock().unwrap().push(iter);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_slice_processed(&self, time: &i64, slice_idx: usize, n_events: usize) {
|
||||||
|
self.slices
|
||||||
|
.lock()
|
||||||
|
.unwrap()
|
||||||
|
.push((*time, slice_idx, n_events));
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_converged(&self, iters: usize, _final_step: (f64, f64), converged: bool) {
|
||||||
|
self.converged.lock().unwrap().push((iters, converged));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn every_observer_callback_fires() {
|
||||||
|
let recorder = Recorder::default();
|
||||||
|
let mut h = History::builder().observer(recorder.clone()).build();
|
||||||
|
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.record_winner(&"b", &"c", 2).unwrap();
|
||||||
|
h.record_winner(&"c", &"a", 3).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
!recorder.iterations.lock().unwrap().is_empty(),
|
||||||
|
"on_iteration_end never fired"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
!recorder.converged.lock().unwrap().is_empty(),
|
||||||
|
"on_converged never fired"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
!recorder.slices.lock().unwrap().is_empty(),
|
||||||
|
"on_slice_processed never fired — the defect this test exists for"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slice_callbacks_report_the_slice_they_swept() {
|
||||||
|
let recorder = Recorder::default();
|
||||||
|
let mut h = History::builder().observer(recorder.clone()).build();
|
||||||
|
|
||||||
|
h.record_winner(&"a", &"b", 10).unwrap();
|
||||||
|
h.record_winner(&"a", &"b", 20).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let slices = recorder.slices.lock().unwrap();
|
||||||
|
|
||||||
|
// Only the times actually in the history, and each with its own events.
|
||||||
|
for &(time, idx, events) in slices.iter() {
|
||||||
|
assert!(time == 10 || time == 20, "unexpected slice time {time}");
|
||||||
|
assert!(idx < 2, "slice index {idx} out of range");
|
||||||
|
assert_eq!(events, 1, "each slice holds exactly one event");
|
||||||
|
}
|
||||||
|
|
||||||
|
// Both slices must be reported, not just one end of the sweep.
|
||||||
|
assert!(
|
||||||
|
slices.iter().any(|&(t, ..)| t == 10),
|
||||||
|
"slice 10 never reported"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
slices.iter().any(|&(t, ..)| t == 20),
|
||||||
|
"slice 20 never reported"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn a_single_slice_history_still_reports_its_sweep() {
|
||||||
|
let recorder = Recorder::default();
|
||||||
|
let mut h = History::builder().observer(recorder.clone()).build();
|
||||||
|
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let slices = recorder.slices.lock().unwrap();
|
||||||
|
assert!(
|
||||||
|
!slices.is_empty(),
|
||||||
|
"the single-slice path must report its sweep too"
|
||||||
|
);
|
||||||
|
assert!(slices.iter().all(|&(t, idx, _)| t == 1 && idx == 0));
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user