Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
901f60972e | ||
|
|
8116fd081f | ||
|
|
17d072b2ae | ||
|
|
3dd659307a | ||
|
|
7e289ee834 | ||
|
|
564969ee5d | ||
|
|
683813ec10 | ||
|
|
2a48d10aa9 | ||
|
|
d4f91fd221 | ||
|
|
8c087ad015 | ||
|
|
7341669d1a | ||
|
|
2fff745c3b | ||
|
|
3c2f9ac64c | ||
|
|
507894dae7 | ||
|
|
bb2a845882 | ||
|
|
87fca8dcca | ||
|
|
ef62b57a08 | ||
|
|
b2a7ade10c | ||
|
|
617bc07f6f | ||
|
|
1a88678384 | ||
|
|
5f46296671 | ||
|
|
2745fbb622 | ||
|
|
6d2573b92e | ||
|
|
1ac3b21db5 | ||
|
|
4e043364fd | ||
|
|
aff3fb948d | ||
|
|
06ed24b240 | ||
|
|
9b2c2b38c8 | ||
|
|
07285283b6 | ||
|
|
b73cf0145a | ||
|
|
56ff01074f | ||
|
|
8d47e54a8a | ||
|
|
7de092ba12 | ||
|
|
eeb43e3be1 | ||
|
|
69ddebe21d | ||
|
|
9c39d1e681 | ||
|
|
50e11cfbfa | ||
|
|
d4af048914 | ||
|
|
bf9d964cae | ||
|
|
187aede924 | ||
|
|
4fde482e48 | ||
|
|
9e8515b7cd | ||
|
|
9506fed4b3 | ||
|
|
6030dc78de | ||
|
|
355cdb7e05 | ||
|
|
06b6a68499 | ||
|
|
c088214fed | ||
|
|
0f1a1b8911 | ||
|
|
0d32690fcc | ||
|
|
6b8bd786d7 | ||
|
|
f4e2922d59 |
@@ -0,0 +1,15 @@
|
|||||||
|
# `Cargo.toml` sets `publish = ["kellnr"]`, so `cargo publish` targets the
|
||||||
|
# private registry and refuses crates.io. Cargo needs that registry's index
|
||||||
|
# declared to resolve the name.
|
||||||
|
#
|
||||||
|
# Committed rather than left to a per-user `~/.cargo/config.toml` so the repo
|
||||||
|
# is self-contained: a fresh clone, a new machine, or CI would otherwise fail
|
||||||
|
# with
|
||||||
|
#
|
||||||
|
# error: registry index was not found in any configuration: `kellnr`
|
||||||
|
#
|
||||||
|
# Index URL only — it is not a secret. Publish tokens live in
|
||||||
|
# `~/.cargo/credentials.toml` (per-user, never committed) or, in CI, in
|
||||||
|
# `CARGO_REGISTRIES_KELLNR_TOKEN`.
|
||||||
|
[registries.kellnr]
|
||||||
|
index = "sparse+https://crates.aceofba.se/api/v1/crates/"
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
name: CI
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
|
||||||
|
env:
|
||||||
|
CARGO_TERM_COLOR: always
|
||||||
|
RUSTFLAGS: -D warnings
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
test:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
include:
|
||||||
|
# The build most consumers get.
|
||||||
|
- name: default
|
||||||
|
features: ""
|
||||||
|
profile: ""
|
||||||
|
# Most numerical goldens need `approx` for assert_ulps_eq.
|
||||||
|
- name: approx
|
||||||
|
features: "--features approx"
|
||||||
|
profile: ""
|
||||||
|
# The parallel path, including tests/determinism.rs.
|
||||||
|
- name: rayon
|
||||||
|
features: "--features approx,rayon"
|
||||||
|
profile: ""
|
||||||
|
# Critical: debug_assert! is compiled out here, which is where the
|
||||||
|
# tie/p_draw and score_sigma validation actually has to hold.
|
||||||
|
- name: release
|
||||||
|
features: "--features approx"
|
||||||
|
profile: "--release"
|
||||||
|
name: test (${{ matrix.name }})
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: dtolnay/rust-toolchain@stable
|
||||||
|
- uses: Swatinem/rust-cache@v2
|
||||||
|
- run: cargo test ${{ matrix.profile }} ${{ matrix.features }}
|
||||||
|
- run: cargo test ${{ matrix.profile }} ${{ matrix.features }} --doc
|
||||||
|
|
||||||
|
determinism:
|
||||||
|
name: determinism across thread counts
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: dtolnay/rust-toolchain@stable
|
||||||
|
- uses: Swatinem/rust-cache@v2
|
||||||
|
# Posteriors must be bit-identical regardless of how many rayon workers
|
||||||
|
# run the color-group sweep.
|
||||||
|
- run: |
|
||||||
|
for threads in 1 2 4 8; do
|
||||||
|
echo "== RAYON_NUM_THREADS=$threads =="
|
||||||
|
RAYON_NUM_THREADS=$threads cargo test --release \
|
||||||
|
--features approx,rayon --test determinism
|
||||||
|
done
|
||||||
|
|
||||||
|
lint:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: dtolnay/rust-toolchain@stable
|
||||||
|
with:
|
||||||
|
components: clippy
|
||||||
|
- uses: Swatinem/rust-cache@v2
|
||||||
|
- run: cargo clippy --all-targets --all-features -- -D warnings
|
||||||
|
|
||||||
|
format:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
# rustfmt.toml uses nightly-only options (imports_granularity).
|
||||||
|
- uses: dtolnay/rust-toolchain@nightly
|
||||||
|
with:
|
||||||
|
components: rustfmt
|
||||||
|
- run: cargo +nightly fmt --check
|
||||||
|
|
||||||
|
msrv:
|
||||||
|
name: minimum supported Rust version
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: dtolnay/rust-toolchain@1.85.0
|
||||||
|
- uses: Swatinem/rust-cache@v2
|
||||||
|
- run: cargo check --all-targets --features approx,rayon
|
||||||
@@ -7,3 +7,4 @@
|
|||||||
NOTEPAD.md
|
NOTEPAD.md
|
||||||
|
|
||||||
/.claude
|
/.claude
|
||||||
|
proptest-regressions/
|
||||||
|
|||||||
+142
@@ -2,6 +2,144 @@
|
|||||||
|
|
||||||
All notable changes to this project will be documented in this file.
|
All notable changes to this project will be documented in this file.
|
||||||
|
|
||||||
|
## 0.4.2 - 2026-09-07
|
||||||
|
|
||||||
|
### Bug Fixes
|
||||||
|
|
||||||
|
- fix: replace the erfc approximation with libm, for free
|
||||||
|
- fix: route every transcendental through libm, and combine sigmas with hypot
|
||||||
|
|
||||||
|
### Testing
|
||||||
|
|
||||||
|
- test: localise the erfc_inv tail residual to the caller's argument
|
||||||
|
|
||||||
|
## 0.4.1 - 2026-09-07
|
||||||
|
|
||||||
|
### Bug Fixes
|
||||||
|
|
||||||
|
- fix: correct erfc_inv's sign error and keep evidence in log space
|
||||||
|
|
||||||
|
### Miscellaneous Tasks
|
||||||
|
|
||||||
|
- chore: Release trueskill-tt version 0.4.1
|
||||||
|
|
||||||
|
### Testing
|
||||||
|
|
||||||
|
- test: pin quality()'s N-group closed form, closing the README cross-check
|
||||||
|
|
||||||
|
## 0.4.0 - 2026-09-07
|
||||||
|
|
||||||
|
### Breaking Changes
|
||||||
|
|
||||||
|
- feat!: N-team outcome prediction with draw mass, replacing the 2-team panic
|
||||||
|
- refactor!: close the remaining API gaps from #21
|
||||||
|
- fix!: apply competitor configuration whenever it is supplied
|
||||||
|
|
||||||
|
### Bug Fixes
|
||||||
|
|
||||||
|
- fix(release): skip the changelog hook during a dry run
|
||||||
|
- fix: stop destroying tail precision in evidence and truncation
|
||||||
|
- fix: reject convergence options that silently disable inference
|
||||||
|
|
||||||
|
### Documentation
|
||||||
|
|
||||||
|
- docs: correct drifted documentation and compile the README in CI
|
||||||
|
|
||||||
|
### Features
|
||||||
|
|
||||||
|
- feat: add expected information gain for active matchup selection
|
||||||
|
- feat: let observers be shared, boxed, or borrowed
|
||||||
|
|
||||||
|
### Miscellaneous Tasks
|
||||||
|
|
||||||
|
- chore: Release trueskill-tt version 0.4.0
|
||||||
|
|
||||||
|
## 0.3.0 - 2026-09-01
|
||||||
|
|
||||||
|
### Breaking Changes
|
||||||
|
|
||||||
|
- refactor!: make Competitor::message an Option, and compute_elapsed loud
|
||||||
|
- refactor!: replace emptiness-as-sentinel with Option for results and weights
|
||||||
|
- refactor!: remove ConvergenceReport::slices_skipped
|
||||||
|
|
||||||
|
### Bug Fixes
|
||||||
|
|
||||||
|
- fix: enforce EventBuilder weight/team length in release
|
||||||
|
|
||||||
|
### Documentation
|
||||||
|
|
||||||
|
- docs: complete the public API documentation contract
|
||||||
|
|
||||||
|
### Features
|
||||||
|
|
||||||
|
- feat: allow drift to vary per competitor via Member::with_drift_scale
|
||||||
|
|
||||||
|
### Miscellaneous Tasks
|
||||||
|
|
||||||
|
- chore: ignore proptest regression seed files
|
||||||
|
- chore: Release trueskill-tt version 0.3.0
|
||||||
|
|
||||||
|
### Performance
|
||||||
|
|
||||||
|
- perf: stop cloning inference inputs in OwnedGame and ingestion
|
||||||
|
- perf: make the per-slice SkillStore compact instead of dense
|
||||||
|
|
||||||
|
### Testing
|
||||||
|
|
||||||
|
- test: add property-based tests, a shared finiteness helper, and boundary inputs
|
||||||
|
|
||||||
|
## 0.2.0 - 2026-08-27
|
||||||
|
|
||||||
|
### Breaking Changes
|
||||||
|
|
||||||
|
- refactor!: remove the inert online flag
|
||||||
|
|
||||||
|
### Bug Fixes
|
||||||
|
|
||||||
|
- fix: reject ties without draw probability; never report NaN as converged
|
||||||
|
- fix(quality): support any number of rating groups
|
||||||
|
- fix(evidence): accumulate in log space and floor the per-link value
|
||||||
|
- fix(history): stop reprocessing the slice that was just appended to
|
||||||
|
- fix(rayon): remove the aliasing unsafe from the parallel sweep
|
||||||
|
- fix: close out four small issues and pin #27's repro
|
||||||
|
|
||||||
|
### Documentation
|
||||||
|
|
||||||
|
- docs: refresh README and CLAUDE.md; add ingest benchmark
|
||||||
|
- docs: spec for filtered (forward-only) estimates
|
||||||
|
- docs: implementation plan for filtered estimates
|
||||||
|
- docs: state filtered accessor cost and evidence semantics precisely
|
||||||
|
- docs(cargo): correct the licence note — kellnr does not require one
|
||||||
|
|
||||||
|
### Features
|
||||||
|
|
||||||
|
- feat: add filtered_log_evidence
|
||||||
|
- feat: add filtered learning curves
|
||||||
|
|
||||||
|
### Miscellaneous Tasks
|
||||||
|
|
||||||
|
- chore: add CI, crate metadata, and crate-level documentation
|
||||||
|
- chore: target releases at the private kellnr registry
|
||||||
|
- chore: keep the 48 MB ATP dataset out of the published crate
|
||||||
|
- chore: dual-license MIT OR Apache-2.0
|
||||||
|
- chore: Release trueskill-tt version 0.2.0
|
||||||
|
|
||||||
|
### Performance
|
||||||
|
|
||||||
|
- perf(gaussian): drop the sqrt round-trip from variance-space operations
|
||||||
|
|
||||||
|
### Refactor
|
||||||
|
|
||||||
|
- refactor: unify convergence defaults, validate builders, clear dead code
|
||||||
|
|
||||||
|
### Styling
|
||||||
|
|
||||||
|
- style: make NaN rejection explicit in score_sigma validation
|
||||||
|
|
||||||
|
### Testing
|
||||||
|
|
||||||
|
- test: pin the invariants that make filtered estimates trustworthy
|
||||||
|
|
||||||
## 0.1.2 - 2026-06-12
|
## 0.1.2 - 2026-06-12
|
||||||
|
|
||||||
### Bug Fixes
|
### Bug Fixes
|
||||||
@@ -32,6 +170,10 @@ All notable changes to this project will be documented in this file.
|
|||||||
- feat(outcome): per-event score_sigma override on Outcome::Scored
|
- feat(outcome): per-event score_sigma override on Outcome::Scored
|
||||||
- feat(event_builder): expose scores_with_sigma fluent method
|
- feat(event_builder): expose scores_with_sigma fluent method
|
||||||
|
|
||||||
|
### Miscellaneous Tasks
|
||||||
|
|
||||||
|
- chore: Release trueskill-tt version 0.1.2
|
||||||
|
|
||||||
### Refactor
|
### Refactor
|
||||||
|
|
||||||
- refactor: dedupe Game::likelihoods and likelihoods_scored via run_chain
|
- refactor: dedupe Game::likelihoods and likelihoods_scored via run_chain
|
||||||
|
|||||||
@@ -5,42 +5,116 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
|
|||||||
## Commands
|
## Commands
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
cargo build # Build the library
|
just test # Full suite across every feature combination CI checks
|
||||||
cargo test --lib # Run all library tests
|
just check # Fast inner loop: cargo test --features approx
|
||||||
cargo test --lib <test_name> # Run a single test by name
|
just lint # clippy, warnings denied
|
||||||
cargo test --lib -- --nocapture # Run tests with stdout output
|
just fmt # ALWAYS nightly — rustfmt.toml uses nightly-only options
|
||||||
cargo clippy # Lint
|
just determinism # Bit-identical posteriors at RAYON_NUM_THREADS 1/2/4/8
|
||||||
cargo bench # Run benchmarks (criterion)
|
just ci # Everything CI runs
|
||||||
|
cargo test --lib <test_name> # A single test by name
|
||||||
|
cargo bench # Criterion benchmarks
|
||||||
```
|
```
|
||||||
|
|
||||||
The `approx` feature enables `approx::AbsDiffEq` for `Gaussian`:
|
**Run tests in release too.** `debug_assert!` is compiled out there, and that
|
||||||
```bash
|
is where several defects have hidden — a debug-only run is not evidence.
|
||||||
cargo test --features approx
|
`just test` includes a release job.
|
||||||
```
|
|
||||||
|
### Feature flags
|
||||||
|
|
||||||
|
- `approx` — `approx::AbsDiffEq` etc. for `Gaussian`. Most numerical goldens need it.
|
||||||
|
- `rayon` — opt-in parallel within-slice sweep and per-slice query passes.
|
||||||
|
|
||||||
## Architecture
|
## Architecture
|
||||||
|
|
||||||
This is a Rust port of [TrueSkillThroughTime.py](https://github.com/glandfried/TrueSkillThroughTime.py) — a Bayesian skill rating system that tracks skill evolution over time using Gaussian message passing.
|
A Rust port of [TrueSkillThroughTime.py](https://github.com/glandfried/TrueSkillThroughTime.py):
|
||||||
|
Bayesian skill rating that infers skill at every point in time, propagating
|
||||||
|
evidence both forward and backward across a history.
|
||||||
|
|
||||||
### Data flow
|
### Data flow
|
||||||
|
|
||||||
|
Ingestion (public types, `event.rs`):
|
||||||
|
|
||||||
```
|
```
|
||||||
History → Batch[] → Game[] → teams/players
|
Event<T, K> → Team<K>[] → Member<K>[]
|
||||||
```
|
```
|
||||||
|
|
||||||
- **`History`** (`history.rs`) — top-level container. Organizes games by time into `Batch`es, runs forward/backward message passing across batches, and exposes `learning_curves()` and `log_evidence()`.
|
`History::add_events` flattens that into indices; teams survive only as
|
||||||
- **`Batch`** (`batch.rs`) — all games at a single time step. Runs `iteration()` to update skill estimates via `Game::posteriors()`, collecting `Skill` distributions per player.
|
grouping, not as a value. Inference then runs on the internal shapes:
|
||||||
- **`Game`** (`game.rs`) — a single match. Given teams (slices of `Gaussian`), computes posterior skill distributions using Gaussian factor graphs and `message.rs` helpers.
|
|
||||||
- **`Agent`** (`agent.rs`) — wraps a `Player` with temporal state (`last_time`, `message`). `receive()` applies time-decay (`gamma`) when the player reappears after a gap.
|
|
||||||
- **`Player`** (`player.rs`) — static configuration: prior `Gaussian`, `beta` (performance noise), `gamma` (skill drift per time unit).
|
|
||||||
- **`Gaussian`** (`gaussian.rs`) — core probability type. Stored as natural parameters (`pi = 1/sigma²`, `tau = mu/sigma²`). Arithmetic ops implement message multiplication/division in the factor graph.
|
|
||||||
- **`message.rs`** — `TeamMessage` and `DiffMessage`: intermediate factor graph messages used inside `Game`.
|
|
||||||
- **`MarginFactor`** (`factor/margin.rs`) — Gaussian observation factor on a diff variable; engaged by `Outcome::Scored`.
|
|
||||||
- **`lib.rs`** — exports the public API (`Game`, `Gaussian`, `History`, `Player`) and standalone functions (`quality()`, `pdf()`, `cdf()`, `erfc()`). Also defines global defaults: `MU=0.0`, `SIGMA=6.0`, `BETA=1.0`, `GAMMA=0.03`, `P_DRAW=0.0`, `EPSILON=1e-6`, `ITERATIONS=30`.
|
|
||||||
|
|
||||||
### Key design points
|
```
|
||||||
|
History → TimeSlice[] → Event[] → Item[]
|
||||||
|
↓
|
||||||
|
Game (factor graph) → Schedule → BuiltinFactor[]
|
||||||
|
```
|
||||||
|
|
||||||
- `History` uses `IndexMap<K>` (defined in `lib.rs`) to map arbitrary player keys to `Agent` state.
|
- **`History`** (`history.rs`) — top level. Interns keys, groups events into
|
||||||
- Convergence is measured by the maximum `delta()` across all skill distributions; iteration stops when below `EPSILON` or after `ITERATIONS` rounds.
|
`TimeSlice`s by time, runs the forward/backward sweep in `converge()`, and
|
||||||
- The `approx` feature gates `AbsDiffEq` on `Gaussian` for use in tests — the feature is optional and only needed for approximate equality assertions.
|
answers `learning_curves()`, `current_skill()`, `log_evidence()`,
|
||||||
- `time` in `History`/`Batch` is currently an `f64`; the README notes it needs to become an enum to support richer temporal states.
|
`predict_quality()`, `predict_outcome()`. Built via `HistoryBuilder`.
|
||||||
|
- **`TimeSlice`** (`time_slice.rs`) — all events at one time. Owns a
|
||||||
|
`SkillStore` and a `ScratchArena`; `iteration()` sweeps its events, using
|
||||||
|
`ColorGroups` to partition independent ones.
|
||||||
|
- **`Event`** — two distinct types, do not confuse them. The *public* ingestion
|
||||||
|
`Event<T, K>` is in `event.rs` (with `Team`/`Member`); the *internal*
|
||||||
|
`pub(crate) Event` in `time_slice.rs` is one match during inference, where
|
||||||
|
`compute()` runs inference reading skills immutably and `apply()` folds the
|
||||||
|
result back. That split is what lets a color group run in parallel with no
|
||||||
|
`unsafe`.
|
||||||
|
- **`Game`** (`game.rs`) — a single match's factor graph. `run_chain` builds the
|
||||||
|
diff chain between rank-adjacent teams and drives it to convergence.
|
||||||
|
- **`Gaussian`** (`gaussian.rs`) — natural parameters (`pi = 1/sigma²`,
|
||||||
|
`tau = mu/sigma²`). `Mul`/`Div` are the EP product/cavity: pure adds and
|
||||||
|
subtracts. Variance-space ops (`Add`, `Sub`, `exclude`, `forget`) go through
|
||||||
|
`from_mv`/`variance()` and take no square root.
|
||||||
|
- **`factor/`** — `TeamSumFactor`, `RankDiffFactor`, `TruncFactor` (ranked),
|
||||||
|
`MarginFactor` (scored), over a flat `VarStore`. `BuiltinFactor` dispatches
|
||||||
|
by enum rather than `dyn`.
|
||||||
|
- **`Schedule`** (`schedule.rs`) — drives factor propagation. `EpsilonOrMax` is
|
||||||
|
the only implementation.
|
||||||
|
- **`Competitor`** (`competitor.rs`) — per-history temporal state (`message`,
|
||||||
|
`last_time`). **`Rating`** (`rating.rs`) — static config (prior, `beta`, drift).
|
||||||
|
- **`storage/`** — `SkillStore` (per slice, `pub(crate)`) and `CompetitorStore`
|
||||||
|
(per history, public), both indexed by `Index`. The module is `pub`, but only
|
||||||
|
`CompetitorStore` is reachable from outside the crate.
|
||||||
|
- **`KeyTable`** (`key_table.rs`) — user key ↔ `Index`, both directions O(1).
|
||||||
|
- **`Drift`** (`drift.rs`) / **`Time`** (`time.rs`) — traits. `Time` is a *trait*
|
||||||
|
(`i64`, `Untimed`), not an enum.
|
||||||
|
- **`lib.rs`** — public exports, global defaults (`MU`, `SIGMA`, `BETA`,
|
||||||
|
`GAMMA`, `P_DRAW`, `EPSILON`, `ITERATIONS`), and the standalone `quality()`.
|
||||||
|
The `cdf()` / `erfc()` helpers live here too but are `pub(crate)` and private
|
||||||
|
respectively — not public API.
|
||||||
|
|
||||||
|
### Invariants worth knowing
|
||||||
|
|
||||||
|
- **A tie needs `p_draw > 0`.** With `p_draw == 0.0` the truncation margin is
|
||||||
|
zero and the two-sided tie update evaluates `0/0`. Ingestion rejects such
|
||||||
|
events with `InferenceError::TieWithoutDrawProbability`. This includes
|
||||||
|
`Outcome::winner(w, n)` for `n >= 3`, which ties every loser.
|
||||||
|
- **NaN is never convergence.** Comparisons against NaN are all false, so
|
||||||
|
`tuple_gt` reads NaN as "below epsilon". Use `step_converged` /
|
||||||
|
`step_is_finite`, never `!tuple_gt(..)` alone.
|
||||||
|
- **Evidence accumulates in log space.** A linear product over a long diff
|
||||||
|
chain underflows to zero, and `ln(0)` is `-inf`.
|
||||||
|
- **Colors are contiguous.** `recompute_color_groups` reorders events so each
|
||||||
|
color occupies one range; `ColorGroups::groups_are_contiguous` asserts it.
|
||||||
|
- **Transcendentals go through `libm`, not `std`.** IEEE 754 pins the basic
|
||||||
|
operations and `sqrt` but says nothing about `exp`/`log`/`erf`, and `std`
|
||||||
|
delegates to the *system* math library — measured, `f64::exp` and `libm::exp`
|
||||||
|
disagree on 9.7% of inputs by one ULP. Since inference is an iterative fixed
|
||||||
|
point, one ULP can change an iteration count. Use `libm::exp` / `libm::log` in
|
||||||
|
inference code; `f64::sqrt` is fine (IEEE specifies it). Tests may use either.
|
||||||
|
- **The crate is `#![forbid(unsafe_code)]`.** Keep it that way.
|
||||||
|
- **Ingestion order must not change the answer.** Events added one at a time
|
||||||
|
must converge to the same fixed point as the same events batched — see
|
||||||
|
`tests/ingestion_equivalence.rs`.
|
||||||
|
|
||||||
|
### Testing notes
|
||||||
|
|
||||||
|
- Numerical goldens are cross-validated against the Python/Julia reference.
|
||||||
|
Some are *convergence residuals*, not exact values; treat a small movement
|
||||||
|
as suspicious but check whether the new value is closer to the analytic
|
||||||
|
truth (symmetric fixtures converge to their prior mean exactly) before
|
||||||
|
assuming a regression.
|
||||||
|
- `tests/degenerate_inputs.rs` covers empty/boundary/error paths,
|
||||||
|
`tests/ingestion_equivalence.rs` covers batching order, `tests/quality.rs`
|
||||||
|
covers N-group quality, `tests/determinism.rs` covers thread counts.
|
||||||
|
|||||||
+34
-1
@@ -1,7 +1,30 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "trueskill-tt"
|
name = "trueskill-tt"
|
||||||
version = "0.1.2"
|
version = "0.4.2"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
|
rust-version = "1.85"
|
||||||
|
description = "TrueSkill Through Time: Bayesian skill rating that tracks how skill evolves over time, via Gaussian message passing"
|
||||||
|
repository = "https://git.aceofba.se/logaritmisk/trueskill-tt"
|
||||||
|
authors = ["Anders Olsson"]
|
||||||
|
# Publishing is restricted to the private kellnr registry; this also makes
|
||||||
|
# an accidental `cargo publish` to crates.io a hard error rather than a
|
||||||
|
# irreversible mistake. Index is declared in `.cargo/config.toml`.
|
||||||
|
publish = ["kellnr"]
|
||||||
|
readme = "README.md"
|
||||||
|
keywords = ["trueskill", "rating", "bayesian", "elo", "skill"]
|
||||||
|
categories = ["algorithms", "science", "game-development"]
|
||||||
|
license = "MIT OR Apache-2.0"
|
||||||
|
# `examples/atp.csv` is a 48 MB tennis dataset — 99% of the packaged crate,
|
||||||
|
# for a library whose source is 312 KB. `examples/atp.rs` opens it by
|
||||||
|
# relative path at runtime, so excluding the data still compiles; the
|
||||||
|
# example just needs the file fetched from the repo to run.
|
||||||
|
exclude = [
|
||||||
|
"/docs",
|
||||||
|
"/benches/*.txt",
|
||||||
|
"/temp",
|
||||||
|
"/.gitea",
|
||||||
|
"/examples/atp.csv",
|
||||||
|
]
|
||||||
|
|
||||||
[lib]
|
[lib]
|
||||||
bench = false
|
bench = false
|
||||||
@@ -22,8 +45,13 @@ harness = false
|
|||||||
name = "scored"
|
name = "scored"
|
||||||
harness = false
|
harness = false
|
||||||
|
|
||||||
|
[[bench]]
|
||||||
|
name = "ingest"
|
||||||
|
harness = false
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
approx = { version = "0.5.1", optional = true }
|
approx = { version = "0.5.1", optional = true }
|
||||||
|
libm = "0.2.16"
|
||||||
rayon = { version = "1", optional = true }
|
rayon = { version = "1", optional = true }
|
||||||
smallvec = "1"
|
smallvec = "1"
|
||||||
|
|
||||||
@@ -35,9 +63,14 @@ rayon = ["dep:rayon"]
|
|||||||
criterion = "0.5"
|
criterion = "0.5"
|
||||||
plotters = { version = "0.3", default-features = false, features = ["svg_backend", "all_elements", "all_series"] }
|
plotters = { version = "0.3", default-features = false, features = ["svg_backend", "all_elements", "all_series"] }
|
||||||
plotters-backend = "0.3"
|
plotters-backend = "0.3"
|
||||||
|
proptest = "1.11.0"
|
||||||
time = { version = "0.3", features = ["parsing"] }
|
time = { version = "0.3", features = ["parsing"] }
|
||||||
trueskill-tt = { path = ".", features = ["approx"] }
|
trueskill-tt = { path = ".", features = ["approx"] }
|
||||||
|
|
||||||
|
# Debug symbols in release are for `just flame` (cargo-flamegraph), which needs
|
||||||
|
# them to symbolicate. Profile settings in a library are ignored by downstream
|
||||||
|
# consumers, so these only affect local builds — this is deliberate, not an
|
||||||
|
# oversight.
|
||||||
[profile.release]
|
[profile.release]
|
||||||
debug = true
|
debug = true
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,39 @@
|
|||||||
alias b := bench
|
alias b := bench
|
||||||
|
alias t := test
|
||||||
|
|
||||||
|
# Run the full test suite across the feature combinations CI checks.
|
||||||
|
test:
|
||||||
|
cargo test
|
||||||
|
cargo test --features approx
|
||||||
|
cargo test --features approx,rayon
|
||||||
|
cargo test --release --features approx
|
||||||
|
|
||||||
|
# Fast inner-loop tests.
|
||||||
|
check:
|
||||||
|
cargo test --features approx
|
||||||
|
|
||||||
|
# Posteriors must be bit-identical across rayon worker counts.
|
||||||
|
determinism:
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
for threads in 1 2 4 8; do
|
||||||
|
echo "== RAYON_NUM_THREADS=$threads =="
|
||||||
|
RAYON_NUM_THREADS=$threads cargo test --release \
|
||||||
|
--features approx,rayon --test determinism
|
||||||
|
done
|
||||||
|
|
||||||
|
lint:
|
||||||
|
cargo clippy --all-targets --all-features -- -D warnings
|
||||||
|
|
||||||
|
# Always nightly: rustfmt.toml uses nightly-only options.
|
||||||
|
fmt:
|
||||||
|
cargo +nightly fmt
|
||||||
|
|
||||||
|
fmt-check:
|
||||||
|
cargo +nightly fmt --check
|
||||||
|
|
||||||
|
# Everything CI runs.
|
||||||
|
ci: fmt-check lint test determinism
|
||||||
|
|
||||||
store:
|
store:
|
||||||
cargo bench -- --save-baseline base
|
cargo bench -- --save-baseline base
|
||||||
@@ -8,3 +43,49 @@ bench:
|
|||||||
|
|
||||||
flame:
|
flame:
|
||||||
cargo flamegraph --root --example atp
|
cargo flamegraph --root --example atp
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Release workflow
|
||||||
|
#
|
||||||
|
# Publishing goes to the private kellnr registry only: `Cargo.toml` sets
|
||||||
|
# `publish = ["kellnr"]`, so an accidental `cargo publish` to crates.io is a
|
||||||
|
# hard error rather than an irreversible mistake. The index is declared in the
|
||||||
|
# committed `.cargo/config.toml`; the token is per-user and lives in
|
||||||
|
# `~/.cargo/credentials.toml` (`cargo login --registry kellnr`).
|
||||||
|
#
|
||||||
|
# Step 1: just release-plan [level] — dry run, no writes
|
||||||
|
# Step 2: just release [level] — bump, changelog, tag, publish, push
|
||||||
|
#
|
||||||
|
# LEVEL is the cargo-release bump level (default `minor`). On 0.x:
|
||||||
|
# minor -> breaking bump (0.1.2 -> 0.2.0) <- any public-API change
|
||||||
|
# patch -> additive only (0.1.2 -> 0.1.3)
|
||||||
|
# major -> reserved for the 1.0.0 jump
|
||||||
|
#
|
||||||
|
# `release.toml` regenerates CHANGELOG.md with git-cliff in a pre-release hook
|
||||||
|
# and keeps push = false; this recipe pushes last, after publish has succeeded.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
# Dry-run preview of the next release. Inspect the version bump and the
|
||||||
|
# "Publishing ..." line before running `just release`.
|
||||||
|
release-plan level="minor":
|
||||||
|
cargo release {{level}}
|
||||||
|
|
||||||
|
# Cut a release from a clean main: gate -> bump -> tag -> publish -> push.
|
||||||
|
release level="minor":
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
if [[ "$(git branch --show-current)" != "main" ]]; then
|
||||||
|
echo "error: run 'just release' from the 'main' branch" >&2; exit 1
|
||||||
|
fi
|
||||||
|
if [[ -n "$(git status --porcelain)" ]]; then
|
||||||
|
echo "error: working tree is dirty — commit or stash first" >&2; exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
# cargo-release only verify-compiles the packaged crate; it does not run the
|
||||||
|
# suite, and publishing is irreversible. Run the same gate CI does, which
|
||||||
|
# includes the release profile where debug_assert! is compiled out.
|
||||||
|
just ci
|
||||||
|
|
||||||
|
cargo release {{level}} --execute --no-confirm
|
||||||
|
git push --follow-tags
|
||||||
|
|||||||
+201
@@ -0,0 +1,201 @@
|
|||||||
|
Apache License
|
||||||
|
Version 2.0, January 2004
|
||||||
|
http://www.apache.org/licenses/
|
||||||
|
|
||||||
|
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||||
|
|
||||||
|
1. Definitions.
|
||||||
|
|
||||||
|
"License" shall mean the terms and conditions for use, reproduction,
|
||||||
|
and distribution as defined by Sections 1 through 9 of this document.
|
||||||
|
|
||||||
|
"Licensor" shall mean the copyright owner or entity authorized by
|
||||||
|
the copyright owner that is granting the License.
|
||||||
|
|
||||||
|
"Legal Entity" shall mean the union of the acting entity and all
|
||||||
|
other entities that control, are controlled by, or are under common
|
||||||
|
control with that entity. For the purposes of this definition,
|
||||||
|
"control" means (i) the power, direct or indirect, to cause the
|
||||||
|
direction or management of such entity, whether by contract or
|
||||||
|
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||||
|
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||||
|
|
||||||
|
"You" (or "Your") shall mean an individual or Legal Entity
|
||||||
|
exercising permissions granted by this License.
|
||||||
|
|
||||||
|
"Source" form shall mean the preferred form for making modifications,
|
||||||
|
including but not limited to software source code, documentation
|
||||||
|
source, and configuration files.
|
||||||
|
|
||||||
|
"Object" form shall mean any form resulting from mechanical
|
||||||
|
transformation or translation of a Source form, including but
|
||||||
|
not limited to compiled object code, generated documentation,
|
||||||
|
and conversions to other media types.
|
||||||
|
|
||||||
|
"Work" shall mean the work of authorship, whether in Source or
|
||||||
|
Object form, made available under the License, as indicated by a
|
||||||
|
copyright notice that is included in or attached to the work
|
||||||
|
(an example is provided in the Appendix below).
|
||||||
|
|
||||||
|
"Derivative Works" shall mean any work, whether in Source or Object
|
||||||
|
form, that is based on (or derived from) the Work and for which the
|
||||||
|
editorial revisions, annotations, elaborations, or other modifications
|
||||||
|
represent, as a whole, an original work of authorship. For the purposes
|
||||||
|
of this License, Derivative Works shall not include works that remain
|
||||||
|
separable from, or merely link (or bind by name) to the interfaces of,
|
||||||
|
the Work and Derivative Works thereof.
|
||||||
|
|
||||||
|
"Contribution" shall mean any work of authorship, including
|
||||||
|
the original version of the Work and any modifications or additions
|
||||||
|
to that Work or Derivative Works thereof, that is intentionally
|
||||||
|
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||||
|
or by an individual or Legal Entity authorized to submit on behalf of
|
||||||
|
the copyright owner. For the purposes of this definition, "submitted"
|
||||||
|
means any form of electronic, verbal, or written communication sent
|
||||||
|
to the Licensor or its representatives, including but not limited to
|
||||||
|
communication on electronic mailing lists, source code control systems,
|
||||||
|
and issue tracking systems that are managed by, or on behalf of, the
|
||||||
|
Licensor for the purpose of discussing and improving the Work, but
|
||||||
|
excluding communication that is conspicuously marked or otherwise
|
||||||
|
designated in writing by the copyright owner as "Not a Contribution."
|
||||||
|
|
||||||
|
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||||
|
on behalf of whom a Contribution has been received by Licensor and
|
||||||
|
subsequently incorporated within the Work.
|
||||||
|
|
||||||
|
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
copyright license to reproduce, prepare Derivative Works of,
|
||||||
|
publicly display, publicly perform, sublicense, and distribute the
|
||||||
|
Work and such Derivative Works in Source or Object form.
|
||||||
|
|
||||||
|
3. Grant of Patent License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
(except as stated in this section) patent license to make, have made,
|
||||||
|
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||||
|
where such license applies only to those patent claims licensable
|
||||||
|
by such Contributor that are necessarily infringed by their
|
||||||
|
Contribution(s) alone or by combination of their Contribution(s)
|
||||||
|
with the Work to which such Contribution(s) was submitted. If You
|
||||||
|
institute patent litigation against any entity (including a
|
||||||
|
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||||
|
or a Contribution incorporated within the Work constitutes direct
|
||||||
|
or contributory patent infringement, then any patent licenses
|
||||||
|
granted to You under this License for that Work shall terminate
|
||||||
|
as of the date such litigation is filed.
|
||||||
|
|
||||||
|
4. Redistribution. You may reproduce and distribute copies of the
|
||||||
|
Work or Derivative Works thereof in any medium, with or without
|
||||||
|
modifications, and in Source or Object form, provided that You
|
||||||
|
meet the following conditions:
|
||||||
|
|
||||||
|
(a) You must give any other recipients of the Work or
|
||||||
|
Derivative Works a copy of this License; and
|
||||||
|
|
||||||
|
(b) You must cause any modified files to carry prominent notices
|
||||||
|
stating that You changed the files; and
|
||||||
|
|
||||||
|
(c) You must retain, in the Source form of any Derivative Works
|
||||||
|
that You distribute, all copyright, patent, trademark, and
|
||||||
|
attribution notices from the Source form of the Work,
|
||||||
|
excluding those notices that do not pertain to any part of
|
||||||
|
the Derivative Works; and
|
||||||
|
|
||||||
|
(d) If the Work includes a "NOTICE" text file as part of its
|
||||||
|
distribution, then any Derivative Works that You distribute must
|
||||||
|
include a readable copy of the attribution notices contained
|
||||||
|
within such NOTICE file, excluding those notices that do not
|
||||||
|
pertain to any part of the Derivative Works, in at least one
|
||||||
|
of the following places: within a NOTICE text file distributed
|
||||||
|
as part of the Derivative Works; within the Source form or
|
||||||
|
documentation, if provided along with the Derivative Works; or,
|
||||||
|
within a display generated by the Derivative Works, if and
|
||||||
|
wherever such third-party notices normally appear. The contents
|
||||||
|
of the NOTICE file are for informational purposes only and
|
||||||
|
do not modify the License. You may add Your own attribution
|
||||||
|
notices within Derivative Works that You distribute, alongside
|
||||||
|
or as an addendum to the NOTICE text from the Work, provided
|
||||||
|
that such additional attribution notices cannot be construed
|
||||||
|
as modifying the License.
|
||||||
|
|
||||||
|
You may add Your own copyright statement to Your modifications and
|
||||||
|
may provide additional or different license terms and conditions
|
||||||
|
for use, reproduction, or distribution of Your modifications, or
|
||||||
|
for any such Derivative Works as a whole, provided Your use,
|
||||||
|
reproduction, and distribution of the Work otherwise complies with
|
||||||
|
the conditions stated in this License.
|
||||||
|
|
||||||
|
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||||
|
any Contribution intentionally submitted for inclusion in the Work
|
||||||
|
by You to the Licensor shall be under the terms and conditions of
|
||||||
|
this License, without any additional terms or conditions.
|
||||||
|
Notwithstanding the above, nothing herein shall supersede or modify
|
||||||
|
the terms of any separate license agreement you may have executed
|
||||||
|
with Licensor regarding such Contributions.
|
||||||
|
|
||||||
|
6. Trademarks. This License does not grant permission to use the trade
|
||||||
|
names, trademarks, service marks, or product names of the Licensor,
|
||||||
|
except as required for reasonable and customary use in describing the
|
||||||
|
origin of the Work and reproducing the content of the NOTICE file.
|
||||||
|
|
||||||
|
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||||
|
agreed to in writing, Licensor provides the Work (and each
|
||||||
|
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||||
|
implied, including, without limitation, any warranties or conditions
|
||||||
|
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||||
|
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||||
|
appropriateness of using or redistributing the Work and assume any
|
||||||
|
risks associated with Your exercise of permissions under this License.
|
||||||
|
|
||||||
|
8. Limitation of Liability. In no event and under no legal theory,
|
||||||
|
whether in tort (including negligence), contract, or otherwise,
|
||||||
|
unless required by applicable law (such as deliberate and grossly
|
||||||
|
negligent acts) or agreed to in writing, shall any Contributor be
|
||||||
|
liable to You for damages, including any direct, indirect, special,
|
||||||
|
incidental, or consequential damages of any character arising as a
|
||||||
|
result of this License or out of the use or inability to use the
|
||||||
|
Work (including but not limited to damages for loss of goodwill,
|
||||||
|
work stoppage, computer failure or malfunction, or any and all
|
||||||
|
other commercial damages or losses), even if such Contributor
|
||||||
|
has been advised of the possibility of such damages.
|
||||||
|
|
||||||
|
9. Accepting Warranty or Additional Liability. While redistributing
|
||||||
|
the Work or Derivative Works thereof, You may choose to offer,
|
||||||
|
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||||
|
or other liability obligations and/or rights consistent with this
|
||||||
|
License. However, in accepting such obligations, You may act only
|
||||||
|
on Your own behalf and on Your sole responsibility, not on behalf
|
||||||
|
of any other Contributor, and only if You agree to indemnify,
|
||||||
|
defend, and hold each Contributor harmless for any liability
|
||||||
|
incurred by, or claims asserted against, such Contributor by reason
|
||||||
|
of your accepting any such warranty or additional liability.
|
||||||
|
|
||||||
|
END OF TERMS AND CONDITIONS
|
||||||
|
|
||||||
|
APPENDIX: How to apply the Apache License to your work.
|
||||||
|
|
||||||
|
To apply the Apache License to your work, attach the following
|
||||||
|
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||||
|
replaced with your own identifying information. (Don't include
|
||||||
|
the brackets!) The text should be enclosed in the appropriate
|
||||||
|
comment syntax for the file format. We also recommend that a
|
||||||
|
file or class name and description of purpose be included on the
|
||||||
|
same "printed page" as the copyright notice for easier
|
||||||
|
identification within third-party archives.
|
||||||
|
|
||||||
|
Copyright [yyyy] [name of copyright owner]
|
||||||
|
|
||||||
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
you may not use this file except in compliance with the License.
|
||||||
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
|
Unless required by applicable law or agreed to in writing, software
|
||||||
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
See the License for the specific language governing permissions and
|
||||||
|
limitations under the License.
|
||||||
+19
@@ -0,0 +1,19 @@
|
|||||||
|
Copyright (c) 2026 Anders Olsson
|
||||||
|
|
||||||
|
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||||
|
of this software and associated documentation files (the "Software"), to deal
|
||||||
|
in the Software without restriction, including without limitation the rights
|
||||||
|
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||||
|
copies of the Software, and to permit persons to whom the Software is
|
||||||
|
furnished to do so, subject to the following conditions:
|
||||||
|
|
||||||
|
The above copyright notice and this permission notice shall be included in
|
||||||
|
all copies or substantial portions of the Software.
|
||||||
|
|
||||||
|
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||||
|
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||||
|
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||||
|
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||||
|
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||||
|
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
|
||||||
|
THE SOFTWARE.
|
||||||
@@ -13,64 +13,136 @@ Rust port of [TrueSkillThroughTime.py](https://github.com/glandfried/TrueSkillTh
|
|||||||
|
|
||||||
## Drift
|
## Drift
|
||||||
|
|
||||||
Skill drift models how a player's true skill can change between appearances. Each time a player reappears after a gap, their skill uncertainty is widened by the drift model before the new evidence is incorporated.
|
Skill drift models how a competitor's true skill can change between appearances.
|
||||||
|
Each time they reappear after a gap, their skill uncertainty is widened by the
|
||||||
|
drift model before the new evidence is incorporated.
|
||||||
|
|
||||||
Drift is represented by the `Drift` trait:
|
Drift is represented by the `Drift` trait (`src/drift.rs`), generic over the
|
||||||
|
history's time type:
|
||||||
|
|
||||||
```rust
|
```text
|
||||||
pub trait Drift: Copy + Debug {
|
pub trait Drift<T: Time>: Copy + Debug + Send + Sync {
|
||||||
fn variance_delta(&self, elapsed: i64) -> f64;
|
fn variance_delta(&self, from: &T, to: &T) -> f64;
|
||||||
|
fn variance_for_elapsed(&self, elapsed: i64) -> f64;
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
`variance_delta` returns the amount to add to `σ²` given the elapsed time since the player last played. Internally, `Gaussian::forget` uses this to compute the new sigma: `σ_new = sqrt(σ² + variance_delta)`.
|
Both methods return the amount to add to `σ²`, not to `σ`. `variance_delta`
|
||||||
|
works from two timestamps; `variance_for_elapsed` takes an already-computed
|
||||||
|
elapsed count, and is used on the paths that cache it. `Gaussian::forget`
|
||||||
|
applies the result entirely in variance space — `from_mv(mu, variance() +
|
||||||
|
variance_delta)` — taking no square root.
|
||||||
|
|
||||||
|
That block is a quotation rather than a doctest. The custom-drift example below
|
||||||
|
is compiled by CI, so it is what actually pins the signature.
|
||||||
|
|
||||||
### ConstantDrift
|
### ConstantDrift
|
||||||
|
|
||||||
The built-in `ConstantDrift` implements a linear random walk — skill uncertainty grows proportionally to time:
|
The built-in `ConstantDrift` implements a linear random walk — skill uncertainty
|
||||||
|
grows proportionally to time:
|
||||||
|
|
||||||
```
|
```text
|
||||||
variance_delta = elapsed * γ²
|
variance_delta = elapsed * γ²
|
||||||
```
|
```
|
||||||
|
|
||||||
This is the standard TrueSkill Through Time model. Use it by passing a `ConstantDrift(gamma)` when constructing a `Player`:
|
This is the standard TrueSkill Through Time model. Pass a `ConstantDrift(gamma)`
|
||||||
|
when constructing a `Rating`:
|
||||||
|
|
||||||
```rust
|
```rust
|
||||||
use trueskill_tt::{Player, Gaussian, drift::ConstantDrift};
|
use trueskill_tt::{ConstantDrift, Gaussian, Rating};
|
||||||
|
|
||||||
// gamma = 0.1 means skill can shift ~0.1 per time unit
|
// gamma = 0.1 means skill can shift ~0.1 per time unit.
|
||||||
let player = Player::new(Gaussian::from_ms(0.0, 6.0), 1.0, ConstantDrift(0.1));
|
let rating: Rating<i64, ConstantDrift> =
|
||||||
|
Rating::new(Gaussian::from_ms(0.0, 6.0), 1.0, ConstantDrift(0.1));
|
||||||
|
|
||||||
|
assert_eq!(rating.drift().0, 0.1);
|
||||||
```
|
```
|
||||||
|
|
||||||
|
The type annotation is load-bearing: `ConstantDrift` implements `Drift<T>` for
|
||||||
|
every `T: Time`, so without it `T` is ambiguous.
|
||||||
|
|
||||||
### Custom drift
|
### Custom drift
|
||||||
|
|
||||||
Implement `Drift` to express any other model. For example, a drift that saturates after a long absence (uncertainty grows with the square root of elapsed time instead of linearly):
|
Implement `Drift<T>` to express any other model. For example, a drift that
|
||||||
|
saturates after a long absence, with uncertainty growing as the square root of
|
||||||
|
elapsed time instead of linearly:
|
||||||
|
|
||||||
```rust
|
```rust
|
||||||
use trueskill_tt::drift::Drift;
|
use trueskill_tt::{Drift, Gaussian, History, Rating, Time};
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug)]
|
#[derive(Clone, Copy, Debug)]
|
||||||
struct SqrtDrift {
|
struct SqrtDrift {
|
||||||
gamma: f64,
|
gamma: f64,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Drift for SqrtDrift {
|
impl<T: Time> Drift<T> for SqrtDrift {
|
||||||
fn variance_delta(&self, elapsed: i64) -> f64 {
|
fn variance_delta(&self, from: &T, to: &T) -> f64 {
|
||||||
(elapsed as f64).sqrt() * self.gamma * self.gamma
|
let elapsed = from.elapsed_to(to).max(0) as f64;
|
||||||
|
elapsed.sqrt() * self.gamma * self.gamma
|
||||||
|
}
|
||||||
|
|
||||||
|
fn variance_for_elapsed(&self, elapsed: i64) -> f64 {
|
||||||
|
(elapsed.max(0) as f64).sqrt() * self.gamma * self.gamma
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let player = Player::new(Gaussian::from_ms(0.0, 6.0), 1.0, SqrtDrift { gamma: 0.5 });
|
// On a single Rating:
|
||||||
|
let rating: Rating<i64, SqrtDrift> =
|
||||||
|
Rating::new(Gaussian::from_ms(0.0, 6.0), 1.0, SqrtDrift { gamma: 0.5 });
|
||||||
|
|
||||||
|
// Or for a whole History, via the builder:
|
||||||
|
let history = History::builder().drift(SqrtDrift { gamma: 0.5 }).build();
|
||||||
|
|
||||||
|
assert_eq!(rating.beta(), 1.0);
|
||||||
|
assert_eq!(history.log_evidence(), 0.0);
|
||||||
```
|
```
|
||||||
|
|
||||||
To use a custom drift type with `History`, use the `.drift()` builder method instead of `.gamma()`:
|
`HistoryBuilder::drift` is the only way to set a history's drift model; there is
|
||||||
|
no `gamma()` shorthand. The default is `ConstantDrift(GAMMA)`.
|
||||||
|
|
||||||
|
### Per-competitor drift
|
||||||
|
|
||||||
|
A `History` has one drift model, but individual competitors can scale it.
|
||||||
|
`Member::with_drift_scale(s)` multiplies the drift *variance* that competitor
|
||||||
|
accumulates, so `s` is in the same units as `gamma`: `ConstantDrift(g)` at
|
||||||
|
scale `s` behaves exactly as `ConstantDrift(g * s)` would, for that competitor
|
||||||
|
alone.
|
||||||
|
|
||||||
|
`0.0` pins a competitor still. That is what makes a **fixed reference point**
|
||||||
|
expressible in the same graph as moving competitors — a bot at a known
|
||||||
|
strength, a rating floor, a course difficulty:
|
||||||
|
|
||||||
```rust
|
```rust
|
||||||
let h = History::builder()
|
use trueskill_tt::{ConstantDrift, Event, History, Member, Outcome, Team};
|
||||||
.drift(SqrtDrift { gamma: 0.5 })
|
|
||||||
.build();
|
let mut h = History::builder().drift(ConstantDrift(0.1)).build();
|
||||||
|
|
||||||
|
h.add_events(vec![Event {
|
||||||
|
time: 0,
|
||||||
|
teams: [
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
// A course does not improve. Pin it, and the round's evidence
|
||||||
|
// lands on the player instead of being split between the two.
|
||||||
|
Team::with_members([Member::new("layout_7").with_drift_scale(0.0)]),
|
||||||
|
]
|
||||||
|
.into_iter()
|
||||||
|
.collect(),
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}])
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Like `with_prior`, the scale is **competitor configuration captured at first
|
||||||
|
appearance** — setting it on a key the history already knows has no effect. It
|
||||||
|
must be finite and non-negative; ingestion otherwise fails with
|
||||||
|
`InferenceError::InvalidParameter`.
|
||||||
|
|
||||||
|
Note that the fluent `EventBuilder` (`h.event(t).team([...])`) sets weights but
|
||||||
|
not `drift_scale` or `prior`; those need the typed `Event` / `Team` / `Member`
|
||||||
|
shape shown above.
|
||||||
|
|
||||||
## Scored outcomes
|
## Scored outcomes
|
||||||
|
|
||||||
Use `Outcome::scores([...])` when you have continuous per-team scores rather
|
Use `Outcome::scores([...])` when you have continuous per-team scores rather
|
||||||
@@ -80,7 +152,7 @@ soft Gaussian evidence about the latent performance diff. Configure
|
|||||||
(smaller σ = more trust).
|
(smaller σ = more trust).
|
||||||
|
|
||||||
```rust
|
```rust
|
||||||
use trueskill_tt::{History, Outcome};
|
use trueskill_tt::History;
|
||||||
|
|
||||||
let mut h = History::builder().score_sigma(2.0).build();
|
let mut h = History::builder().score_sigma(2.0).build();
|
||||||
h.event(1)
|
h.event(1)
|
||||||
@@ -92,12 +164,99 @@ h.event(1)
|
|||||||
h.converge().unwrap();
|
h.converge().unwrap();
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Prediction
|
||||||
|
|
||||||
|
`predict_outcome` gives the full distribution over finishing orders. Each entry
|
||||||
|
is a rank vector in the same shape `Outcome::ranking` takes — equal ranks mean a
|
||||||
|
tie — so an outcome feeds straight back into inference.
|
||||||
|
|
||||||
|
```rust
|
||||||
|
use trueskill_tt::History;
|
||||||
|
|
||||||
|
let mut h = History::builder().p_draw(0.1).build();
|
||||||
|
h.record_winner(&"alice", &"bob", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let p = h.predict_outcome(&[&[&"alice"], &[&"bob"]]).unwrap();
|
||||||
|
|
||||||
|
// Probabilities are exhaustive and disjoint, so they sum to one.
|
||||||
|
assert!((p.total() - 1.0).abs() < 1e-6);
|
||||||
|
|
||||||
|
let (best, likelihood) = p.most_likely().unwrap();
|
||||||
|
println!("most likely: {best:?} at {likelihood:.3}");
|
||||||
|
println!("draw: {:.3}", p.probability_of(&[0, 0]));
|
||||||
|
```
|
||||||
|
|
||||||
|
Supports any number of teams. Because the outcome space grows factorially, the
|
||||||
|
full distribution is capped at `MAX_PREDICTED_TEAMS`; two cheaper entry points
|
||||||
|
stay available at any size:
|
||||||
|
|
||||||
|
- `predict_win_probabilities(teams)` — `P(team i finishes strictly first)`,
|
||||||
|
quadratic in team count.
|
||||||
|
- `predict_ranking(teams, ranks)` — one specific finishing order.
|
||||||
|
|
||||||
|
Unknown keys are an error, not a silent omission: a team the history has never
|
||||||
|
seen cannot produce a confident-looking probability.
|
||||||
|
|
||||||
|
## Which match to play next
|
||||||
|
|
||||||
|
`quality()` measures whether a matchup is *fair*. That is not the same as
|
||||||
|
whether it is *informative*, and the two only coincide for two evenly matched
|
||||||
|
competitors. When each observation costs something, ask
|
||||||
|
`expected_information_gain` instead — the outcome-weighted divergence between
|
||||||
|
what you believe now and what you would believe afterwards.
|
||||||
|
|
||||||
|
```rust
|
||||||
|
use trueskill_tt::History;
|
||||||
|
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
for t in 1..=10 {
|
||||||
|
h.record_winner(&"veteran", &"regular", t).unwrap();
|
||||||
|
h.record_winner(&"regular", &"veteran", t + 100).unwrap();
|
||||||
|
}
|
||||||
|
h.record_winner(&"veteran", &"newcomer", 500).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let settled = h.expected_information_gain(&[&[&"veteran"], &[&"regular"]]).unwrap();
|
||||||
|
let unknown = h.expected_information_gain(&[&[&"veteran"], &[&"newcomer"]]).unwrap();
|
||||||
|
|
||||||
|
// Playing the newcomer teaches you more than replaying a settled rivalry.
|
||||||
|
assert!(unknown > settled);
|
||||||
|
```
|
||||||
|
|
||||||
|
The result is in nats, and is bounded by the entropy of the outcome: at most
|
||||||
|
`ln 2 ≈ 0.693` for a two-way result, `ln 3` once draws are possible, `ln k` for
|
||||||
|
`k` outcomes. A value near zero means you already know how it ends.
|
||||||
|
|
||||||
|
This costs one full inference pass **per possible outcome**, so it is far more
|
||||||
|
expensive than `quality()`. Scoring every pairing among `n` competitors is
|
||||||
|
`O(n² × outcomes)` passes — shortlist with `quality()` or
|
||||||
|
`predict_win_probabilities` first, then score only the shortlist.
|
||||||
|
|
||||||
## Todo
|
## Todo
|
||||||
|
|
||||||
- [x] Implement approx for Gaussian
|
- [x] Implement approx for Gaussian
|
||||||
- [x] Add more tests from `TrueSkillThroughTime.jl`
|
- [x] Add more tests from `TrueSkillThroughTime.jl`
|
||||||
- [ ] Add tests for `quality()` (Use [sublee/trueskill](https://github.com/sublee/trueskill/tree/master) as reference)
|
- [x] Generalise a time axis — `Time` is now a trait (`Untimed`, `i64`), not an enum
|
||||||
- [ ] Benchmark Batch::iteration()
|
- [x] Add examples (`examples/atp.rs`, `examples/scored.rs`)
|
||||||
- [ ] Time needs to be an enum so we can have multiple states (see `batch::compute_elapsed()`)
|
- [x] Add Observer (`Observer` / `NullObserver`)
|
||||||
- [ ] Add examples (use same TrueSkillThroughTime.(py|jl))
|
- [x] Benchmark the inference loop (`benches/batch.rs`, `benches/history_converge.rs`, `benches/ingest.rs`)
|
||||||
- [ ] Add Observer (see [argmin](https://docs.rs/argmin/latest/argmin/core/trait.Observe.html) for inspiration)
|
- [x] N-team `predict_outcome` with draw mass, and `expected_information_gain`
|
||||||
|
- [x] Cross-check `quality()` against [sublee/trueskill](https://github.com/sublee/trueskill/tree/master) — N identical teams follow the closed form `(1/5)^((n-1)/2)` for the conventional parameters, asserted for n = 2..10, and the n=3/n=5 values (0.200, 0.040) match the reference package
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
Licensed under either of
|
||||||
|
|
||||||
|
- Apache License, Version 2.0 ([LICENSE-APACHE](LICENSE-APACHE) or
|
||||||
|
<http://www.apache.org/licenses/LICENSE-2.0>)
|
||||||
|
- MIT license ([LICENSE-MIT](LICENSE-MIT) or
|
||||||
|
<http://opensource.org/licenses/MIT>)
|
||||||
|
|
||||||
|
at your option.
|
||||||
|
|
||||||
|
### Contribution
|
||||||
|
|
||||||
|
Unless you explicitly state otherwise, any contribution intentionally submitted
|
||||||
|
for inclusion in the work by you, as defined in the Apache-2.0 license, shall be
|
||||||
|
dual licensed as above, without any additional terms or conditions.
|
||||||
|
|||||||
+1
-1
@@ -36,7 +36,7 @@ fn criterion_benchmark(criterion: &mut Criterion) {
|
|||||||
let kinds = vec![EventKind::Ranked; composition.len()];
|
let kinds = vec![EventKind::Ranked; composition.len()];
|
||||||
|
|
||||||
let mut time_slice = TimeSlice::new(1, P_DRAW, ConvergenceOptions::default());
|
let mut time_slice = TimeSlice::new(1, P_DRAW, ConvergenceOptions::default());
|
||||||
time_slice.add_events(composition, results, weights, kinds, &agents);
|
time_slice.add_events(composition, Some(results), Some(weights), kinds, &agents);
|
||||||
|
|
||||||
criterion.bench_function("Batch::iteration", |b| {
|
criterion.bench_function("Batch::iteration", |b| {
|
||||||
b.iter(|| time_slice.iteration(0, &agents))
|
b.iter(|| time_slice.iteration(0, &agents))
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
//! Ingestion cost: one event per call versus one batched call.
|
||||||
|
//!
|
||||||
|
//! The rest of the suite only measures batched construction, which is why a
|
||||||
|
//! quadratic in the incremental path went unnoticed — `record_winner` and
|
||||||
|
//! `event(..).commit()` each ingest a single event, so a caller looping over a
|
||||||
|
//! match feed takes that path.
|
||||||
|
|
||||||
|
use std::hint::black_box;
|
||||||
|
|
||||||
|
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{Event, History, Member, Outcome, Team};
|
||||||
|
|
||||||
|
fn events(n: usize, time: i64) -> Vec<Event<i64, String>> {
|
||||||
|
(0..n)
|
||||||
|
.map(|i| Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(format!("p{}", 2 * i))]),
|
||||||
|
Team::with_members([Member::new(format!("p{}", 2 * i + 1))]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
})
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn bench_ingest(c: &mut Criterion) {
|
||||||
|
let mut group = c.benchmark_group("ingest");
|
||||||
|
|
||||||
|
for n in [250usize, 500, 1000] {
|
||||||
|
group.bench_with_input(BenchmarkId::new("one-at-a-time", n), &n, |b, &n| {
|
||||||
|
b.iter_batched(
|
||||||
|
|| events(n, 0),
|
||||||
|
|evs| {
|
||||||
|
let mut h: History<i64, _, _, String> = History::builder_with_key().build();
|
||||||
|
for ev in evs {
|
||||||
|
h.add_events(std::iter::once(ev)).unwrap();
|
||||||
|
}
|
||||||
|
black_box(h.time_slices_len())
|
||||||
|
},
|
||||||
|
criterion::BatchSize::SmallInput,
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
group.bench_with_input(BenchmarkId::new("single-batch", n), &n, |b, &n| {
|
||||||
|
b.iter_batched(
|
||||||
|
|| events(n, 0),
|
||||||
|
|evs| {
|
||||||
|
let mut h: History<i64, _, _, String> = History::builder_with_key().build();
|
||||||
|
h.add_events(evs).unwrap();
|
||||||
|
black_box(h.time_slices_len())
|
||||||
|
},
|
||||||
|
criterion::BatchSize::SmallInput,
|
||||||
|
);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
group.finish();
|
||||||
|
}
|
||||||
|
|
||||||
|
criterion_group!(benches, bench_ingest);
|
||||||
|
criterion_main!(benches);
|
||||||
@@ -44,6 +44,11 @@ split_commits = false
|
|||||||
# Assigns commits to groups.
|
# Assigns commits to groups.
|
||||||
# Optionally sets the commit's scope and can decide to exclude commits from further processing.
|
# Optionally sets the commit's scope and can decide to exclude commits from further processing.
|
||||||
commit_parsers = [
|
commit_parsers = [
|
||||||
|
# Must precede the type parsers below: a `feat!`/`fix!`/`refactor!` subject
|
||||||
|
# matches those too, and the first match wins. Without this a breaking
|
||||||
|
# change renders as an ordinary line of its own type.
|
||||||
|
{ message = "^[a-z]+(\\(.+\\))?!:", group = "Breaking Changes" },
|
||||||
|
{ body = "BREAKING CHANGE", group = "Breaking Changes" },
|
||||||
{ message = "^feat", group = "Features" },
|
{ message = "^feat", group = "Features" },
|
||||||
{ message = "^fix", group = "Bug Fixes" },
|
{ message = "^fix", group = "Bug Fixes" },
|
||||||
{ message = "^doc", group = "Documentation" },
|
{ message = "^doc", group = "Documentation" },
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,342 @@
|
|||||||
|
# Filtered (Forward-Only) Estimates
|
||||||
|
|
||||||
|
Closes [#19](https://git.aceofba.se/logaritmisk/trueskill-tt/issues/19).
|
||||||
|
|
||||||
|
## Summary
|
||||||
|
|
||||||
|
`HistoryBuilder::online(true)` is inert. It flips a flag that reaches
|
||||||
|
`Item::within_prior` (`src/time_slice.rs:70-71`), which reads
|
||||||
|
`Skill.online` (`src/time_slice.rs:25`) — a field initialised to `N_INF`
|
||||||
|
(`src/time_slice.rs:41`) and never assigned anywhere. The online path
|
||||||
|
therefore builds every rating from the improper Gaussian, and
|
||||||
|
`log_evidence()` silently reports `n × ln(0.5)`: every game scored as a
|
||||||
|
coin flip, finite and plausible-looking.
|
||||||
|
|
||||||
|
This spec replaces the field and the flag with a **read-only forward-only
|
||||||
|
pass** over the converged history, exposed as three new public methods.
|
||||||
|
The pass reuses the production within-slice sweep verbatim rather than
|
||||||
|
reimplementing inference, and stores nothing on `Skill`.
|
||||||
|
|
||||||
|
## Background
|
||||||
|
|
||||||
|
### Why a stored field cannot hold this quantity
|
||||||
|
|
||||||
|
The issue proposes populating `skill.online` during the forward pass,
|
||||||
|
alongside `new_forward_info` (`src/time_slice.rs:576`). That would not
|
||||||
|
work, and understanding why determines the whole design.
|
||||||
|
|
||||||
|
`new_forward_info` sets `skill.forward` from
|
||||||
|
`agents[a].receive_for_elapsed(...)`, whose `message` was written by the
|
||||||
|
previous slice's `forward_prior_out` (`src/time_slice.rs:549`):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
skill.forward * skill.likelihood
|
||||||
|
```
|
||||||
|
|
||||||
|
`History::iteration` (`src/history.rs:255`) alternates a backward sweep
|
||||||
|
over slices and a forward sweep. From the second iteration onward, the
|
||||||
|
`skill.likelihood` feeding that message has already absorbed backward
|
||||||
|
information from the preceding backward sweep. So after `converge()`,
|
||||||
|
**`skill.forward` is a smoothed quantity, not a filtering one** — and any
|
||||||
|
field written from it inherits the same contamination on every sweep
|
||||||
|
after the first.
|
||||||
|
|
||||||
|
### The neighbouring trap
|
||||||
|
|
||||||
|
The same reasoning applies to the existing `forward: bool` parameter on
|
||||||
|
`log_evidence_internal` (`src/history.rs:395`). It is a genuine filtering
|
||||||
|
quantity only on a history that has never been converged. That is why the
|
||||||
|
test at `src/history.rs:1183` can assert
|
||||||
|
|
||||||
|
```rust
|
||||||
|
assert_ulps_eq!(trueskill_log_evidence, trueskill_log_evidence_online, epsilon = 1e-6);
|
||||||
|
```
|
||||||
|
|
||||||
|
— the fixture is never converged, so the forward message still equals the
|
||||||
|
cavity prior. (Note also that the local binding is named `..._online`
|
||||||
|
while the flag it passes is `forward`; the two senses were already
|
||||||
|
muddled.)
|
||||||
|
|
||||||
|
Fixing `forward: bool` is **out of scope** here; see *Out-of-scope
|
||||||
|
follow-ups*.
|
||||||
|
|
||||||
|
### Why this is worth implementing rather than deleting
|
||||||
|
|
||||||
|
The forward-only estimate has a second consumer beyond prequential model
|
||||||
|
comparison. `learning_curve()` returns post-convergence posteriors, so
|
||||||
|
every point is smoothed — the estimate at a given date incorporates
|
||||||
|
rounds played years later. On [ustat](https://git.aceofba.se/logaritmisk/ustat)'s
|
||||||
|
real data (prior μ=0, σ=6) that produces curves which start already
|
||||||
|
spread apart and barely move:
|
||||||
|
|
||||||
|
```
|
||||||
|
player first point final point
|
||||||
|
Eskil mu +3.72 sigma 1.17 mu +4.61 sigma 1.21
|
||||||
|
Anders Olsson mu +1.61 sigma 0.90 mu +1.16 sigma 0.82
|
||||||
|
LUDVIGSSON mu -2.09 sigma 1.08 mu -2.61 sigma 1.13
|
||||||
|
Anners mu -2.85 sigma 1.27 mu -2.86 sigma 1.26
|
||||||
|
```
|
||||||
|
|
||||||
|
σ at the *first* plotted point is 0.90–1.60 against a prior of 6.00. A
|
||||||
|
caller cannot reconstruct the filtered view from the public API today
|
||||||
|
except by refitting over `events[0..k]` for every k — O(n²) fits for
|
||||||
|
something one forward pass already computes.
|
||||||
|
|
||||||
|
## Scope
|
||||||
|
|
||||||
|
### What ships
|
||||||
|
|
||||||
|
1. A read-only forward-only pass on `History`, walking slices in time
|
||||||
|
order and carrying its own forward messages.
|
||||||
|
2. Three public methods: `filtered_log_evidence`,
|
||||||
|
`filtered_learning_curves`, `filtered_learning_curve`.
|
||||||
|
3. Removal of `Skill.online`, `History.online`, `HistoryBuilder.online`,
|
||||||
|
`HistoryBuilder::online()`, and the `online: bool` parameter threaded
|
||||||
|
through `Item::within_prior`, `Event::within_priors`, and
|
||||||
|
`TimeSlice::log_evidence`.
|
||||||
|
4. `#[derive(Clone)]` on `Event`, `Team`, `Item`; `iterate_to_convergence`
|
||||||
|
loses its `#[cfg(test)]` gate.
|
||||||
|
5. A CHANGELOG entry recording the API break.
|
||||||
|
|
||||||
|
### What does not ship
|
||||||
|
|
||||||
|
- No change to `log_evidence()`, `log_evidence_for()`, `learning_curve()`,
|
||||||
|
`learning_curves()`, or `current_skill()`. Their values are unchanged
|
||||||
|
by this work.
|
||||||
|
- No fix to the `forward: bool` flag described above.
|
||||||
|
- No caching of pass results. Each call runs a full pass; the doc
|
||||||
|
comments say so.
|
||||||
|
- No `rayon` parallelism across slices — the pass is sequentially
|
||||||
|
dependent by construction.
|
||||||
|
- No prior-predictive accessor. The pass computes the pre-event forward
|
||||||
|
message internally, but only the filtered posterior is exposed until a
|
||||||
|
second caller needs otherwise.
|
||||||
|
|
||||||
|
## Design
|
||||||
|
|
||||||
|
### Naming
|
||||||
|
|
||||||
|
`filtered_*`, not `online_*`. "Filtered" is the standard term for the
|
||||||
|
forward-only estimate, and the crate already uses "online" for a second,
|
||||||
|
unrelated thing — incremental ingestion, which `benches/baseline.txt:128`
|
||||||
|
calls the "online-add" path. Two senses of one word in one crate is how
|
||||||
|
the present bug reads as plausible.
|
||||||
|
|
||||||
|
### The pass
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub(crate) struct FilteredStep {
|
||||||
|
log_evidence: f64,
|
||||||
|
posteriors: Vec<(Index, Gaussian)>,
|
||||||
|
}
|
||||||
|
|
||||||
|
fn filtered_pass(&self) -> Vec<(T, FilteredStep)>
|
||||||
|
```
|
||||||
|
|
||||||
|
`posteriors` doubles as the outgoing forward message: the scratch sweep never
|
||||||
|
writes `backward`, so it stays `N_INF`, and `Skill::posterior()` and
|
||||||
|
`forward_prior_out` are then the same product.
|
||||||
|
|
||||||
|
Walk `self.time_slices` in order, carrying
|
||||||
|
`messages: HashMap<Index, Gaussian>` — the forward message out of each
|
||||||
|
competitor's most recent appearance. For each slice:
|
||||||
|
|
||||||
|
1. **Build a scratch clone.** Same `time`, `p_draw`, `convergence`, and
|
||||||
|
cloned `events` with every `item.likelihood` reset to `N_INF`. Fresh
|
||||||
|
`SkillStore` in which, for each agent present in the real slice:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
forward = match messages.get(&agent) {
|
||||||
|
Some(msg) => msg.forget(rating.drift.variance_for_elapsed(skill.elapsed)),
|
||||||
|
None => rating.prior,
|
||||||
|
}
|
||||||
|
backward = N_INF
|
||||||
|
likelihood = N_INF
|
||||||
|
elapsed = skill.elapsed // copied from the real slice
|
||||||
|
```
|
||||||
|
|
||||||
|
This mirrors `Competitor::receive_for_elapsed` (`src/competitor.rs:39`)
|
||||||
|
exactly, including its `message != N_INF` fallback to the prior.
|
||||||
|
`skill.elapsed` is reused rather than recomputed: it is maintained by
|
||||||
|
`add_events_with_prior` across out-of-order ingestion, and production
|
||||||
|
convergence already trusts it.
|
||||||
|
|
||||||
|
2. **Run the real sweep.** `scratch.iterate_to_convergence(agents)`
|
||||||
|
(`src/time_slice.rs:516`), unmodified. Fidelity comes from reusing the
|
||||||
|
production path rather than a parallel reimplementation — in
|
||||||
|
particular, a competitor appearing in two events at the same time is
|
||||||
|
handled by the same within-slice EP that `converge()` uses, not
|
||||||
|
approximated the way the current `online`/`forward` evidence paths are
|
||||||
|
(they run each event independently and sum).
|
||||||
|
|
||||||
|
3. **Harvest.** With `backward == N_INF` acting as the multiplicative
|
||||||
|
identity, `Skill::posterior()` is exactly forward × likelihood — the
|
||||||
|
filtered posterior. Slice evidence is
|
||||||
|
`scratch.events.iter().map(|e| e.log_evidence).sum()`; `apply`
|
||||||
|
(`src/time_slice.rs:162`) writes that field on every event during the
|
||||||
|
sweep.
|
||||||
|
|
||||||
|
4. **Carry forward.** `messages.insert(a, scratch.forward_prior_out(&a))`
|
||||||
|
for each agent in the slice.
|
||||||
|
|
||||||
|
Steps 1–4 are the forward half of `History::iteration`
|
||||||
|
(`src/history.rs:283-297`) with the backward half never run. The pass
|
||||||
|
touches no field of `self`.
|
||||||
|
|
||||||
|
### Public API
|
||||||
|
|
||||||
|
```rust
|
||||||
|
impl<T: Time, D: Drift<T>, O: Observer<T>, K: Eq + Hash + Clone> History<T, D, O, K> {
|
||||||
|
pub fn filtered_log_evidence(&self) -> f64;
|
||||||
|
pub fn filtered_learning_curves(&self) -> HashMap<K, Vec<(T, Gaussian)>>;
|
||||||
|
pub fn filtered_learning_curve<Q>(&self, key: &Q) -> Vec<(T, Gaussian)>
|
||||||
|
where
|
||||||
|
K: Borrow<Q>,
|
||||||
|
Q: Hash + Eq + ?Sized;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
All take `&self` — the pass mutates nothing. Shapes deliberately mirror
|
||||||
|
`learning_curve` / `learning_curves` (`src/history.rs:325`, `:381`) so a
|
||||||
|
caller can plot smoothed and filtered curves on one chart with the same
|
||||||
|
handling code.
|
||||||
|
|
||||||
|
`filtered_learning_curve` runs the same full pass as the plural form and
|
||||||
|
collects one key; the cost is identical, only the collection differs.
|
||||||
|
Callers wanting several keys should use the plural form. Documented on
|
||||||
|
both methods.
|
||||||
|
|
||||||
|
Because the pass carries its own messages and re-runs inference, its
|
||||||
|
results **do not depend on whether `converge()` has been called**. That
|
||||||
|
is the property a stored field cannot have, and it is asserted as a test.
|
||||||
|
|
||||||
|
### Removal inventory
|
||||||
|
|
||||||
|
| Location | Change |
|
||||||
|
|---|---|
|
||||||
|
| `src/time_slice.rs:25` | delete `pub(crate) online: Gaussian` |
|
||||||
|
| `src/time_slice.rs:41` | delete `online: N_INF` from `Default` |
|
||||||
|
| `src/time_slice.rs:62,70-73` | drop `online` param and its branch from `Item::within_prior` |
|
||||||
|
| `src/time_slice.rs:110,120` | drop `online` param from `Event::within_priors` |
|
||||||
|
| `src/time_slice.rs:585,597,626,634` | drop `online` param from `TimeSlice::log_evidence`; `online \|\| forward` becomes `forward` |
|
||||||
|
| `src/history.rs:32,63,138,158,174,199,226` | delete the two `online` field declarations (`:32`, `:199`) and the five struct-literal copies |
|
||||||
|
| `src/history.rs:90-93` | delete `HistoryBuilder::online()` |
|
||||||
|
| `src/history.rs:402,410` | drop the `self.online` argument |
|
||||||
|
| `src/history.rs:1183-1189` | the `..._online` assertion becomes a `forward`-flag assertion; rename the binding to match what it tests |
|
||||||
|
|
||||||
|
`Skill` loses 16 bytes, which is a small independent win for #17.
|
||||||
|
|
||||||
|
## Testing strategy
|
||||||
|
|
||||||
|
Every new test is mutation-proved before it counts: break the production
|
||||||
|
line it names, watch it fail for the *right* assertion, restore. A test
|
||||||
|
never observed failing is not evidence.
|
||||||
|
|
||||||
|
### The red test
|
||||||
|
|
||||||
|
On the issue's own fixture — five 1v1 games, same winner each time —
|
||||||
|
`filtered_log_evidence()` must land strictly between the two known
|
||||||
|
endpoints:
|
||||||
|
|
||||||
|
```
|
||||||
|
5 × ln(0.5) = -3.4657... (today's inert value)
|
||||||
|
< filtered
|
||||||
|
< -0.4012... (batch / smoothed evidence)
|
||||||
|
```
|
||||||
|
|
||||||
|
Two-sided, so neither "still inert" nor "accidentally smoothed" can pass.
|
||||||
|
The lower bound is right for a real reason: game one genuinely *is* a
|
||||||
|
coin flip under filtering, games two through five are not.
|
||||||
|
|
||||||
|
### Invariants
|
||||||
|
|
||||||
|
1. **Invariant to `converge()`** — `filtered_log_evidence()` and
|
||||||
|
`filtered_learning_curves()` agree before and after `converge()`. This
|
||||||
|
is exactly what `skill.forward` fails, and what makes a stored field
|
||||||
|
the wrong mechanism.
|
||||||
|
|
||||||
|
Agreement is to tolerance, not bit-identity, and the reason is worth
|
||||||
|
recording. `iteration` calls `recompute_color_groups`
|
||||||
|
(`src/time_slice.rs:369`) only when `from == 0`, so a slice built by
|
||||||
|
repeated appends keeps insertion order until the first `converge()`
|
||||||
|
reorders it. The scratch clone inherits whichever order it finds, and
|
||||||
|
greedy coloring over a permuted input can group differently, giving a
|
||||||
|
different within-slice sweep order — same EP fixed point, different
|
||||||
|
path to it. Follow the house pattern in
|
||||||
|
`tests/ingestion_equivalence.rs`: converge tightly (`max_iter: 2_000`,
|
||||||
|
`epsilon: 1e-12`) and compare within `1e-8`.
|
||||||
|
2. **Invariant to ingestion order** — events added one at a time produce
|
||||||
|
the same filtered results as the same events batched. Extends the
|
||||||
|
existing invariant in `tests/ingestion_equivalence.rs`.
|
||||||
|
3. **Single-slice exactness** — for a history with one time slice there
|
||||||
|
is no future to propagate back, so filtered results equal smoothed
|
||||||
|
results exactly.
|
||||||
|
4. **Uncertainty ordering** — for a competitor with many later games, σ
|
||||||
|
at the first filtered point is greater than σ at the first smoothed
|
||||||
|
point, and less than the prior σ. This is the ustat complaint restated
|
||||||
|
as an assertion.
|
||||||
|
5. **Degenerate inputs** — empty history yields `0.0` and empty maps;
|
||||||
|
unknown key yields an empty curve. Added to
|
||||||
|
`tests/degenerate_inputs.rs`.
|
||||||
|
|
||||||
|
### Regression net
|
||||||
|
|
||||||
|
The existing suite must be unchanged by the removals: `log_evidence()`,
|
||||||
|
`log_evidence_for()`, and every numerical golden keep their current
|
||||||
|
values, since the default `online` was already `false` and the flag was
|
||||||
|
inert.
|
||||||
|
|
||||||
|
## Verification gates
|
||||||
|
|
||||||
|
- `just test` — full matrix, including the release job. `debug_assert!`
|
||||||
|
is compiled out in release, and that is where defects in this crate
|
||||||
|
have hidden before.
|
||||||
|
- `just lint` — clippy, warnings denied.
|
||||||
|
- `just fmt` — nightly.
|
||||||
|
- `just determinism` — the new pass must not perturb bit-identical
|
||||||
|
posteriors across `RAYON_NUM_THREADS` 1/2/4/8.
|
||||||
|
- `#![forbid(unsafe_code)]` stays.
|
||||||
|
|
||||||
|
## Risks
|
||||||
|
|
||||||
|
- **Clone cost.** One slice's events are cloned per slice visited. At
|
||||||
|
ustat scale this is negligible, but the pass is O(events) allocation on
|
||||||
|
top of O(events) inference. Accepted: fidelity to the production sweep
|
||||||
|
is worth more than avoiding the clone, and no caller is on a hot path.
|
||||||
|
- **`iterate_to_convergence` leaving test-only status.** Its doc comment
|
||||||
|
claims "only used by tests"; that comment must be updated, or it
|
||||||
|
becomes the next piece of load-bearing prose that is quietly false.
|
||||||
|
- **Event order is inherited, not normalised.** The scratch clone takes
|
||||||
|
the real slice's current event order, which differs pre- and
|
||||||
|
post-`converge()` for incrementally-ingested slices (see *Invariants*).
|
||||||
|
Results agree to within convergence tolerance rather than exactly.
|
||||||
|
Normalising the order in the scratch builder would buy bit-identity at
|
||||||
|
the cost of diverging from what the real sweep does; not worth it.
|
||||||
|
|
||||||
|
**Measured after implementation, this risk is smaller than stated.**
|
||||||
|
Flipping the scratch's `color_groups_dirty` from `true` to `false`
|
||||||
|
switches it between the grouped sweep (`sweep_color_groups`) and the
|
||||||
|
sequential fallback across its entire convergence loop — a far larger
|
||||||
|
perturbation than a permuted event order — and the ingestion-order
|
||||||
|
invariance test stays green at `1e-8` under `max_iter: 2_000`,
|
||||||
|
`epsilon: 1e-12`. EP reaches the same fixed point regardless of sweep
|
||||||
|
order once driven far enough. The tolerance caveat is correct but
|
||||||
|
conservative. Note the flag itself is load-bearing: with it `false` the
|
||||||
|
scratch would take the sequential path always, diverging from the
|
||||||
|
production sweep it exists to mirror.
|
||||||
|
- **Divergence risk.** If `TimeSlice`'s sweep gains state that the
|
||||||
|
scratch construction does not initialise, the pass silently reads a
|
||||||
|
default. The scratch builder must construct `Skill` field-by-field
|
||||||
|
rather than via `..Default::default()`, so adding a field to `Skill`
|
||||||
|
is a compile error here rather than a silent wrong answer.
|
||||||
|
|
||||||
|
## Out-of-scope follow-ups
|
||||||
|
|
||||||
|
File as separate issues:
|
||||||
|
|
||||||
|
1. **`forward: bool` is only a filtering quantity pre-convergence**
|
||||||
|
(`src/history.rs:395`). Either document the constraint or fold the
|
||||||
|
flag into the new pass and delete it.
|
||||||
|
2. **`log_evidence` takes `&mut self`** (`src/history.rs:416`) but
|
||||||
|
mutates nothing. The new `filtered_*` methods take `&self`; the
|
||||||
|
asymmetry is worth removing.
|
||||||
+14
-2
@@ -1,2 +1,14 @@
|
|||||||
publish = false
|
# Publish to the registry named in Cargo.toml's `publish` list (kellnr).
|
||||||
pre-release-hook = ["sh", "-c", "git cliff -o CHANGELOG.md --tag {{version}} && git add CHANGELOG.md"]
|
publish = true
|
||||||
|
# Hold off pushing until tags and publish have both succeeded; `just release`
|
||||||
|
# pushes last.
|
||||||
|
push = false
|
||||||
|
# Regenerate the changelog and stage it so it lands in the release commit.
|
||||||
|
#
|
||||||
|
# Guarded on DRY_RUN because cargo-release runs pre-release hooks during a dry
|
||||||
|
# run too (verified against cargo-release 1.1.5, which exports DRY_RUN=true,
|
||||||
|
# CRATE_NAME, PREV_VERSION and NEW_VERSION to the hook). Without the guard,
|
||||||
|
# `just release-plan` — documented as a preview that writes nothing — writes and
|
||||||
|
# `git add`s CHANGELOG.md, and the clean-tree check in `just release` then
|
||||||
|
# refuses to run. That check is load-bearing: publishing is irreversible.
|
||||||
|
pre-release-hook = ["sh", "-c", '[ "$DRY_RUN" = "true" ] || (git cliff -o CHANGELOG.md --tag {{version}} && git add CHANGELOG.md)']
|
||||||
|
|||||||
@@ -0,0 +1,352 @@
|
|||||||
|
//! Active learning: which comparison teaches you the most.
|
||||||
|
//!
|
||||||
|
//! [`quality`](crate::quality) answers "is this matchup *fair*". That is a
|
||||||
|
//! different question from "is this matchup *informative*", and the two
|
||||||
|
//! coincide only for two evenly matched competitors. When each observation
|
||||||
|
//! costs something — a human click, a scheduled fixture — the question worth
|
||||||
|
//! asking is the second one.
|
||||||
|
//!
|
||||||
|
//! The quantity here is expected information gain: the outcome-weighted
|
||||||
|
//! divergence between what you believe now and what you would believe after
|
||||||
|
//! seeing the result.
|
||||||
|
//!
|
||||||
|
//! ```text
|
||||||
|
//! EIG(matchup) = SUM P(outcome) * KL( posterior_after(outcome) || prior )
|
||||||
|
//! outcome
|
||||||
|
//! ```
|
||||||
|
//!
|
||||||
|
//! It is the mutual information between the observed outcome and the skills,
|
||||||
|
//! which is worth remembering because it pins the scale: information gain
|
||||||
|
//! cannot exceed the entropy of the thing you are about to observe. A contest
|
||||||
|
//! with `k` distinguishable outcomes can teach you at most `ln k` nats,
|
||||||
|
//! whatever the ratings. That ceiling is the sharpest available test of an
|
||||||
|
//! implementation — see [`expected_information_gain`].
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
GameOptions, Gaussian, InferenceError, Outcome, Rating, drift::Drift, predict, time::Time,
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Outcomes below this probability contribute nothing measurable and are not
|
||||||
|
/// worth an inference pass.
|
||||||
|
///
|
||||||
|
/// The contribution of an outcome is `P * KL`, and `KL` is bounded in practice
|
||||||
|
/// by tens of nats, so a probability this small moves the total by less than
|
||||||
|
/// the quadrature error already present in `P` itself.
|
||||||
|
const NEGLIGIBLE: f64 = 1e-12;
|
||||||
|
|
||||||
|
/// `KL(q || p)` for two univariate Gaussians, in nats.
|
||||||
|
///
|
||||||
|
/// Both arguments are proper posteriors from inference, so the degenerate
|
||||||
|
/// cases guarded here (zero or infinite variance) indicate that inference has
|
||||||
|
/// broken down rather than anything a caller did.
|
||||||
|
fn kl_divergence(q: Gaussian, p: Gaussian) -> f64 {
|
||||||
|
let (var_q, var_p) = (q.sigma().powi(2), p.sigma().powi(2));
|
||||||
|
|
||||||
|
if !(var_q.is_finite() && var_p.is_finite()) || var_q <= 0.0 || var_p <= 0.0 {
|
||||||
|
return 0.0;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mean_gap = q.mu() - p.mu();
|
||||||
|
0.5 * (libm::log(var_p / var_q) + (var_q + mean_gap * mean_gap) / var_p - 1.0)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Expected information gain of a hypothetical matchup, in nats.
|
||||||
|
///
|
||||||
|
/// Enumerates the outcomes this matchup could have, runs inference for each to
|
||||||
|
/// get the belief it would produce, and weights the resulting divergence by
|
||||||
|
/// that outcome's probability. A higher value means the result would teach you
|
||||||
|
/// more.
|
||||||
|
///
|
||||||
|
/// # Interpreting the value
|
||||||
|
///
|
||||||
|
/// Nats. The upper bound is the entropy of the outcome variable: at most
|
||||||
|
/// `ln 2 ≈ 0.693` for a two-way result, `ln 3 ≈ 1.099` once draws are
|
||||||
|
/// possible, `ln k` for `k` outcomes. A value near the ceiling means the
|
||||||
|
/// result is close to a coin flip *and* would move the posteriors a long way;
|
||||||
|
/// a value near zero means you already know what will happen, or that the
|
||||||
|
/// result would barely change your beliefs if you saw it.
|
||||||
|
///
|
||||||
|
/// This is not a monotone transform of [`quality`](crate::quality). A lopsided
|
||||||
|
/// matchup between two uncertain competitors scores well on quality-times-
|
||||||
|
/// variance heuristics and poorly here, because the near-certain outcome
|
||||||
|
/// carries almost no information.
|
||||||
|
///
|
||||||
|
/// # Cost
|
||||||
|
///
|
||||||
|
/// One full inference pass per possible outcome, so this is far more expensive
|
||||||
|
/// than `quality()` — which is one closed-form evaluation. The outcome count
|
||||||
|
/// grows quickly with team count (3 outcomes for two teams that can draw, 13
|
||||||
|
/// for three, 75 for four), and scoring every candidate pairing among `n`
|
||||||
|
/// competitors is `O(n² × outcomes)` inference passes.
|
||||||
|
///
|
||||||
|
/// For a selector over many candidates, shortlist with the cheap
|
||||||
|
/// [`quality`](crate::quality) or
|
||||||
|
/// [`predict_win_probabilities`](crate::History::predict_win_probabilities)
|
||||||
|
/// first and score only the shortlist here. The expected-variance-reduction
|
||||||
|
/// proxy sometimes suggested as a cheaper alternative is *not* cheaper: it
|
||||||
|
/// needs the same hypothetical posteriors, so it shares the dominant cost.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// - `NotEnoughTeams` if fewer than two teams are supplied.
|
||||||
|
/// - `EmptyTeam` if any team has no members.
|
||||||
|
/// - `TooManyTeams` if the outcome space is too large to enumerate; see
|
||||||
|
/// [`MAX_PREDICTED_TEAMS`](crate::MAX_PREDICTED_TEAMS).
|
||||||
|
/// - `InvalidProbability` if `options.p_draw` is outside `[0.0, 1.0)`.
|
||||||
|
/// - Anything [`Game::ranked`](crate::Game::ranked) returns for a hypothetical
|
||||||
|
/// outcome.
|
||||||
|
pub fn expected_information_gain<T: Time, D: Drift<T>>(
|
||||||
|
teams: &[&[Rating<T, D>]],
|
||||||
|
options: &GameOptions,
|
||||||
|
) -> Result<f64, InferenceError> {
|
||||||
|
if teams.len() < 2 {
|
||||||
|
return Err(InferenceError::NotEnoughTeams { got: teams.len() });
|
||||||
|
}
|
||||||
|
if teams.len() > crate::MAX_PREDICTED_TEAMS {
|
||||||
|
return Err(InferenceError::TooManyTeams {
|
||||||
|
got: teams.len(),
|
||||||
|
max: crate::MAX_PREDICTED_TEAMS,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if !(0.0..1.0).contains(&options.p_draw) {
|
||||||
|
return Err(InferenceError::InvalidProbability {
|
||||||
|
value: options.p_draw,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
for (idx, team) in teams.iter().enumerate() {
|
||||||
|
if team.is_empty() {
|
||||||
|
return Err(InferenceError::EmptyTeam { team: idx });
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prediction runs on performances: skill inflated by each member's beta.
|
||||||
|
let performances: Vec<Gaussian> = teams
|
||||||
|
.iter()
|
||||||
|
.map(|team| {
|
||||||
|
team.iter()
|
||||||
|
.fold(crate::N00, |acc, rating| acc + rating.performance())
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
// Draw margins per pair, derived from the teams' betas exactly as
|
||||||
|
// inference derives them, so the outcomes weighted here are the outcomes
|
||||||
|
// that would actually be fitted.
|
||||||
|
let beta_sq: Vec<f64> = teams
|
||||||
|
.iter()
|
||||||
|
.map(|team| team.iter().map(|r| r.beta().powi(2)).sum())
|
||||||
|
.collect();
|
||||||
|
let p_draw = options.p_draw;
|
||||||
|
let margins = predict::Margins::new(teams.len(), |i, j| {
|
||||||
|
if p_draw == 0.0 {
|
||||||
|
0.0
|
||||||
|
} else {
|
||||||
|
crate::compute_margin(p_draw, (beta_sq[i] + beta_sq[j]).sqrt())
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
let mut gain = 0.0;
|
||||||
|
|
||||||
|
for (ranks, probability) in predict::outcome_distribution(&performances, &margins) {
|
||||||
|
if probability <= NEGLIGIBLE {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
let game = crate::Game::ranked(teams, Outcome::ranking(ranks), options)?;
|
||||||
|
let posteriors = game.posteriors();
|
||||||
|
|
||||||
|
// Beliefs factorise across competitors, so the joint divergence is the
|
||||||
|
// sum of the per-competitor ones.
|
||||||
|
let divergence: f64 = teams
|
||||||
|
.iter()
|
||||||
|
.zip(&posteriors)
|
||||||
|
.flat_map(|(team, posterior)| team.iter().zip(posterior))
|
||||||
|
.map(|(rating, &after)| kl_divergence(after, rating.prior()))
|
||||||
|
.sum();
|
||||||
|
|
||||||
|
gain += probability * divergence;
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(gain)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::{BETA, ConstantDrift, GAMMA};
|
||||||
|
|
||||||
|
type R = Rating<i64, ConstantDrift>;
|
||||||
|
|
||||||
|
fn rating(mu: f64, sigma: f64) -> R {
|
||||||
|
R::new(Gaussian::from_ms(mu, sigma), BETA, ConstantDrift(GAMMA))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn options(p_draw: f64) -> GameOptions {
|
||||||
|
GameOptions {
|
||||||
|
p_draw,
|
||||||
|
..GameOptions::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn eig(teams: &[&[R]], p_draw: f64) -> f64 {
|
||||||
|
expected_information_gain(teams, &options(p_draw)).unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The analytic ceiling. Information gain is the mutual information between
|
||||||
|
/// the outcome and the skills, so it cannot exceed the entropy of the
|
||||||
|
/// outcome variable — whatever the ratings. This is the check a subtly
|
||||||
|
/// wrong implementation fails while still returning plausible numbers: an
|
||||||
|
/// early prototype of this returned 4.77 nats from a sign error and passed
|
||||||
|
/// every monotonicity test.
|
||||||
|
#[test]
|
||||||
|
fn never_exceeds_the_entropy_of_the_outcome() {
|
||||||
|
let ceiling_two = std::f64::consts::LN_2;
|
||||||
|
|
||||||
|
for (a, b) in [
|
||||||
|
(rating(0.0, 6.0), rating(0.0, 6.0)),
|
||||||
|
(rating(0.0, 0.5), rating(0.0, 0.5)),
|
||||||
|
(rating(12.0, 6.0), rating(-12.0, 6.0)),
|
||||||
|
(rating(40.0, 1.0), rating(-40.0, 1.0)),
|
||||||
|
(rating(3.0, 6.0), rating(-2.0, 0.1)),
|
||||||
|
(rating(0.0, 25.0), rating(0.0, 25.0)),
|
||||||
|
] {
|
||||||
|
let g = eig(&[&[a], &[b]], 0.0);
|
||||||
|
assert!(
|
||||||
|
g >= 0.0 && g <= ceiling_two,
|
||||||
|
"EIG {g} outside [0, ln 2] for mu=({}, {}) sigma=({}, {})",
|
||||||
|
a.prior().mu(),
|
||||||
|
b.prior().mu(),
|
||||||
|
a.prior().sigma(),
|
||||||
|
b.prior().sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// With draws enabled there are three outcomes, so the ceiling rises to
|
||||||
|
/// `ln 3` — and the two-outcome bound no longer applies.
|
||||||
|
#[test]
|
||||||
|
fn the_ceiling_follows_the_outcome_count() {
|
||||||
|
let ceiling_three = 3.0f64.ln();
|
||||||
|
for sigma in [0.5, 3.0, 6.0, 25.0] {
|
||||||
|
let g = eig(&[&[rating(0.0, sigma)], &[rating(0.0, sigma)]], 0.25);
|
||||||
|
assert!(
|
||||||
|
g >= 0.0 && g <= ceiling_three,
|
||||||
|
"EIG {g} outside [0, ln 3] at sigma {sigma}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// An even matchup between uncertain competitors is the informative one.
|
||||||
|
/// A hopelessly lopsided matchup teaches you almost nothing, because you
|
||||||
|
/// already know how it ends.
|
||||||
|
#[test]
|
||||||
|
fn an_even_matchup_beats_a_lopsided_one() {
|
||||||
|
let even = eig(&[&[rating(0.0, 6.0)], &[rating(0.0, 6.0)]], 0.0);
|
||||||
|
let lopsided = eig(&[&[rating(12.0, 6.0)], &[rating(-12.0, 6.0)]], 0.0);
|
||||||
|
assert!(
|
||||||
|
even > lopsided,
|
||||||
|
"even {even} should beat lopsided {lopsided}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Certainty is the thing information gain is measuring the absence of:
|
||||||
|
/// the less you know, the more there is to learn.
|
||||||
|
#[test]
|
||||||
|
fn gain_falls_as_certainty_rises() {
|
||||||
|
let mut previous = f64::INFINITY;
|
||||||
|
for sigma in [12.0, 6.0, 3.0, 1.0, 0.5, 0.1] {
|
||||||
|
let g = eig(&[&[rating(0.0, sigma)], &[rating(0.0, sigma)]], 0.0);
|
||||||
|
assert!(
|
||||||
|
g < previous,
|
||||||
|
"sigma {sigma}: {g} did not fall below {previous}"
|
||||||
|
);
|
||||||
|
previous = g;
|
||||||
|
}
|
||||||
|
assert!(previous >= 0.0);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The heuristic this replaces is `quality * sigma_a^2 * sigma_b^2`. It is
|
||||||
|
/// not a monotone transform of information gain — it ranks a lopsided
|
||||||
|
/// matchup above a confident even one, and EIG ranks them the other way.
|
||||||
|
/// Pinning the disagreement down is what stops a future "simplification"
|
||||||
|
/// from quietly reverting to the heuristic.
|
||||||
|
#[test]
|
||||||
|
fn disagrees_with_the_quality_times_variance_heuristic() {
|
||||||
|
let heuristic = |a: &R, b: &R| {
|
||||||
|
crate::quality(&[&[a.prior()], &[b.prior()]], BETA)
|
||||||
|
* a.prior().sigma().powi(2)
|
||||||
|
* b.prior().sigma().powi(2)
|
||||||
|
};
|
||||||
|
|
||||||
|
let (confident_a, confident_b) = (rating(0.0, 0.5), rating(0.0, 0.5));
|
||||||
|
let (lopsided_a, lopsided_b) = (rating(12.0, 6.0), rating(-12.0, 6.0));
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
heuristic(&lopsided_a, &lopsided_b) > heuristic(&confident_a, &confident_b),
|
||||||
|
"the heuristic should prefer the lopsided matchup"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
eig(&[&[confident_a], &[confident_b]], 0.0) > eig(&[&[lopsided_a], &[lopsided_b]], 0.0),
|
||||||
|
"information gain should prefer the even matchup"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn supports_more_than_two_teams() {
|
||||||
|
let teams: Vec<Vec<R>> = vec![
|
||||||
|
vec![rating(0.0, 6.0)],
|
||||||
|
vec![rating(0.0, 6.0)],
|
||||||
|
vec![rating(0.0, 6.0)],
|
||||||
|
];
|
||||||
|
let refs: Vec<&[R]> = teams.iter().map(Vec::as_slice).collect();
|
||||||
|
let g = expected_information_gain(&refs, &options(0.0)).unwrap();
|
||||||
|
// Six distinguishable orderings with no draws.
|
||||||
|
assert!(
|
||||||
|
g > 0.0 && g <= 6.0f64.ln(),
|
||||||
|
"three-team EIG {g} out of range"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn multi_member_teams_are_supported() {
|
||||||
|
let a = [rating(0.0, 6.0), rating(1.0, 4.0)];
|
||||||
|
let b = [rating(0.0, 6.0)];
|
||||||
|
let g = expected_information_gain(&[&a, &b], &options(0.0)).unwrap();
|
||||||
|
assert!(g > 0.0 && g <= std::f64::consts::LN_2, "{g}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn degenerate_shapes_are_errors() {
|
||||||
|
let a = [rating(0.0, 6.0)];
|
||||||
|
assert!(matches!(
|
||||||
|
expected_information_gain(&[&a], &options(0.0)),
|
||||||
|
Err(InferenceError::NotEnoughTeams { got: 1 })
|
||||||
|
));
|
||||||
|
let empty: [R; 0] = [];
|
||||||
|
assert!(matches!(
|
||||||
|
expected_information_gain(&[&a, &empty], &options(0.0)),
|
||||||
|
Err(InferenceError::EmptyTeam { team: 1 })
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
expected_information_gain(&[&a, &a], &options(1.5)),
|
||||||
|
Err(InferenceError::InvalidProbability { .. })
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn kl_divergence_is_zero_for_identical_beliefs() {
|
||||||
|
let g = Gaussian::from_ms(3.0, 2.0);
|
||||||
|
assert!(kl_divergence(g, g).abs() < 1e-15);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn kl_divergence_is_non_negative_and_grows_with_separation() {
|
||||||
|
let prior = Gaussian::from_ms(0.0, 3.0);
|
||||||
|
let mut previous = 0.0;
|
||||||
|
for mu in [0.0, 0.5, 1.0, 2.0, 4.0] {
|
||||||
|
let d = kl_divergence(Gaussian::from_ms(mu, 3.0), prior);
|
||||||
|
assert!(d >= 0.0, "negative divergence at mu {mu}: {d}");
|
||||||
|
assert!(d >= previous, "not increasing at mu {mu}");
|
||||||
|
previous = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+46
-11
@@ -26,39 +26,75 @@ pub(crate) struct ColorGroups {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl ColorGroups {
|
impl ColorGroups {
|
||||||
#[allow(dead_code)]
|
|
||||||
pub(crate) fn new() -> Self {
|
pub(crate) fn new() -> Self {
|
||||||
Self::default()
|
Self::default()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)]
|
|
||||||
pub(crate) fn n_colors(&self) -> usize {
|
|
||||||
self.groups.len()
|
|
||||||
}
|
|
||||||
|
|
||||||
#[allow(dead_code)]
|
|
||||||
pub(crate) fn is_empty(&self) -> bool {
|
pub(crate) fn is_empty(&self) -> bool {
|
||||||
self.groups.is_empty()
|
self.groups.is_empty()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Total event count across all colors.
|
/// Number of distinct colors in the partition. Test-only.
|
||||||
#[allow(dead_code)]
|
#[cfg(test)]
|
||||||
|
pub(crate) fn n_colors(&self) -> usize {
|
||||||
|
self.groups.len()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Total event count across all colors. Test-only.
|
||||||
|
#[cfg(test)]
|
||||||
pub(crate) fn total_events(&self) -> usize {
|
pub(crate) fn total_events(&self) -> usize {
|
||||||
self.groups.iter().map(|g| g.len()).sum()
|
self.groups.iter().map(|g| g.len()).sum()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Contiguous index range for one color after events have been reordered
|
/// Contiguous index range for one color after events have been reordered
|
||||||
/// into color-contiguous positions by `TimeSlice::recompute_color_groups`.
|
/// into color-contiguous positions by `TimeSlice::recompute_color_groups`.
|
||||||
#[allow(dead_code)]
|
|
||||||
pub(crate) fn color_range(&self, color_idx: usize) -> std::ops::Range<usize> {
|
pub(crate) fn color_range(&self, color_idx: usize) -> std::ops::Range<usize> {
|
||||||
let group = &self.groups[color_idx];
|
let group = &self.groups[color_idx];
|
||||||
if group.is_empty() {
|
if group.is_empty() {
|
||||||
return 0..0;
|
return 0..0;
|
||||||
}
|
}
|
||||||
|
|
||||||
let start = *group.first().unwrap();
|
let start = *group.first().unwrap();
|
||||||
let end = *group.last().unwrap() + 1;
|
let end = *group.last().unwrap() + 1;
|
||||||
|
|
||||||
|
debug_assert_eq!(
|
||||||
|
end - start,
|
||||||
|
group.len(),
|
||||||
|
"color {color_idx} is not contiguous; its range would overlap other colors"
|
||||||
|
);
|
||||||
|
|
||||||
start..end
|
start..end
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Whether every color occupies a contiguous, ascending range of event
|
||||||
|
/// indices, and no two colors overlap.
|
||||||
|
///
|
||||||
|
/// The parallel sweep derives one `&mut` sub-slice per color from these
|
||||||
|
/// ranges and relies on them being disjoint. That disjointness is what
|
||||||
|
/// makes concurrent writes to distinct skills sound, so it is checked
|
||||||
|
/// rather than assumed.
|
||||||
|
pub(crate) fn groups_are_contiguous(&self) -> bool {
|
||||||
|
let mut expected_start = 0;
|
||||||
|
|
||||||
|
for group in &self.groups {
|
||||||
|
if group.is_empty() {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
let ascending_run = group
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.all(|(offset, &idx)| idx == group[0] + offset);
|
||||||
|
|
||||||
|
if !ascending_run || group[0] != expected_start {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
expected_start += group.len();
|
||||||
|
}
|
||||||
|
|
||||||
|
true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Compute color groups greedily.
|
/// Compute color groups greedily.
|
||||||
@@ -67,7 +103,6 @@ impl ColorGroups {
|
|||||||
/// `Index` values that event touches. The returned `ColorGroups` has one
|
/// `Index` values that event touches. The returned `ColorGroups` has one
|
||||||
/// inner `Vec<usize>` per color, containing event indices in the order
|
/// inner `Vec<usize>` per color, containing event indices in the order
|
||||||
/// they were assigned.
|
/// they were assigned.
|
||||||
#[allow(dead_code)]
|
|
||||||
pub(crate) fn color_greedy<I, F>(n_events: usize, index_set: F) -> ColorGroups
|
pub(crate) fn color_greedy<I, F>(n_events: usize, index_set: F) -> ColorGroups
|
||||||
where
|
where
|
||||||
F: Fn(usize) -> I,
|
F: Fn(usize) -> I,
|
||||||
|
|||||||
+22
-16
@@ -1,5 +1,4 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
N_INF,
|
|
||||||
drift::{ConstantDrift, Drift},
|
drift::{ConstantDrift, Drift},
|
||||||
gaussian::Gaussian,
|
gaussian::Gaussian,
|
||||||
rating::Rating,
|
rating::Rating,
|
||||||
@@ -8,12 +7,19 @@ use crate::{
|
|||||||
|
|
||||||
/// Per-history, temporal state for someone competing.
|
/// Per-history, temporal state for someone competing.
|
||||||
///
|
///
|
||||||
/// Renamed from `Agent` in T2; the former `.player` field is now
|
/// The mutable half of a competitor: `Rating` holds their static
|
||||||
/// `.rating` to match the `Player → Rating` rename.
|
/// configuration, this holds what inference learns as it sweeps.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct Competitor<T: Time = i64, D: Drift<T> = ConstantDrift> {
|
pub struct Competitor<T: Time = i64, D: Drift<T> = ConstantDrift> {
|
||||||
pub rating: Rating<T, D>,
|
pub rating: Rating<T, D>,
|
||||||
pub message: Gaussian,
|
/// The forward message carried from this competitor's last appearance, or
|
||||||
|
/// `None` before they have appeared anywhere.
|
||||||
|
///
|
||||||
|
/// Previously an improper `N_INF` served as the unset sentinel, which made
|
||||||
|
/// "no message yet" indistinguishable from "a legitimately improper
|
||||||
|
/// message" at the type level and required every reader to know the
|
||||||
|
/// convention.
|
||||||
|
pub message: Option<Gaussian>,
|
||||||
pub last_time: Option<T>,
|
pub last_time: Option<T>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -21,14 +27,16 @@ impl<T: Time, D: Drift<T>> Competitor<T, D> {
|
|||||||
/// Compute the message received at time `now`, with drift accumulated
|
/// Compute the message received at time `now`, with drift accumulated
|
||||||
/// from `self.last_time` (if any) to `now`.
|
/// from `self.last_time` (if any) to `now`.
|
||||||
pub(crate) fn receive(&self, now: &T) -> Gaussian {
|
pub(crate) fn receive(&self, now: &T) -> Gaussian {
|
||||||
if self.message != N_INF {
|
match self.message {
|
||||||
|
Some(message) => {
|
||||||
let elapsed_variance = match &self.last_time {
|
let elapsed_variance = match &self.last_time {
|
||||||
Some(last) => self.rating.drift.variance_delta(last, now),
|
Some(last) => self.rating.drift_variance_delta(last, now),
|
||||||
None => 0.0,
|
None => 0.0,
|
||||||
};
|
};
|
||||||
self.message.forget(elapsed_variance)
|
|
||||||
} else {
|
message.forget(elapsed_variance)
|
||||||
self.rating.prior
|
}
|
||||||
|
None => self.rating.prior,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -37,11 +45,9 @@ impl<T: Time, D: Drift<T>> Competitor<T, D> {
|
|||||||
/// Used in convergence sweeps where the elapsed was cached at slice-construction time
|
/// Used in convergence sweeps where the elapsed was cached at slice-construction time
|
||||||
/// and should not be recomputed from `last_time` (which may have shifted).
|
/// and should not be recomputed from `last_time` (which may have shifted).
|
||||||
pub(crate) fn receive_for_elapsed(&self, elapsed: i64) -> Gaussian {
|
pub(crate) fn receive_for_elapsed(&self, elapsed: i64) -> Gaussian {
|
||||||
if self.message != N_INF {
|
match self.message {
|
||||||
self.message
|
Some(message) => message.forget(self.rating.drift_variance_for_elapsed(elapsed)),
|
||||||
.forget(self.rating.drift.variance_for_elapsed(elapsed))
|
None => self.rating.prior,
|
||||||
} else {
|
|
||||||
self.rating.prior
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -50,7 +56,7 @@ impl Default for Competitor<i64, ConstantDrift> {
|
|||||||
fn default() -> Self {
|
fn default() -> Self {
|
||||||
Self {
|
Self {
|
||||||
rating: Rating::default(),
|
rating: Rating::default(),
|
||||||
message: N_INF,
|
message: None,
|
||||||
last_time: None,
|
last_time: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -63,7 +69,7 @@ where
|
|||||||
C: Iterator<Item = &'a mut Competitor<T, D>>,
|
C: Iterator<Item = &'a mut Competitor<T, D>>,
|
||||||
{
|
{
|
||||||
for c in competitors {
|
for c in competitors {
|
||||||
c.message = N_INF;
|
c.message = None;
|
||||||
if last_time {
|
if last_time {
|
||||||
c.last_time = None;
|
c.last_time = None;
|
||||||
}
|
}
|
||||||
|
|||||||
+31
-1
@@ -20,6 +20,37 @@ pub struct ConvergenceOptions {
|
|||||||
pub alpha: f64,
|
pub alpha: f64,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl ConvergenceOptions {
|
||||||
|
/// Reject values that would make inference silently meaningless.
|
||||||
|
///
|
||||||
|
/// `HistoryBuilder::convergence` asserts these eagerly, but the fields are
|
||||||
|
/// public and `GameOptions` carries a `ConvergenceOptions` — so a caller
|
||||||
|
/// can hand `Game::ranked` a set the builder never saw. In release the
|
||||||
|
/// engine's `debug_assert!`s are gone, and an `alpha` of zero leaves every
|
||||||
|
/// EP update unapplied: inference returns the priors, with every likelihood
|
||||||
|
/// uninformative and nothing to indicate anything went wrong.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// `InvalidParameter` if `alpha` is outside `(0.0, 1.0]` or `epsilon` is
|
||||||
|
/// negative. NaN fails both comparisons and is rejected.
|
||||||
|
pub(crate) fn validate(&self) -> Result<(), crate::InferenceError> {
|
||||||
|
if !(self.alpha > 0.0 && self.alpha <= 1.0) {
|
||||||
|
return Err(crate::InferenceError::InvalidParameter {
|
||||||
|
name: "alpha",
|
||||||
|
value: self.alpha,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if self.epsilon.is_nan() || self.epsilon < 0.0 {
|
||||||
|
return Err(crate::InferenceError::InvalidParameter {
|
||||||
|
name: "epsilon",
|
||||||
|
value: self.epsilon,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl Default for ConvergenceOptions {
|
impl Default for ConvergenceOptions {
|
||||||
fn default() -> Self {
|
fn default() -> Self {
|
||||||
Self {
|
Self {
|
||||||
@@ -38,7 +69,6 @@ pub struct ConvergenceReport {
|
|||||||
pub log_evidence: f64,
|
pub log_evidence: f64,
|
||||||
pub converged: bool,
|
pub converged: bool,
|
||||||
pub per_iteration_time: SmallVec<[Duration; 32]>,
|
pub per_iteration_time: SmallVec<[Duration; 32]>,
|
||||||
pub slices_skipped: usize,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
|||||||
+91
-13
@@ -1,6 +1,7 @@
|
|||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq)]
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum InferenceError {
|
pub enum InferenceError {
|
||||||
/// Expected and actual lengths of some array-shaped input differ.
|
/// Expected and actual lengths of some array-shaped input differ.
|
||||||
MismatchedShape {
|
MismatchedShape {
|
||||||
@@ -8,17 +9,61 @@ pub enum InferenceError {
|
|||||||
expected: usize,
|
expected: usize,
|
||||||
got: usize,
|
got: usize,
|
||||||
},
|
},
|
||||||
|
/// An `Outcome` of the wrong variant was supplied for the requested inference.
|
||||||
|
WrongOutcomeKind {
|
||||||
|
context: &'static str,
|
||||||
|
expected: &'static str,
|
||||||
|
got: &'static str,
|
||||||
|
},
|
||||||
/// A probability value is outside `[0, 1]`.
|
/// A probability value is outside `[0, 1]`.
|
||||||
InvalidProbability { value: f64 },
|
InvalidProbability { value: f64 },
|
||||||
/// A scalar parameter is outside its valid range.
|
/// A scalar parameter is outside its valid range.
|
||||||
InvalidParameter { name: &'static str, value: f64 },
|
InvalidParameter { name: &'static str, value: f64 },
|
||||||
/// Convergence exceeded `max_iter` without falling below `epsilon`.
|
/// An event contains tied teams, but the draw probability is zero.
|
||||||
ConvergenceFailed {
|
///
|
||||||
last_step: (f64, f64),
|
/// A zero draw probability asserts that draws cannot occur, so a tied
|
||||||
iterations: usize,
|
/// result has no representable likelihood. Configure a positive `p_draw`
|
||||||
|
/// (via `HistoryBuilder::p_draw` or `GameOptions::p_draw`) to admit ties.
|
||||||
|
TieWithoutDrawProbability { teams: (usize, usize) },
|
||||||
|
/// Inference produced a non-finite value (NaN or infinity).
|
||||||
|
///
|
||||||
|
/// Indicates numerical breakdown; the resulting skills are meaningless
|
||||||
|
/// and must not be treated as a converged estimate.
|
||||||
|
NonFiniteResult {
|
||||||
|
context: &'static str,
|
||||||
|
step: (f64, f64),
|
||||||
},
|
},
|
||||||
/// Negative precision: a Gaussian with `pi < 0` slipped into an API call.
|
/// One batch declared two different values for the same competitor's
|
||||||
NegativePrecision { pi: f64 },
|
/// configuration.
|
||||||
|
///
|
||||||
|
/// `prior` and `drift_scale` configure a competitor, not an event, so a
|
||||||
|
/// batch that sets one of them twice with different values has no
|
||||||
|
/// well-defined meaning: events within a batch are not ordered, so
|
||||||
|
/// "last one wins" would make the result depend on iteration order.
|
||||||
|
/// Declaring the same value repeatedly is fine and is the expected shape
|
||||||
|
/// when a competitor's configuration is a property of the domain.
|
||||||
|
ConflictingCompetitorConfig {
|
||||||
|
competitor: usize,
|
||||||
|
field: &'static str,
|
||||||
|
},
|
||||||
|
/// A prediction referenced a key the history has no skill for.
|
||||||
|
///
|
||||||
|
/// Reported rather than skipped: dropping unknown keys turns a team of
|
||||||
|
/// strangers into a confident-looking probability about nobody.
|
||||||
|
UnknownKey { team: usize, member: usize },
|
||||||
|
/// A prediction was given a team with no members.
|
||||||
|
EmptyTeam { team: usize },
|
||||||
|
/// Fewer than two teams were supplied to a prediction.
|
||||||
|
NotEnoughTeams { got: usize },
|
||||||
|
/// The full outcome distribution was requested for too many teams.
|
||||||
|
///
|
||||||
|
/// Each realisation sorts into exactly one (order, tie-pattern) event, so
|
||||||
|
/// the space holds `n! * 2^(n-1)` members — 1_920 at five teams, 23_040 at
|
||||||
|
/// six, 322_560 at seven. Past `max` this stops being something to
|
||||||
|
/// enumerate on a caller's behalf; ask for individual rankings with
|
||||||
|
/// `predict_ranking`, or for `predict_win_probabilities`, both of which
|
||||||
|
/// stay cheap at any team count.
|
||||||
|
TooManyTeams { got: usize, max: usize },
|
||||||
}
|
}
|
||||||
|
|
||||||
impl fmt::Display for InferenceError {
|
impl fmt::Display for InferenceError {
|
||||||
@@ -31,23 +76,56 @@ impl fmt::Display for InferenceError {
|
|||||||
} => {
|
} => {
|
||||||
write!(f, "{kind}: expected length {expected}, got {got}")
|
write!(f, "{kind}: expected length {expected}, got {got}")
|
||||||
}
|
}
|
||||||
|
Self::WrongOutcomeKind {
|
||||||
|
context,
|
||||||
|
expected,
|
||||||
|
got,
|
||||||
|
} => {
|
||||||
|
write!(f, "{context}: expected {expected}, got {got}")
|
||||||
|
}
|
||||||
Self::InvalidProbability { value } => {
|
Self::InvalidProbability { value } => {
|
||||||
write!(f, "probability must be in [0, 1]; got {value}")
|
write!(f, "probability must be in [0, 1]; got {value}")
|
||||||
}
|
}
|
||||||
|
Self::TieWithoutDrawProbability { teams } => {
|
||||||
|
write!(
|
||||||
|
f,
|
||||||
|
"teams {} and {} are tied, but p_draw is 0.0; set a positive draw probability to admit ties",
|
||||||
|
teams.0, teams.1
|
||||||
|
)
|
||||||
|
}
|
||||||
|
Self::NonFiniteResult { context, step } => {
|
||||||
|
write!(
|
||||||
|
f,
|
||||||
|
"{context}: inference produced a non-finite result (step = {step:?})"
|
||||||
|
)
|
||||||
|
}
|
||||||
Self::InvalidParameter { name, value } => {
|
Self::InvalidParameter { name, value } => {
|
||||||
write!(f, "{name} is invalid: {value}")
|
write!(f, "{name} is invalid: {value}")
|
||||||
}
|
}
|
||||||
Self::ConvergenceFailed {
|
Self::ConflictingCompetitorConfig { competitor, field } => {
|
||||||
last_step,
|
|
||||||
iterations,
|
|
||||||
} => {
|
|
||||||
write!(
|
write!(
|
||||||
f,
|
f,
|
||||||
"convergence failed after {iterations} iterations; last step = {last_step:?}"
|
"competitor {competitor}: this batch sets {field} to two different values"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
Self::NegativePrecision { pi } => {
|
Self::UnknownKey { team, member } => {
|
||||||
write!(f, "precision must be non-negative; got {pi}")
|
write!(
|
||||||
|
f,
|
||||||
|
"team {team}, member {member}: no skill recorded for this key"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
Self::EmptyTeam { team } => {
|
||||||
|
write!(f, "team {team} has no members")
|
||||||
|
}
|
||||||
|
Self::NotEnoughTeams { got } => {
|
||||||
|
write!(f, "prediction needs at least 2 teams, got {got}")
|
||||||
|
}
|
||||||
|
Self::TooManyTeams { got, max } => {
|
||||||
|
write!(
|
||||||
|
f,
|
||||||
|
"the outcome distribution over {got} teams is too large to enumerate (limit {max}); \
|
||||||
|
use predict_ranking or predict_win_probabilities instead"
|
||||||
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+49
-6
@@ -1,8 +1,10 @@
|
|||||||
//! Typed event description for bulk ingestion.
|
//! Typed event description for bulk ingestion.
|
||||||
//!
|
//!
|
||||||
//! `Event<T, K>` is the new public event shape (spec Section 4). Replaces
|
//! `Event<T, K>` is the public event shape taken by `History::add_events`. It
|
||||||
//! the nested `Vec<Vec<Vec<Index>>>`, `Vec<Vec<f64>>`, `Vec<Vec<Vec<f64>>>`
|
//! is a typed front end, not a replacement: `add_events` flattens it into the
|
||||||
//! that the old `add_events_with_prior` took.
|
//! nested `Vec<Vec<Vec<Index>>>` / `Vec<Vec<f64>>` / `Vec<Vec<Vec<f64>>>` that
|
||||||
|
//! the internal `add_events_with_prior` chokepoint still takes, and which
|
||||||
|
//! `record_winner` and `record_draw` also route through.
|
||||||
|
|
||||||
use smallvec::SmallVec;
|
use smallvec::SmallVec;
|
||||||
|
|
||||||
@@ -23,6 +25,7 @@ pub struct Team<K> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl<K> Team<K> {
|
impl<K> Team<K> {
|
||||||
|
#[must_use]
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
members: SmallVec::new(),
|
members: SmallVec::new(),
|
||||||
@@ -44,13 +47,28 @@ impl<K> Default for Team<K> {
|
|||||||
|
|
||||||
/// One member of a team, identified by user key `K`.
|
/// One member of a team, identified by user key `K`.
|
||||||
///
|
///
|
||||||
/// `weight` defaults to 1.0; a per-event `prior` can override the competitor's
|
/// `weight` applies per event and defaults to 1.0.
|
||||||
/// current skill estimate for this event only.
|
///
|
||||||
|
/// `prior` and `drift_scale` are **competitor configuration**, not per-event
|
||||||
|
/// values. Setting either applies to the competitor for the whole history, not
|
||||||
|
/// just to this event, and applies whenever it is supplied — including on a key
|
||||||
|
/// the history already knows. Because configuration lives on the competitor and
|
||||||
|
/// `converge` refits from competitor state, configuring one late still refits
|
||||||
|
/// the whole history rather than taking effect only from that event onward.
|
||||||
|
///
|
||||||
|
/// Repeating the same value is inert, which is the expected shape when the
|
||||||
|
/// configuration is a property of the domain. Supplying two *different* values
|
||||||
|
/// for one competitor within a single batch is
|
||||||
|
/// `InferenceError::ConflictingCompetitorConfig`: events in a batch have no
|
||||||
|
/// order, so there would be no well-defined winner.
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct Member<K> {
|
pub struct Member<K> {
|
||||||
pub key: K,
|
pub key: K,
|
||||||
pub weight: f64,
|
pub weight: f64,
|
||||||
pub prior: Option<Gaussian>,
|
pub prior: Option<Gaussian>,
|
||||||
|
/// Multiplier on the drift *variance* this competitor accumulates.
|
||||||
|
/// `None` means 1.0.
|
||||||
|
pub drift_scale: Option<f64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<K> Member<K> {
|
impl<K> Member<K> {
|
||||||
@@ -59,6 +77,7 @@ impl<K> Member<K> {
|
|||||||
key,
|
key,
|
||||||
weight: 1.0,
|
weight: 1.0,
|
||||||
prior: None,
|
prior: None,
|
||||||
|
drift_scale: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -67,10 +86,31 @@ impl<K> Member<K> {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set this competitor's starting skill estimate.
|
||||||
|
///
|
||||||
|
/// Captured at the competitor's first appearance; see the type docs.
|
||||||
pub fn with_prior(mut self, prior: Gaussian) -> Self {
|
pub fn with_prior(mut self, prior: Gaussian) -> Self {
|
||||||
self.prior = Some(prior);
|
self.prior = Some(prior);
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Scale how fast this competitor drifts, relative to the history's drift.
|
||||||
|
///
|
||||||
|
/// The scale multiplies the drift *variance*, so it is in the same units as
|
||||||
|
/// `gamma`: `ConstantDrift(g)` at `scale = s` behaves exactly as
|
||||||
|
/// `ConstantDrift(g * s)` would for this competitor alone.
|
||||||
|
///
|
||||||
|
/// `0.0` pins the competitor still — useful for a reference point that
|
||||||
|
/// shares a scale with moving competitors but should not itself move: a bot
|
||||||
|
/// at a known strength, a rating floor, a course difficulty.
|
||||||
|
///
|
||||||
|
/// Captured at the competitor's first appearance; see the type docs.
|
||||||
|
/// Must be finite and non-negative, or ingestion fails with
|
||||||
|
/// [`InferenceError::InvalidParameter`](crate::InferenceError::InvalidParameter).
|
||||||
|
pub fn with_drift_scale(mut self, scale: f64) -> Self {
|
||||||
|
self.drift_scale = Some(scale);
|
||||||
|
self
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Convenience: a member is a user key with default weight 1.0 and no prior.
|
/// Convenience: a member is a user key with default weight 1.0 and no prior.
|
||||||
@@ -91,15 +131,18 @@ mod tests {
|
|||||||
assert_eq!(m.key, "alice");
|
assert_eq!(m.key, "alice");
|
||||||
assert_eq!(m.weight, 1.0);
|
assert_eq!(m.weight, 1.0);
|
||||||
assert!(m.prior.is_none());
|
assert!(m.prior.is_none());
|
||||||
|
assert!(m.drift_scale.is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn member_builder_methods_chain() {
|
fn member_builder_methods_chain() {
|
||||||
let m = Member::new("alice")
|
let m = Member::new("alice")
|
||||||
.with_weight(0.5)
|
.with_weight(0.5)
|
||||||
.with_prior(Gaussian::from_ms(20.0, 5.0));
|
.with_prior(Gaussian::from_ms(20.0, 5.0))
|
||||||
|
.with_drift_scale(0.0);
|
||||||
assert_eq!(m.weight, 0.5);
|
assert_eq!(m.weight, 0.5);
|
||||||
assert!(m.prior.is_some());
|
assert!(m.prior.is_some());
|
||||||
|
assert_eq!(m.drift_scale, Some(0.0));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
+44
-8
@@ -19,6 +19,14 @@ where
|
|||||||
history: &'h mut History<T, D, O, K>,
|
history: &'h mut History<T, D, O, K>,
|
||||||
event: Event<T, K>,
|
event: Event<T, K>,
|
||||||
current_team_idx: Option<usize>,
|
current_team_idx: Option<usize>,
|
||||||
|
/// First validation failure seen while building, surfaced by `commit`.
|
||||||
|
///
|
||||||
|
/// The setters return `Self` so the chain stays fluent; they cannot return
|
||||||
|
/// a `Result` without breaking that. Recording the failure and reporting it
|
||||||
|
/// at `commit` keeps the check enforced in release, where the previous
|
||||||
|
/// `debug_assert!` was compiled out and a mismatched event was ingested
|
||||||
|
/// silently.
|
||||||
|
error: Option<InferenceError>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'h, T, D, O, K> EventBuilder<'h, T, D, O, K>
|
impl<'h, T, D, O, K> EventBuilder<'h, T, D, O, K>
|
||||||
@@ -37,6 +45,7 @@ where
|
|||||||
outcome: Outcome::Ranked(SmallVec::new()),
|
outcome: Outcome::Ranked(SmallVec::new()),
|
||||||
},
|
},
|
||||||
current_team_idx: None,
|
current_team_idx: None,
|
||||||
|
error: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -50,22 +59,36 @@ where
|
|||||||
|
|
||||||
/// Set per-member weights for the most recently added team.
|
/// Set per-member weights for the most recently added team.
|
||||||
///
|
///
|
||||||
/// Panics in debug builds if called before `.team(...)` or if the length
|
/// A length mismatch is recorded and returned by [`EventBuilder::commit`]
|
||||||
/// doesn't match the team's member count.
|
/// as `InferenceError::MismatchedShape`, in both debug and release. The
|
||||||
|
/// weights are not applied in that case, so a partially-weighted team
|
||||||
|
/// cannot reach the history.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if called before any `.team(...)`.
|
||||||
pub fn weights<I: IntoIterator<Item = f64>>(mut self, weights: I) -> Self {
|
pub fn weights<I: IntoIterator<Item = f64>>(mut self, weights: I) -> Self {
|
||||||
let idx = self
|
let idx = self
|
||||||
.current_team_idx
|
.current_team_idx
|
||||||
.expect(".weights(...) called before any .team(...)");
|
.expect(".weights(...) called before any .team(...)");
|
||||||
|
|
||||||
let ws: Vec<f64> = weights.into_iter().collect();
|
let ws: Vec<f64> = weights.into_iter().collect();
|
||||||
let team = &mut self.event.teams[idx];
|
let team = &mut self.event.teams[idx];
|
||||||
debug_assert_eq!(
|
|
||||||
ws.len(),
|
if ws.len() != team.members.len() {
|
||||||
team.members.len(),
|
self.error.get_or_insert(InferenceError::MismatchedShape {
|
||||||
"weights length must match team size"
|
kind: "weights",
|
||||||
);
|
expected: team.members.len(),
|
||||||
|
got: ws.len(),
|
||||||
|
});
|
||||||
|
|
||||||
|
return self;
|
||||||
|
}
|
||||||
|
|
||||||
for (m, w) in team.members.iter_mut().zip(ws) {
|
for (m, w) in team.members.iter_mut().zip(ws) {
|
||||||
m.weight = w;
|
m.weight = w;
|
||||||
}
|
}
|
||||||
|
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,7 +107,10 @@ where
|
|||||||
/// Set explicit per-team continuous scores with a per-event noise override.
|
/// Set explicit per-team continuous scores with a per-event noise override.
|
||||||
///
|
///
|
||||||
/// `sigma` overrides `HistoryBuilder::score_sigma` for this event only.
|
/// `sigma` overrides `HistoryBuilder::score_sigma` for this event only.
|
||||||
/// Must be `> 0.0`; debug-asserts otherwise via `Outcome::scores_with_sigma`.
|
/// Must be `> 0.0`. Constructing the outcome with a non-positive or NaN
|
||||||
|
/// sigma is allowed; the value is rejected with
|
||||||
|
/// `InferenceError::InvalidParameter` when the event is ingested, so
|
||||||
|
/// callers get an error from `commit` rather than a panic.
|
||||||
pub fn scores_with_sigma<I: IntoIterator<Item = f64>>(mut self, scores: I, sigma: f64) -> Self {
|
pub fn scores_with_sigma<I: IntoIterator<Item = f64>>(mut self, scores: I, sigma: f64) -> Self {
|
||||||
self.event.outcome = crate::Outcome::scores_with_sigma(scores, sigma);
|
self.event.outcome = crate::Outcome::scores_with_sigma(scores, sigma);
|
||||||
self
|
self
|
||||||
@@ -103,7 +129,17 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Commit the event to the history.
|
/// Commit the event to the history.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// Returns the first validation failure recorded while building — see
|
||||||
|
/// [`EventBuilder::weights`] — otherwise forwards to
|
||||||
|
/// [`History::add_events`] and returns its errors.
|
||||||
pub fn commit(self) -> Result<(), InferenceError> {
|
pub fn commit(self) -> Result<(), InferenceError> {
|
||||||
|
if let Some(error) = self.error {
|
||||||
|
return Err(error);
|
||||||
|
}
|
||||||
|
|
||||||
self.history.add_events(std::iter::once(self.event))
|
self.history.add_events(std::iter::once(self.event))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+34
-14
@@ -2,7 +2,7 @@ use crate::{
|
|||||||
N_INF,
|
N_INF,
|
||||||
factor::{Factor, VarId, VarStore},
|
factor::{Factor, VarId, VarStore},
|
||||||
gaussian::Gaussian,
|
gaussian::Gaussian,
|
||||||
pdf,
|
ln_pdf,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Gaussian observation factor on a diff variable.
|
/// Gaussian observation factor on a diff variable.
|
||||||
@@ -16,10 +16,11 @@ pub struct MarginFactor {
|
|||||||
pub m_obs: f64,
|
pub m_obs: f64,
|
||||||
pub sigma: f64,
|
pub sigma: f64,
|
||||||
pub(crate) msg: Gaussian,
|
pub(crate) msg: Gaussian,
|
||||||
pub(crate) evidence_cached: Option<f64>,
|
pub(crate) log_evidence_cached: Option<f64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MarginFactor {
|
impl MarginFactor {
|
||||||
|
#[must_use]
|
||||||
pub fn new(diff: VarId, m_obs: f64, sigma: f64) -> Self {
|
pub fn new(diff: VarId, m_obs: f64, sigma: f64) -> Self {
|
||||||
debug_assert!(sigma > 0.0, "score sigma must be positive");
|
debug_assert!(sigma > 0.0, "score sigma must be positive");
|
||||||
Self {
|
Self {
|
||||||
@@ -27,7 +28,7 @@ impl MarginFactor {
|
|||||||
m_obs,
|
m_obs,
|
||||||
sigma,
|
sigma,
|
||||||
msg: N_INF,
|
msg: N_INF,
|
||||||
evidence_cached: None,
|
log_evidence_cached: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -40,8 +41,8 @@ impl MarginFactor {
|
|||||||
let marginal = vars.get(self.diff);
|
let marginal = vars.get(self.diff);
|
||||||
let cavity = marginal / self.msg;
|
let cavity = marginal / self.msg;
|
||||||
|
|
||||||
if self.evidence_cached.is_none() {
|
if self.log_evidence_cached.is_none() {
|
||||||
self.evidence_cached = Some(cavity_evidence(cavity, self.m_obs, self.sigma));
|
self.log_evidence_cached = Some(cavity_log_evidence(cavity, self.m_obs, self.sigma));
|
||||||
}
|
}
|
||||||
|
|
||||||
let new_msg = Gaussian::from_ms(self.m_obs, self.sigma);
|
let new_msg = Gaussian::from_ms(self.m_obs, self.sigma);
|
||||||
@@ -60,13 +61,32 @@ impl Factor for MarginFactor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn log_evidence(&self, _vars: &VarStore) -> f64 {
|
fn log_evidence(&self, _vars: &VarStore) -> f64 {
|
||||||
self.evidence_cached.unwrap_or(1.0).ln()
|
self.log_evidence_cached.unwrap_or(0.0)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn cavity_evidence(cavity: Gaussian, m_obs: f64, sigma: f64) -> f64 {
|
/// `ln` of the observed margin's density under the cavity.
|
||||||
let combined_sigma = (cavity.sigma().powi(2) + sigma.powi(2)).sqrt();
|
///
|
||||||
pdf(m_obs, cavity.mu(), combined_sigma)
|
/// Computed in log space rather than as `pdf(..).ln()`. The density underflows
|
||||||
|
/// to zero past about 38 sigma of separation, and clamping that to
|
||||||
|
/// `f64::MIN_POSITIVE` reported -708 nats however far out the observation
|
||||||
|
/// actually was — 4292 nats adrift at 100 sigma, and unbounded beyond. A score
|
||||||
|
/// far from what the model expected is exactly the observation a log-evidence
|
||||||
|
/// figure exists to notice.
|
||||||
|
fn cavity_log_evidence(cavity: Gaussian, m_obs: f64, sigma: f64) -> f64 {
|
||||||
|
// `hypot`, not `sqrt(a^2 + b^2)`: squaring overflows to infinity above a
|
||||||
|
// sigma of ~1.3e154 and flushes to zero below ~1.5e-154, and `Gaussian`'s
|
||||||
|
// constructors are public so a caller can reach both.
|
||||||
|
let combined_sigma = cavity.sigma().hypot(sigma);
|
||||||
|
let value = ln_pdf(m_obs, cavity.mu(), combined_sigma);
|
||||||
|
|
||||||
|
// A degenerate cavity (infinite sigma) is the only way to reach a
|
||||||
|
// non-finite result; fall back to the old floor rather than emit -inf.
|
||||||
|
if value.is_finite() {
|
||||||
|
value
|
||||||
|
} else {
|
||||||
|
f64::MIN_POSITIVE.ln()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -108,16 +128,16 @@ mod tests {
|
|||||||
let mut vars = VarStore::new();
|
let mut vars = VarStore::new();
|
||||||
let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0));
|
let diff = vars.alloc(Gaussian::from_ms(0.0, 6.0));
|
||||||
let mut f = MarginFactor::new(diff, 5.0, 1.0);
|
let mut f = MarginFactor::new(diff, 5.0, 1.0);
|
||||||
assert!(f.evidence_cached.is_none());
|
assert!(f.log_evidence_cached.is_none());
|
||||||
|
|
||||||
f.propagate(&mut vars);
|
f.propagate(&mut vars);
|
||||||
let z = f.evidence_cached.unwrap();
|
let z = f.log_evidence_cached.unwrap();
|
||||||
// pdf(5, 0, sqrt(37)) ≈ 0.046783
|
// ln pdf(5, 0, sqrt(37)) = ln(0.046783...)
|
||||||
assert!((z - 0.04678300292616668).abs() < 1e-10);
|
assert!((z.exp() - 0.04678300292616668).abs() < 1e-10);
|
||||||
|
|
||||||
// Subsequent propagations don't change it.
|
// Subsequent propagations don't change it.
|
||||||
f.propagate(&mut vars);
|
f.propagate(&mut vars);
|
||||||
assert_eq!(f.evidence_cached.unwrap(), z);
|
assert_eq!(f.log_evidence_cached.unwrap(), z);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ pub struct VarStore {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl VarStore {
|
impl VarStore {
|
||||||
|
#[must_use]
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self::default()
|
Self::default()
|
||||||
}
|
}
|
||||||
@@ -28,10 +29,12 @@ impl VarStore {
|
|||||||
self.marginals.clear();
|
self.marginals.clear();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.marginals.len()
|
self.marginals.len()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn is_empty(&self) -> bool {
|
pub fn is_empty(&self) -> bool {
|
||||||
self.marginals.is_empty()
|
self.marginals.is_empty()
|
||||||
}
|
}
|
||||||
@@ -42,6 +45,7 @@ impl VarStore {
|
|||||||
id
|
id
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn get(&self, id: VarId) -> Gaussian {
|
pub fn get(&self, id: VarId) -> Gaussian {
|
||||||
self.marginals[id.0 as usize]
|
self.marginals[id.0 as usize]
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,12 +5,12 @@ use crate::factor::{Factor, VarId, VarStore};
|
|||||||
/// On each propagation:
|
/// On each propagation:
|
||||||
/// - Reads marginals at `team_a` and `team_b` (which already incorporate any
|
/// - Reads marginals at `team_a` and `team_b` (which already incorporate any
|
||||||
/// incoming messages from neighboring factors).
|
/// incoming messages from neighboring factors).
|
||||||
/// - Computes `new_diff = team_a - team_b` (variance addition; see Gaussian::Sub).
|
/// - Computes `new_diff = team_a - team_b` (variance addition; see `Gaussian::Sub`).
|
||||||
/// - Writes the new marginal to `diff`.
|
/// - Writes the new marginal to `diff`.
|
||||||
/// - Returns the delta against the previous diff value.
|
/// - Returns the delta against the previous diff value.
|
||||||
///
|
///
|
||||||
/// This factor does NOT store an outgoing message; the diff variable is
|
/// This factor does NOT store an outgoing message; the diff variable is
|
||||||
/// effectively replaced on each propagation. The TruncFactor on the same diff
|
/// effectively replaced on each propagation. The `TruncFactor` on the same diff
|
||||||
/// var holds the EP-divide message that produces the cavity.
|
/// var holds the EP-divide message that produces the cavity.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct RankDiffFactor {
|
pub struct RankDiffFactor {
|
||||||
|
|||||||
+109
-19
@@ -1,7 +1,8 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
N_INF, approx, cdf,
|
N_INF, approx,
|
||||||
factor::{Factor, VarId, VarStore},
|
factor::{Factor, VarId, VarStore},
|
||||||
gaussian::Gaussian,
|
gaussian::Gaussian,
|
||||||
|
ln_interval, ln_sf,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// EP truncation factor on a diff variable.
|
/// EP truncation factor on a diff variable.
|
||||||
@@ -15,20 +16,21 @@ pub struct TruncFactor {
|
|||||||
pub diff: VarId,
|
pub diff: VarId,
|
||||||
pub margin: f64,
|
pub margin: f64,
|
||||||
pub tie: bool,
|
pub tie: bool,
|
||||||
/// Outgoing message to the diff variable (initial: N_INF, the EP identity).
|
/// Outgoing message to the diff variable (initial: `N_INF`, the EP identity).
|
||||||
pub(crate) msg: Gaussian,
|
pub(crate) msg: Gaussian,
|
||||||
/// Cached evidence (linear, not log) computed from the cavity on first propagation.
|
/// Cached evidence (linear, not log) computed from the cavity on first propagation.
|
||||||
pub(crate) evidence_cached: Option<f64>,
|
pub(crate) log_evidence_cached: Option<f64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl TruncFactor {
|
impl TruncFactor {
|
||||||
|
#[must_use]
|
||||||
pub fn new(diff: VarId, margin: f64, tie: bool) -> Self {
|
pub fn new(diff: VarId, margin: f64, tie: bool) -> Self {
|
||||||
Self {
|
Self {
|
||||||
diff,
|
diff,
|
||||||
margin,
|
margin,
|
||||||
tie,
|
tie,
|
||||||
msg: N_INF,
|
msg: N_INF,
|
||||||
evidence_cached: None,
|
log_evidence_cached: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -41,8 +43,8 @@ impl TruncFactor {
|
|||||||
let marginal = vars.get(self.diff);
|
let marginal = vars.get(self.diff);
|
||||||
let cavity = marginal / self.msg;
|
let cavity = marginal / self.msg;
|
||||||
|
|
||||||
if self.evidence_cached.is_none() {
|
if self.log_evidence_cached.is_none() {
|
||||||
self.evidence_cached = Some(cavity_evidence(cavity, self.margin, self.tie));
|
self.log_evidence_cached = Some(cavity_log_evidence(cavity, self.margin, self.tie));
|
||||||
}
|
}
|
||||||
|
|
||||||
let trunc = approx(cavity, self.margin, self.tie);
|
let trunc = approx(cavity, self.margin, self.tie);
|
||||||
@@ -67,16 +69,33 @@ impl Factor for TruncFactor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn log_evidence(&self, _vars: &VarStore) -> f64 {
|
fn log_evidence(&self, _vars: &VarStore) -> f64 {
|
||||||
self.evidence_cached.unwrap_or(1.0).ln()
|
self.log_evidence_cached.unwrap_or(0.0)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// P(diff > margin) for non-tie, P(|diff| < margin) for tie.
|
/// `ln P(diff > margin)` for a win, `ln P(|diff| < margin)` for a tie.
|
||||||
fn cavity_evidence(diff: Gaussian, margin: f64, tie: bool) -> f64 {
|
///
|
||||||
if tie {
|
/// Computed in log space throughout. Two earlier shapes both lost the tail:
|
||||||
cdf(margin, diff.mu(), diff.sigma()) - cdf(-margin, diff.mu(), diff.sigma())
|
/// `1 - cdf(..)` cancelled away every digit of an unlikely outcome, and even
|
||||||
|
/// once that was fixed the linear probability underflows to zero past about 38
|
||||||
|
/// sigma, where clamping reported -708 nats regardless of the truth. An upset
|
||||||
|
/// is the observation a log-evidence figure exists to notice, so it has to stay
|
||||||
|
/// exact precisely where it is smallest.
|
||||||
|
fn cavity_log_evidence(diff: Gaussian, margin: f64, tie: bool) -> f64 {
|
||||||
|
let (mu, sigma) = (diff.mu(), diff.sigma());
|
||||||
|
|
||||||
|
let value = if tie {
|
||||||
|
ln_interval(-margin, margin, mu, sigma)
|
||||||
} else {
|
} else {
|
||||||
1.0 - cdf(margin, diff.mu(), diff.sigma())
|
ln_sf(margin, mu, sigma)
|
||||||
|
};
|
||||||
|
|
||||||
|
// A degenerate cavity is the only route to a non-finite result; keep the
|
||||||
|
// old floor for it rather than letting -inf poison the whole history's sum.
|
||||||
|
if value.is_finite() {
|
||||||
|
value
|
||||||
|
} else {
|
||||||
|
f64::MIN_POSITIVE.ln()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -108,19 +127,90 @@ mod tests {
|
|||||||
let diff = vars.alloc(Gaussian::from_ms(2.0, 3.0));
|
let diff = vars.alloc(Gaussian::from_ms(2.0, 3.0));
|
||||||
|
|
||||||
let mut f = TruncFactor::new(diff, 0.0, false);
|
let mut f = TruncFactor::new(diff, 0.0, false);
|
||||||
assert!(f.evidence_cached.is_none());
|
assert!(f.log_evidence_cached.is_none());
|
||||||
|
|
||||||
f.propagate(&mut vars);
|
f.propagate(&mut vars);
|
||||||
assert!(f.evidence_cached.is_some());
|
assert!(f.log_evidence_cached.is_some());
|
||||||
let first = f.evidence_cached.unwrap();
|
let first = f.log_evidence_cached.unwrap();
|
||||||
|
|
||||||
// Evidence should be P(diff > 0) for diff ~ N(2, 9) ≈ 0.748
|
// Evidence should be P(diff > 0) for diff ~ N(2, 9) ≈ 0.748
|
||||||
assert!(first > 0.7);
|
assert!(first.exp() > 0.7);
|
||||||
assert!(first < 0.8);
|
assert!(first.exp() < 0.8);
|
||||||
|
|
||||||
// Subsequent propagations don't change it.
|
// Subsequent propagations don't change it.
|
||||||
f.propagate(&mut vars);
|
f.propagate(&mut vars);
|
||||||
assert_eq!(f.evidence_cached.unwrap(), first);
|
assert_eq!(f.log_evidence_cached.unwrap(), first);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The defect this guards: `1 - cdf` collapsed to zero for a surprising
|
||||||
|
/// result, the clamp turned that into `f64::MIN_POSITIVE`, and
|
||||||
|
/// `log_evidence` reported ln of *that* — about -708 whatever the truth
|
||||||
|
/// was. An upset is the observation a model-comparison score exists to
|
||||||
|
/// notice, so it was wrong exactly where it mattered.
|
||||||
|
#[test]
|
||||||
|
fn evidence_of_an_upset_is_not_flattened_to_the_clamp_floor() {
|
||||||
|
// diff ~ N(-9, 1) with margin 0: the favoured side lost by nine sigma.
|
||||||
|
let evidence = cavity_log_evidence(Gaussian::from_ms(-9.0, 1.0), 0.0, false).exp();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
evidence > f64::MIN_POSITIVE,
|
||||||
|
"evidence collapsed onto the clamp floor: {evidence}"
|
||||||
|
);
|
||||||
|
// P(X > 0) for X ~ N(-9, 1) is the standard normal tail at 9 sigma.
|
||||||
|
assert!(
|
||||||
|
(evidence - 1.128_588e-19).abs() / 1.128_588e-19 < 1e-6,
|
||||||
|
"expected ~1.13e-19, got {evidence}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(evidence.ln() + 43.628).abs() < 1e-2,
|
||||||
|
"log evidence {} should be about -43.6, not -708",
|
||||||
|
evidence.ln()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Evidence must stay finite and positive however extreme the mismatch,
|
||||||
|
/// since `log_evidence` sums across the whole history and one `-inf` or
|
||||||
|
/// `NaN` poisons all of it.
|
||||||
|
///
|
||||||
|
/// Finiteness alone is too weak a bar — the clamped version was finite too,
|
||||||
|
/// and wrong by hundreds of nats. `log_evidence_tracks_the_analytic_tail`
|
||||||
|
/// below is the assertion that actually holds this up.
|
||||||
|
#[test]
|
||||||
|
fn evidence_stays_positive_and_finite_at_any_separation() {
|
||||||
|
for mu in [-300.0f64, -50.0, -9.0, 0.0, 9.0, 50.0, 300.0] {
|
||||||
|
for tie in [false, true] {
|
||||||
|
let ln_e = cavity_log_evidence(Gaussian::from_ms(mu, 1.0), 1.0, tie);
|
||||||
|
assert!(
|
||||||
|
ln_e.is_finite() && ln_e <= 0.0,
|
||||||
|
"mu={mu} tie={tie}: log evidence {ln_e} is not a log-probability"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The clamp used to floor everything past ~38 sigma at `ln(MIN_POSITIVE)`
|
||||||
|
/// = -708, however far out the real observation was. In log space the
|
||||||
|
/// answer is a polynomial and stays exact: at 1000 sigma the truth is about
|
||||||
|
/// -500_000 nats, and -708 is not a rounding error.
|
||||||
|
#[test]
|
||||||
|
fn log_evidence_tracks_the_analytic_tail() {
|
||||||
|
for mu in [-40.0f64, -60.0, -100.0, -1000.0] {
|
||||||
|
// P(diff > 0) for diff ~ N(mu, 1), mu far below zero.
|
||||||
|
let got = cavity_log_evidence(Gaussian::from_ms(mu, 1.0), 0.0, false);
|
||||||
|
|
||||||
|
// ln Phi(mu) ~ -mu^2/2 - ln(-mu) - ln(sqrt(2 pi)) for mu << 0.
|
||||||
|
let z = -mu;
|
||||||
|
let approx = -0.5 * z * z - z.ln() - (2.0 * std::f64::consts::PI).sqrt().ln();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
got < f64::MIN_POSITIVE.ln(),
|
||||||
|
"mu={mu}: {got} is still stuck on the old clamp floor"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(got - approx).abs() / approx.abs() < 1e-3,
|
||||||
|
"mu={mu}: got {got}, asymptotic expectation {approx}"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -132,7 +222,7 @@ mod tests {
|
|||||||
f.propagate(&mut vars);
|
f.propagate(&mut vars);
|
||||||
|
|
||||||
// For diff ~ N(0, 4), tie=true with margin=1: P(-1 < diff < 1) ≈ 0.383
|
// For diff ~ N(0, 4), tie=true with margin=1: P(-1 < diff < 1) ≈ 0.383
|
||||||
let ev = f.evidence_cached.unwrap();
|
let ev = f.log_evidence_cached.unwrap().exp();
|
||||||
assert!(ev > 0.35 && ev < 0.42);
|
assert!(ev > 0.35 && ev < 0.42);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +0,0 @@
|
|||||||
//! Factor-graph public API.
|
|
||||||
//!
|
|
||||||
//! Power users can construct custom factor graphs via `Game::custom` (T2
|
|
||||||
//! minimal; full ergonomics in T4) and drive them with custom `Schedule`
|
|
||||||
//! implementations.
|
|
||||||
|
|
||||||
pub use crate::{
|
|
||||||
factor::{
|
|
||||||
BuiltinFactor, Factor, VarId, VarStore, margin::MarginFactor, rank_diff::RankDiffFactor,
|
|
||||||
team_sum::TeamSumFactor, trunc::TruncFactor,
|
|
||||||
},
|
|
||||||
schedule::{EpsilonOrMax, Schedule, ScheduleReport},
|
|
||||||
};
|
|
||||||
+121
-73
@@ -37,10 +37,17 @@ impl DiffFactor {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn evidence(&self) -> f64 {
|
/// Log of this link's cached evidence.
|
||||||
|
///
|
||||||
|
/// Accumulating in log space keeps a long diff chain from underflowing:
|
||||||
|
/// each link contributes a probability in `(0, 1]`, so the linear product
|
||||||
|
/// over an n-team game decays geometrically and flushes to zero — and
|
||||||
|
/// `ln(0.0)` is `-inf` — well within the team counts a large free-for-all
|
||||||
|
/// reaches.
|
||||||
|
pub(crate) fn log_evidence(&self) -> f64 {
|
||||||
match self {
|
match self {
|
||||||
Self::Trunc(f) => f.evidence_cached.unwrap_or(1.0),
|
Self::Trunc(f) => f.log_evidence_cached.unwrap_or(0.0),
|
||||||
Self::Margin(f) => f.evidence_cached.unwrap_or(1.0),
|
Self::Margin(f) => f.log_evidence_cached.unwrap_or(0.0),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -81,18 +88,14 @@ impl Default for GameOptions {
|
|||||||
/// Owned variant of `Game` returned by public constructors.
|
/// Owned variant of `Game` returned by public constructors.
|
||||||
///
|
///
|
||||||
/// Unlike `Game<'a, T, D>` (which borrows its result/weights slices from
|
/// Unlike `Game<'a, T, D>` (which borrows its result/weights slices from
|
||||||
/// History's internal state), `OwnedGame<T, D>` owns its inputs so it can
|
/// History's internal state), `OwnedGame<T, D>` owns the team ratings, so it
|
||||||
/// be returned freely from public constructors.
|
/// can be returned freely from public constructors. The inference inputs
|
||||||
|
/// themselves are not retained — nothing reads them back.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
#[allow(dead_code)]
|
|
||||||
pub struct OwnedGame<T: Time, D: Drift<T>> {
|
pub struct OwnedGame<T: Time, D: Drift<T>> {
|
||||||
teams: Vec<Vec<Rating<T, D>>>,
|
teams: Vec<Vec<Rating<T, D>>>,
|
||||||
result: Vec<f64>,
|
|
||||||
weights: Vec<Vec<f64>>,
|
|
||||||
p_draw: f64,
|
|
||||||
pub(crate) convergence: crate::ConvergenceOptions,
|
|
||||||
pub(crate) likelihoods: Vec<Vec<Gaussian>>,
|
pub(crate) likelihoods: Vec<Vec<Gaussian>>,
|
||||||
pub(crate) evidence: f64,
|
pub(crate) log_evidence: f64,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T: Time, D: Drift<T>> OwnedGame<T, D> {
|
impl<T: Time, D: Drift<T>> OwnedGame<T, D> {
|
||||||
@@ -104,24 +107,15 @@ impl<T: Time, D: Drift<T>> OwnedGame<T, D> {
|
|||||||
convergence: crate::ConvergenceOptions,
|
convergence: crate::ConvergenceOptions,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let mut arena = ScratchArena::new();
|
let mut arena = ScratchArena::new();
|
||||||
let g = Game::ranked_with_arena(
|
|
||||||
teams.clone(),
|
// `Game` takes the teams by value and is dropped here, so take the vec
|
||||||
&result,
|
// back out of it rather than handing it a clone.
|
||||||
&weights,
|
let g = Game::ranked_with_arena(teams, &result, &weights, p_draw, convergence, &mut arena);
|
||||||
p_draw,
|
|
||||||
convergence,
|
|
||||||
&mut arena,
|
|
||||||
);
|
|
||||||
let likelihoods = g.likelihoods;
|
|
||||||
let evidence = g.evidence;
|
|
||||||
Self {
|
Self {
|
||||||
teams,
|
teams: g.teams,
|
||||||
result,
|
likelihoods: g.likelihoods,
|
||||||
weights,
|
log_evidence: g.log_evidence,
|
||||||
p_draw,
|
|
||||||
convergence,
|
|
||||||
likelihoods,
|
|
||||||
evidence,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -133,27 +127,24 @@ impl<T: Time, D: Drift<T>> OwnedGame<T, D> {
|
|||||||
convergence: crate::ConvergenceOptions,
|
convergence: crate::ConvergenceOptions,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let mut arena = ScratchArena::new();
|
let mut arena = ScratchArena::new();
|
||||||
|
|
||||||
let g = Game::scored_with_arena(
|
let g = Game::scored_with_arena(
|
||||||
teams.clone(),
|
teams,
|
||||||
&scores,
|
&scores,
|
||||||
&weights,
|
&weights,
|
||||||
score_sigma,
|
score_sigma,
|
||||||
convergence,
|
convergence,
|
||||||
&mut arena,
|
&mut arena,
|
||||||
);
|
);
|
||||||
let likelihoods = g.likelihoods;
|
|
||||||
let evidence = g.evidence;
|
|
||||||
Self {
|
Self {
|
||||||
teams,
|
teams: g.teams,
|
||||||
result: scores,
|
likelihoods: g.likelihoods,
|
||||||
weights,
|
log_evidence: g.log_evidence,
|
||||||
p_draw: 0.0,
|
|
||||||
convergence,
|
|
||||||
likelihoods,
|
|
||||||
evidence,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn posteriors(&self) -> Vec<Vec<Gaussian>> {
|
pub fn posteriors(&self) -> Vec<Vec<Gaussian>> {
|
||||||
self.likelihoods
|
self.likelihoods
|
||||||
.iter()
|
.iter()
|
||||||
@@ -162,8 +153,9 @@ impl<T: Time, D: Drift<T>> OwnedGame<T, D> {
|
|||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn log_evidence(&self) -> f64 {
|
pub fn log_evidence(&self) -> f64 {
|
||||||
self.evidence.ln()
|
self.log_evidence
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -175,7 +167,7 @@ pub struct Game<'a, T: Time = i64, D: Drift<T> = crate::drift::ConstantDrift> {
|
|||||||
p_draw: f64,
|
p_draw: f64,
|
||||||
pub(crate) convergence: crate::ConvergenceOptions,
|
pub(crate) convergence: crate::ConvergenceOptions,
|
||||||
pub(crate) likelihoods: Vec<Vec<Gaussian>>,
|
pub(crate) likelihoods: Vec<Vec<Gaussian>>,
|
||||||
pub(crate) evidence: f64,
|
pub(crate) log_evidence: f64,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
||||||
@@ -222,7 +214,7 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
p_draw,
|
p_draw,
|
||||||
convergence,
|
convergence,
|
||||||
likelihoods: Vec::new(),
|
likelihoods: Vec::new(),
|
||||||
evidence: 0.0,
|
log_evidence: 0.0,
|
||||||
};
|
};
|
||||||
|
|
||||||
this.likelihoods(arena);
|
this.likelihoods(arena);
|
||||||
@@ -261,7 +253,7 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
p_draw: 0.0,
|
p_draw: 0.0,
|
||||||
convergence,
|
convergence,
|
||||||
likelihoods: Vec::new(),
|
likelihoods: Vec::new(),
|
||||||
evidence: 0.0,
|
log_evidence: 0.0,
|
||||||
};
|
};
|
||||||
|
|
||||||
this.likelihoods_scored(arena, score_sigma);
|
this.likelihoods_scored(arena, score_sigma);
|
||||||
@@ -355,7 +347,7 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
arena.lhood_lose[n_teams - 1] = pw_last - links[n_diffs - 1].msg();
|
arena.lhood_lose[n_teams - 1] = pw_last - links[n_diffs - 1].msg();
|
||||||
}
|
}
|
||||||
|
|
||||||
let evidence: f64 = links.iter().map(|l| l.evidence()).product();
|
let log_evidence: f64 = links.iter().map(DiffFactor::log_evidence).sum();
|
||||||
|
|
||||||
// Inverse permutation: inv_buf[orig_i] = sorted_i.
|
// Inverse permutation: inv_buf[orig_i] = sorted_i.
|
||||||
arena.inv_buf.resize(n_teams, 0);
|
arena.inv_buf.resize(n_teams, 0);
|
||||||
@@ -371,10 +363,9 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
.map(|(orig_i, (players, weights))| {
|
.map(|(orig_i, (players, weights))| {
|
||||||
let si = arena.inv_buf[orig_i];
|
let si = arena.inv_buf[orig_i];
|
||||||
let m = arena.lhood_win[si] * arena.lhood_lose[si];
|
let m = arena.lhood_win[si] * arena.lhood_lose[si];
|
||||||
let performance = players
|
// Already folded into `team_prior` at the top of the chain,
|
||||||
.iter()
|
// indexed by sorted position.
|
||||||
.zip(weights.iter())
|
let performance = arena.team_prior[si];
|
||||||
.fold(N00, |p, (player, &w)| p + (player.performance() * w));
|
|
||||||
players
|
players
|
||||||
.iter()
|
.iter()
|
||||||
.zip(weights.iter())
|
.zip(weights.iter())
|
||||||
@@ -386,11 +377,11 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
|
|
||||||
(evidence, likelihoods)
|
(log_evidence, likelihoods)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn likelihoods(&mut self, arena: &mut ScratchArena) {
|
fn likelihoods(&mut self, arena: &mut ScratchArena) {
|
||||||
let (evidence, likelihoods) = self.run_chain(arena, |i, sort_buf, vars| {
|
let (log_evidence, likelihoods) = self.run_chain(arena, |i, sort_buf, vars| {
|
||||||
let tie = self.result[sort_buf[i]] == self.result[sort_buf[i + 1]];
|
let tie = self.result[sort_buf[i]] == self.result[sort_buf[i + 1]];
|
||||||
let margin = if self.p_draw == 0.0 {
|
let margin = if self.p_draw == 0.0 {
|
||||||
0.0
|
0.0
|
||||||
@@ -405,20 +396,21 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
let vid = vars.alloc(N_INF);
|
let vid = vars.alloc(N_INF);
|
||||||
DiffFactor::Trunc(TruncFactor::new(vid, margin, tie))
|
DiffFactor::Trunc(TruncFactor::new(vid, margin, tie))
|
||||||
});
|
});
|
||||||
self.evidence = evidence;
|
self.log_evidence = log_evidence;
|
||||||
self.likelihoods = likelihoods;
|
self.likelihoods = likelihoods;
|
||||||
}
|
}
|
||||||
|
|
||||||
fn likelihoods_scored(&mut self, arena: &mut ScratchArena, score_sigma: f64) {
|
fn likelihoods_scored(&mut self, arena: &mut ScratchArena, score_sigma: f64) {
|
||||||
let (evidence, likelihoods) = self.run_chain(arena, |i, sort_buf, vars| {
|
let (log_evidence, likelihoods) = self.run_chain(arena, |i, sort_buf, vars| {
|
||||||
let m_obs = self.result[sort_buf[i]] - self.result[sort_buf[i + 1]];
|
let m_obs = self.result[sort_buf[i]] - self.result[sort_buf[i + 1]];
|
||||||
let vid = vars.alloc(N_INF);
|
let vid = vars.alloc(N_INF);
|
||||||
DiffFactor::Margin(MarginFactor::new(vid, m_obs, score_sigma))
|
DiffFactor::Margin(MarginFactor::new(vid, m_obs, score_sigma))
|
||||||
});
|
});
|
||||||
self.evidence = evidence;
|
self.log_evidence = log_evidence;
|
||||||
self.likelihoods = likelihoods;
|
self.likelihoods = likelihoods;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn posteriors(&self) -> Vec<Vec<Gaussian>> {
|
pub fn posteriors(&self) -> Vec<Vec<Gaussian>> {
|
||||||
self.likelihoods
|
self.likelihoods
|
||||||
.iter()
|
.iter()
|
||||||
@@ -432,17 +424,30 @@ impl<'a, T: Time, D: Drift<T>> Game<'a, T, D> {
|
|||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn log_evidence(&self) -> f64 {
|
pub fn log_evidence(&self) -> f64 {
|
||||||
self.evidence.ln()
|
self.log_evidence
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// - `InvalidParameter` if `options.convergence` is out of range — an
|
||||||
|
/// `alpha` of zero would leave every EP update unapplied and silently
|
||||||
|
/// return the priors.
|
||||||
|
/// - `InvalidProbability` if `options.p_draw` is outside `[0.0, 1.0)`.
|
||||||
|
/// - `MismatchedShape` if the outcome's rank count differs from `teams.len()`.
|
||||||
|
/// - `WrongOutcomeKind` if `outcome` is not `Outcome::Ranked`.
|
||||||
|
/// - `TieWithoutDrawProbability` if the outcome ties two teams while
|
||||||
|
/// `p_draw` is zero: the truncation margin is then zero and the two-sided
|
||||||
|
/// tie update evaluates `0/0`.
|
||||||
pub fn ranked(
|
pub fn ranked(
|
||||||
teams: &[&[Rating<T, D>]],
|
teams: &[&[Rating<T, D>]],
|
||||||
outcome: crate::Outcome,
|
outcome: crate::Outcome,
|
||||||
options: &GameOptions,
|
options: &GameOptions,
|
||||||
) -> Result<OwnedGame<T, D>, crate::InferenceError> {
|
) -> Result<OwnedGame<T, D>, crate::InferenceError> {
|
||||||
|
options.convergence.validate()?;
|
||||||
if !(0.0..1.0).contains(&options.p_draw) {
|
if !(0.0..1.0).contains(&options.p_draw) {
|
||||||
return Err(crate::InferenceError::InvalidProbability {
|
return Err(crate::InferenceError::InvalidProbability {
|
||||||
value: options.p_draw,
|
value: options.p_draw,
|
||||||
@@ -458,11 +463,22 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
|
|
||||||
let ranks = outcome
|
let ranks = outcome
|
||||||
.as_ranks()
|
.as_ranks()
|
||||||
.ok_or(crate::InferenceError::MismatchedShape {
|
.ok_or(crate::InferenceError::WrongOutcomeKind {
|
||||||
kind: "Game::ranked requires Outcome::Ranked",
|
context: "Game::ranked",
|
||||||
expected: 0,
|
expected: "Outcome::Ranked",
|
||||||
got: 0,
|
got: "Outcome::Scored",
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
|
let tied = if options.p_draw == 0.0 {
|
||||||
|
crate::first_tied_pair(ranks)
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
|
||||||
|
if let Some(teams) = tied {
|
||||||
|
return Err(crate::InferenceError::TieWithoutDrawProbability { teams });
|
||||||
|
}
|
||||||
|
|
||||||
let max_rank = ranks.iter().copied().max().unwrap_or(0) as f64;
|
let max_rank = ranks.iter().copied().max().unwrap_or(0) as f64;
|
||||||
let result: Vec<f64> = ranks.iter().map(|&r| max_rank - r as f64).collect();
|
let result: Vec<f64> = ranks.iter().map(|&r| max_rank - r as f64).collect();
|
||||||
let teams_owned: Vec<Vec<Rating<T, D>>> = teams.iter().map(|t| t.to_vec()).collect();
|
let teams_owned: Vec<Vec<Rating<T, D>>> = teams.iter().map(|t| t.to_vec()).collect();
|
||||||
@@ -477,11 +493,18 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// - `InvalidParameter` if `options.score_sigma` is not strictly positive
|
||||||
|
/// or is NaN, or if `options.convergence` is out of range.
|
||||||
|
/// - `MismatchedShape` if the outcome's score count differs from `teams.len()`.
|
||||||
|
/// - `WrongOutcomeKind` if `outcome` is not `Outcome::Scored`.
|
||||||
pub fn scored(
|
pub fn scored(
|
||||||
teams: &[&[Rating<T, D>]],
|
teams: &[&[Rating<T, D>]],
|
||||||
outcome: crate::Outcome,
|
outcome: crate::Outcome,
|
||||||
options: &GameOptions,
|
options: &GameOptions,
|
||||||
) -> Result<OwnedGame<T, D>, crate::InferenceError> {
|
) -> Result<OwnedGame<T, D>, crate::InferenceError> {
|
||||||
|
options.convergence.validate()?;
|
||||||
if options.score_sigma <= 0.0 || options.score_sigma.is_nan() {
|
if options.score_sigma <= 0.0 || options.score_sigma.is_nan() {
|
||||||
return Err(crate::InferenceError::InvalidParameter {
|
return Err(crate::InferenceError::InvalidParameter {
|
||||||
name: "score_sigma",
|
name: "score_sigma",
|
||||||
@@ -497,10 +520,10 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
}
|
}
|
||||||
let scores = outcome
|
let scores = outcome
|
||||||
.as_scores()
|
.as_scores()
|
||||||
.ok_or(crate::InferenceError::MismatchedShape {
|
.ok_or(crate::InferenceError::WrongOutcomeKind {
|
||||||
kind: "Game::scored requires Outcome::Scored",
|
context: "Game::scored",
|
||||||
expected: 0,
|
expected: "Outcome::Scored",
|
||||||
got: 0,
|
got: "Outcome::Ranked",
|
||||||
})?
|
})?
|
||||||
.to_vec();
|
.to_vec();
|
||||||
let teams_owned: Vec<Vec<Rating<T, D>>> = teams.iter().map(|t| t.to_vec()).collect();
|
let teams_owned: Vec<Vec<Rating<T, D>>> = teams.iter().map(|t| t.to_vec()).collect();
|
||||||
@@ -514,16 +537,28 @@ impl<T: Time, D: Drift<T>> Game<'_, T, D> {
|
|||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Convenience wrapper over [`Game::ranked`] for two single-player teams.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// Delegates to [`Game::ranked`], so it returns the same errors — in
|
||||||
|
/// practice `WrongOutcomeKind` for a non-ranked outcome, or
|
||||||
|
/// `TieWithoutDrawProbability` for a draw when `options.p_draw` is zero.
|
||||||
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]))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// Wraps each player in a one-member team and delegates to
|
||||||
|
/// [`Game::ranked`], so it returns the same errors.
|
||||||
pub fn free_for_all(
|
pub fn free_for_all(
|
||||||
players: &[&Rating<T, D>],
|
players: &[&Rating<T, D>],
|
||||||
outcome: crate::Outcome,
|
outcome: crate::Outcome,
|
||||||
@@ -535,11 +570,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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -698,9 +733,15 @@ mod tests {
|
|||||||
let c = p[2][0];
|
let c = p[2][0];
|
||||||
|
|
||||||
// T1 ULP shift: mu rounds to 25.0 (was 24.999999) under natural-parameter storage.
|
// T1 ULP shift: mu rounds to 25.0 (was 24.999999) under natural-parameter storage.
|
||||||
|
//
|
||||||
|
// The 1e-6-place values moved when `erfc_inv`'s sign error was fixed:
|
||||||
|
// this case runs at `p_draw = 0.5`, so it goes through `compute_margin`,
|
||||||
|
// and the margin is now 8.4e-8 from the exact quantile where it was
|
||||||
|
// 1.46e-7. Verified as movement *toward* analytic truth, not a
|
||||||
|
// regression — see `erfc_inv_matches_known_quantiles`.
|
||||||
assert_ulps_eq!(a, Gaussian::from_ms(25.0, 6.092561), epsilon = 1e-6);
|
assert_ulps_eq!(a, Gaussian::from_ms(25.0, 6.092561), epsilon = 1e-6);
|
||||||
assert_ulps_eq!(b, Gaussian::from_ms(33.379314, 6.483575), epsilon = 1e-6);
|
assert_ulps_eq!(b, Gaussian::from_ms(33.379315, 6.483576), epsilon = 1e-6);
|
||||||
assert_ulps_eq!(c, Gaussian::from_ms(16.620685, 6.483575), epsilon = 1e-6);
|
assert_ulps_eq!(c, Gaussian::from_ms(16.620685, 6.483576), epsilon = 1e-6);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -730,8 +771,12 @@ mod tests {
|
|||||||
let a = p[0][0];
|
let a = p[0][0];
|
||||||
let b = p[1][0];
|
let b = p[1][0];
|
||||||
|
|
||||||
assert_ulps_eq!(a, Gaussian::from_ms(24.999999, 6.469480), epsilon = 1e-6);
|
// Two identical competitors drawing must land on their shared prior
|
||||||
assert_ulps_eq!(b, Gaussian::from_ms(24.999999, 6.469480), epsilon = 1e-6);
|
// mean exactly, by symmetry. The reference transcription of 24.999999
|
||||||
|
// is that value rounded to six decimals; asserting it at epsilon 1e-6
|
||||||
|
// left no headroom. The root-free variance path now hits 25.0 exactly.
|
||||||
|
assert_ulps_eq!(a, Gaussian::from_ms(25.0, 6.469480), epsilon = 1e-6);
|
||||||
|
assert_ulps_eq!(b, Gaussian::from_ms(25.0, 6.469480), epsilon = 1e-6);
|
||||||
|
|
||||||
let t_a = R::new(
|
let t_a = R::new(
|
||||||
Gaussian::from_ms(25.0, 3.0),
|
Gaussian::from_ms(25.0, 3.0),
|
||||||
@@ -1124,7 +1169,10 @@ mod tests {
|
|||||||
&GameOptions::default(),
|
&GameOptions::default(),
|
||||||
)
|
)
|
||||||
.unwrap_err();
|
.unwrap_err();
|
||||||
assert!(matches!(err, crate::InferenceError::MismatchedShape { .. }));
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
crate::InferenceError::WrongOutcomeKind { .. }
|
||||||
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -1206,7 +1254,7 @@ mod tests {
|
|||||||
);
|
);
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
p[1][0],
|
p[1][0],
|
||||||
Gaussian::from_ms(19.287197, 7.243465),
|
Gaussian::from_ms(19.287198285, 7.243465848),
|
||||||
epsilon = 1e-6
|
epsilon = 1e-6
|
||||||
);
|
);
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
@@ -1266,7 +1314,7 @@ mod tests {
|
|||||||
|
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
p[0][0],
|
p[0][0],
|
||||||
Gaussian::from_ms(31.674697, 7.501180),
|
Gaussian::from_ms(31.674698083, 7.501180037),
|
||||||
epsilon = 1e-6
|
epsilon = 1e-6
|
||||||
);
|
);
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
|
|||||||
+51
-13
@@ -18,6 +18,7 @@ pub struct Gaussian {
|
|||||||
|
|
||||||
impl Gaussian {
|
impl Gaussian {
|
||||||
/// Construct from mean and standard deviation.
|
/// Construct from mean and standard deviation.
|
||||||
|
#[must_use]
|
||||||
pub const fn from_ms(mu: f64, sigma: f64) -> Self {
|
pub const fn from_ms(mu: f64, sigma: f64) -> Self {
|
||||||
if sigma == f64::INFINITY {
|
if sigma == f64::INFINITY {
|
||||||
Self { pi: 0.0, tau: 0.0 }
|
Self { pi: 0.0, tau: 0.0 }
|
||||||
@@ -35,6 +36,28 @@ impl Gaussian {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Construct from mean and *variance*, skipping the square-root round trip.
|
||||||
|
///
|
||||||
|
/// `from_ms(mu, var.sqrt())` immediately squares the root away again to
|
||||||
|
/// recover `pi = 1/var`. Variance-combining operations (`Add`, `Sub`,
|
||||||
|
/// `exclude`, `forget`) work in variance space throughout, so they go
|
||||||
|
/// through here instead and never take a root.
|
||||||
|
#[inline]
|
||||||
|
pub(crate) fn from_mv(mu: f64, var: f64) -> Self {
|
||||||
|
if var == f64::INFINITY {
|
||||||
|
Self { pi: 0.0, tau: 0.0 }
|
||||||
|
} else if var == 0.0 {
|
||||||
|
// Point mass at mu; see `from_ms` for the tau convention.
|
||||||
|
Self {
|
||||||
|
pi: f64::INFINITY,
|
||||||
|
tau: if mu == 0.0 { 0.0 } else { f64::INFINITY },
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
let pi = 1.0 / var;
|
||||||
|
Self { pi, tau: mu * pi }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Construct directly from natural parameters.
|
/// Construct directly from natural parameters.
|
||||||
#[inline]
|
#[inline]
|
||||||
pub(crate) const fn from_natural(pi: f64, tau: f64) -> Self {
|
pub(crate) const fn from_natural(pi: f64, tau: f64) -> Self {
|
||||||
@@ -42,16 +65,19 @@ impl Gaussian {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[inline]
|
#[inline]
|
||||||
|
#[must_use]
|
||||||
pub fn pi(&self) -> f64 {
|
pub fn pi(&self) -> f64 {
|
||||||
self.pi
|
self.pi
|
||||||
}
|
}
|
||||||
|
|
||||||
#[inline]
|
#[inline]
|
||||||
|
#[must_use]
|
||||||
pub fn tau(&self) -> f64 {
|
pub fn tau(&self) -> f64 {
|
||||||
self.tau
|
self.tau
|
||||||
}
|
}
|
||||||
|
|
||||||
#[inline]
|
#[inline]
|
||||||
|
#[must_use]
|
||||||
pub fn mu(&self) -> f64 {
|
pub fn mu(&self) -> f64 {
|
||||||
// A non-positive precision is an improper (uninformative) Gaussian — its mean is
|
// A non-positive precision is an improper (uninformative) Gaussian — its mean is
|
||||||
// undefined. Treat it like `pi == 0` and return 0. EP message cancellation can land
|
// undefined. Treat it like `pi == 0` and return 0. EP message cancellation can land
|
||||||
@@ -64,7 +90,23 @@ impl Gaussian {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Variance, `1 / pi`, without the root-and-square of `sigma().powi(2)`.
|
||||||
|
///
|
||||||
|
/// Mirrors `sigma()`'s treatment of the improper (`pi <= 0`) and point-mass
|
||||||
|
/// (`pi == inf`) cases.
|
||||||
#[inline]
|
#[inline]
|
||||||
|
pub(crate) fn variance(&self) -> f64 {
|
||||||
|
if self.pi <= 0.0 {
|
||||||
|
f64::INFINITY
|
||||||
|
} else if self.pi.is_infinite() {
|
||||||
|
0.0
|
||||||
|
} else {
|
||||||
|
1.0 / self.pi
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[inline]
|
||||||
|
#[must_use]
|
||||||
pub fn sigma(&self) -> f64 {
|
pub fn sigma(&self) -> f64 {
|
||||||
// A non-positive precision is improper → infinite standard deviation. Guarding
|
// A non-positive precision is improper → infinite standard deviation. Guarding
|
||||||
// `pi <= 0.0` (not just `== 0.0`) keeps `1.0 / pi.sqrt()` from returning NaN when EP
|
// `pi <= 0.0` (not just `== 0.0`) keeps `1.0 / pi.sqrt()` from returning NaN when EP
|
||||||
@@ -86,22 +128,21 @@ impl Gaussian {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn exclude(&self, other: Gaussian) -> Self {
|
pub(crate) fn exclude(&self, other: Gaussian) -> Self {
|
||||||
let var = self.sigma().powi(2) - other.sigma().powi(2);
|
let var = self.variance() - other.variance();
|
||||||
if var <= 0.0 {
|
if var <= 0.0 {
|
||||||
// When sigma_self ≈ sigma_other (including ULP-level rounding differences
|
// When sigma_self ≈ sigma_other (including ULP-level rounding differences
|
||||||
// from the pi→sigma accessor round-trip), the excluded contribution is N00.
|
// from the pi→sigma accessor round-trip), the excluded contribution is N00.
|
||||||
// Computing from_ms(tiny_mu, 0.0) would give {pi:inf, tau:inf}, whose
|
// Computing from_ms(tiny_mu, 0.0) would give {pi:inf, tau:inf}, whose
|
||||||
// mu() = inf/inf = NaN. Returning N00 is correct: when both Gaussians
|
// mu() = inf/inf = NaN. Returning N00 is correct: when both Gaussians
|
||||||
// carry the same variance, the residual is a point mass at 0.
|
// carry the same variance, the residual is a point mass at 0.
|
||||||
return Gaussian::from_ms(0.0, 0.0);
|
return Gaussian::from_mv(0.0, 0.0);
|
||||||
}
|
}
|
||||||
let mu = self.mu() - other.mu();
|
|
||||||
Self::from_ms(mu, var.sqrt())
|
Self::from_mv(self.mu() - other.mu(), var)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn forget(&self, variance_delta: f64) -> Self {
|
pub(crate) fn forget(&self, variance_delta: f64) -> Self {
|
||||||
let var = self.sigma().powi(2) + variance_delta;
|
Self::from_mv(self.mu(), self.variance() + variance_delta)
|
||||||
Self::from_ms(self.mu(), var.sqrt())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// EP damping in natural-parameter space: `α·new + (1−α)·self`.
|
/// EP damping in natural-parameter space: `α·new + (1−α)·self`.
|
||||||
@@ -109,6 +150,7 @@ impl Gaussian {
|
|||||||
/// Used by within-game inference to stabilise oscillating fixed-point
|
/// Used by within-game inference to stabilise oscillating fixed-point
|
||||||
/// loops on hard graphs. `alpha = 1.0` returns `new` exactly;
|
/// loops on hard graphs. `alpha = 1.0` returns `new` exactly;
|
||||||
/// `alpha < 1.0` shrinks each per-step update.
|
/// `alpha < 1.0` shrinks each per-step update.
|
||||||
|
#[must_use]
|
||||||
pub fn damp_natural(self, new: Gaussian, alpha: f64) -> Gaussian {
|
pub fn damp_natural(self, new: Gaussian, alpha: f64) -> Gaussian {
|
||||||
Gaussian::from_natural(
|
Gaussian::from_natural(
|
||||||
alpha * new.pi() + (1.0 - alpha) * self.pi(),
|
alpha * new.pi() + (1.0 - alpha) * self.pi(),
|
||||||
@@ -128,9 +170,7 @@ impl ops::Add<Gaussian> for Gaussian {
|
|||||||
/// Variance addition: (mu1 + mu2, sqrt(σ1² + σ2²)).
|
/// Variance addition: (mu1 + mu2, sqrt(σ1² + σ2²)).
|
||||||
/// Used for combining performance and noise; rare relative to mul/div.
|
/// Used for combining performance and noise; rare relative to mul/div.
|
||||||
fn add(self, rhs: Gaussian) -> Self::Output {
|
fn add(self, rhs: Gaussian) -> Self::Output {
|
||||||
let mu = self.mu() + rhs.mu();
|
Self::from_mv(self.mu() + rhs.mu(), self.variance() + rhs.variance())
|
||||||
let var = self.sigma().powi(2) + rhs.sigma().powi(2);
|
|
||||||
Self::from_ms(mu, var.sqrt())
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -138,9 +178,7 @@ impl ops::Sub<Gaussian> for Gaussian {
|
|||||||
type Output = Gaussian;
|
type Output = Gaussian;
|
||||||
/// (mu1 - mu2, sqrt(σ1² + σ2²)). Same sigma combination as Add.
|
/// (mu1 - mu2, sqrt(σ1² + σ2²)). Same sigma combination as Add.
|
||||||
fn sub(self, rhs: Gaussian) -> Self::Output {
|
fn sub(self, rhs: Gaussian) -> Self::Output {
|
||||||
let mu = self.mu() - rhs.mu();
|
Self::from_mv(self.mu() - rhs.mu(), self.variance() + rhs.variance())
|
||||||
let var = self.sigma().powi(2) + rhs.sigma().powi(2);
|
|
||||||
Self::from_ms(mu, var.sqrt())
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -161,7 +199,7 @@ impl ops::Mul<f64> for Gaussian {
|
|||||||
if scalar == 0.0 {
|
if scalar == 0.0 {
|
||||||
// Scaling by 0 collapses to a point mass at 0 (sigma' = 0, mu' = 0).
|
// Scaling by 0 collapses to a point mass at 0 (sigma' = 0, mu' = 0).
|
||||||
// This is N00, the additive identity, NOT N_INF.
|
// This is N00, the additive identity, NOT N_INF.
|
||||||
return Gaussian::from_ms(0.0, 0.0);
|
return Gaussian::from_mv(0.0, 0.0);
|
||||||
}
|
}
|
||||||
// sigma' = sigma * |scalar| => pi' = pi / scalar²
|
// sigma' = sigma * |scalar| => pi' = pi / scalar²
|
||||||
// mu' = mu * scalar => tau' = tau / scalar
|
// mu' = mu * scalar => tau' = tau / scalar
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
//! 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
|
||||||
|
//! schedules can be written against them.
|
||||||
|
//!
|
||||||
|
//! Building a factor graph by hand goes through `Game::custom`, which is
|
||||||
|
//! deliberately `#[doc(hidden)]`: it works, but its signature is not yet
|
||||||
|
//! considered stable API and so is not listed in these docs.
|
||||||
|
|
||||||
|
pub use crate::{
|
||||||
|
factor::{
|
||||||
|
BuiltinFactor, Factor, VarId, VarStore, margin::MarginFactor, rank_diff::RankDiffFactor,
|
||||||
|
team_sum::TeamSumFactor, trunc::TruncFactor,
|
||||||
|
},
|
||||||
|
schedule::{EpsilonOrMax, Schedule, ScheduleReport},
|
||||||
|
};
|
||||||
+737
-112
File diff suppressed because it is too large
Load Diff
+28
-15
@@ -12,59 +12,72 @@ use crate::Index;
|
|||||||
/// crate. Power users can promote `&K` to `Index` via `get_or_create` and
|
/// crate. Power users can promote `&K` to `Index` via `get_or_create` and
|
||||||
/// skip the lookup on subsequent hot-path calls.
|
/// skip the lookup on subsequent hot-path calls.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct KeyTable<K>(HashMap<K, Index>);
|
pub struct KeyTable<K> {
|
||||||
|
forward: HashMap<K, Index>,
|
||||||
|
/// Reverse mapping, indexed by `Index.0`.
|
||||||
|
///
|
||||||
|
/// Indices are handed out densely and sequentially, so position *is* the
|
||||||
|
/// index and `key()` is a lookup rather than a scan over every entry.
|
||||||
|
reverse: Vec<K>,
|
||||||
|
}
|
||||||
|
|
||||||
impl<K> KeyTable<K>
|
impl<K> KeyTable<K>
|
||||||
where
|
where
|
||||||
K: Eq + Hash,
|
K: Eq + Hash + Clone,
|
||||||
{
|
{
|
||||||
|
#[must_use]
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self(HashMap::new())
|
Self {
|
||||||
|
forward: HashMap::new(),
|
||||||
|
reverse: Vec::new(),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get<Q: ?Sized + Hash + Eq>(&self, k: &Q) -> Option<Index>
|
pub fn get<Q: ?Sized + Hash + Eq>(&self, k: &Q) -> Option<Index>
|
||||||
where
|
where
|
||||||
K: Borrow<Q>,
|
K: Borrow<Q>,
|
||||||
{
|
{
|
||||||
self.0.get(k).cloned()
|
self.forward.get(k).cloned()
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_or_create<Q: ?Sized + Hash + Eq + ToOwned<Owned = K>>(&mut self, k: &Q) -> Index
|
pub fn get_or_create<Q: ?Sized + Hash + Eq + ToOwned<Owned = K>>(&mut self, k: &Q) -> Index
|
||||||
where
|
where
|
||||||
K: Borrow<Q>,
|
K: Borrow<Q>,
|
||||||
{
|
{
|
||||||
if let Some(idx) = self.0.get(k) {
|
if let Some(idx) = self.forward.get(k) {
|
||||||
*idx
|
*idx
|
||||||
} else {
|
} else {
|
||||||
let idx = Index::from(self.0.len());
|
let idx = Index::from(self.reverse.len());
|
||||||
self.0.insert(k.to_owned(), idx);
|
let owned = k.to_owned();
|
||||||
|
self.reverse.push(owned.clone());
|
||||||
|
self.forward.insert(owned, idx);
|
||||||
idx
|
idx
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn key(&self, idx: Index) -> Option<&K> {
|
pub fn key(&self, idx: Index) -> Option<&K> {
|
||||||
self.0
|
self.reverse.get(idx.0)
|
||||||
.iter()
|
|
||||||
.find(|&(_, value)| *value == idx)
|
|
||||||
.map(|(key, _)| key)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn keys(&self) -> impl Iterator<Item = &K> {
|
pub fn keys(&self) -> impl Iterator<Item = &K> {
|
||||||
self.0.keys()
|
self.forward.keys()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.0.len()
|
self.reverse.len()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn is_empty(&self) -> bool {
|
pub fn is_empty(&self) -> bool {
|
||||||
self.0.is_empty()
|
self.reverse.is_empty()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<K> Default for KeyTable<K>
|
impl<K> Default for KeyTable<K>
|
||||||
where
|
where
|
||||||
K: Eq + Hash,
|
K: Eq + Hash + Clone,
|
||||||
{
|
{
|
||||||
fn default() -> Self {
|
fn default() -> Self {
|
||||||
KeyTable::new()
|
KeyTable::new()
|
||||||
|
|||||||
+772
-37
@@ -1,3 +1,104 @@
|
|||||||
|
//! `TrueSkill` Through Time — Bayesian skill rating over a time axis.
|
||||||
|
//!
|
||||||
|
//! Where plain `TrueSkill` gives each competitor one running estimate, `TrueSkill`
|
||||||
|
//! Through Time treats a whole history as a single model and infers skill *at
|
||||||
|
//! every point in time*. Evidence flows both directions: a result today
|
||||||
|
//! sharpens the estimate of who someone was last year, so early estimates stop
|
||||||
|
//! being frozen guesses and comparisons across eras become meaningful.
|
||||||
|
//!
|
||||||
|
//! This is a Rust port of
|
||||||
|
//! [TrueSkillThroughTime.py](https://github.com/glandfried/TrueSkillThroughTime.py).
|
||||||
|
//!
|
||||||
|
//! # Getting started
|
||||||
|
//!
|
||||||
|
//! Record results, converge, then read off skills:
|
||||||
|
//!
|
||||||
|
//! ```
|
||||||
|
//! use trueskill_tt::History;
|
||||||
|
//!
|
||||||
|
//! let mut history = History::default();
|
||||||
|
//!
|
||||||
|
//! history.record_winner(&"alice", &"bob", 1)?;
|
||||||
|
//! history.record_winner(&"bob", &"carol", 2)?;
|
||||||
|
//! history.record_winner(&"alice", &"carol", 3)?;
|
||||||
|
//!
|
||||||
|
//! let report = history.converge()?;
|
||||||
|
//! assert!(report.converged);
|
||||||
|
//!
|
||||||
|
//! let alice = history.current_skill("alice").unwrap();
|
||||||
|
//! assert!(alice.mu() > 0.0, "alice won every game she played");
|
||||||
|
//! # Ok::<(), trueskill_tt::InferenceError>(())
|
||||||
|
//! ```
|
||||||
|
//!
|
||||||
|
//! Teams, weights, explicit rankings and continuous scores go through the
|
||||||
|
//! fluent event builder:
|
||||||
|
//!
|
||||||
|
//! ```
|
||||||
|
//! use trueskill_tt::History;
|
||||||
|
//!
|
||||||
|
//! let mut history = History::builder().p_draw(0.1).build();
|
||||||
|
//!
|
||||||
|
//! history
|
||||||
|
//! .event(1)
|
||||||
|
//! .team(["alice", "bob"])
|
||||||
|
//! .team(["carol", "dave"])
|
||||||
|
//! .ranking([0, 1])
|
||||||
|
//! .commit()?;
|
||||||
|
//!
|
||||||
|
//! history.converge()?;
|
||||||
|
//! # Ok::<(), trueskill_tt::InferenceError>(())
|
||||||
|
//! ```
|
||||||
|
//!
|
||||||
|
//! # Draws need a draw probability
|
||||||
|
//!
|
||||||
|
//! A `p_draw` of zero asserts that draws cannot happen, so a tied result has
|
||||||
|
//! no representable likelihood and is rejected:
|
||||||
|
//!
|
||||||
|
//! ```
|
||||||
|
//! use trueskill_tt::{History, InferenceError};
|
||||||
|
//!
|
||||||
|
//! let mut history = History::default(); // p_draw defaults to 0.0
|
||||||
|
//! let err = history.record_draw(&"alice", &"bob", 1).unwrap_err();
|
||||||
|
//! assert!(matches!(err, InferenceError::TieWithoutDrawProbability { .. }));
|
||||||
|
//! ```
|
||||||
|
//!
|
||||||
|
//! This also applies to [`Outcome::winner`] for three or more teams, which
|
||||||
|
//! ties every loser. Configure a positive `p_draw` for those.
|
||||||
|
//!
|
||||||
|
//! # Core types
|
||||||
|
//!
|
||||||
|
//! - [`History`] — the top-level container: ingests events, runs
|
||||||
|
//! forward/backward message passing, and answers queries.
|
||||||
|
//! - [`Gaussian`] — the probability type, stored in natural parameters
|
||||||
|
//! (`pi = 1/sigma²`, `tau = mu/sigma²`) so message passing is add/subtract.
|
||||||
|
//! - [`Game`] — one match in isolation, for scoring a hypothetical without a
|
||||||
|
//! history.
|
||||||
|
//! - [`Outcome`] — how a match ended: ranks, or continuous scores.
|
||||||
|
//! - [`Rating`] — a competitor's static configuration (prior, `beta`, drift).
|
||||||
|
//!
|
||||||
|
//! # Feature flags
|
||||||
|
//!
|
||||||
|
//! - `approx` — implements [`approx`](https://docs.rs/approx) equality traits
|
||||||
|
//! for [`Gaussian`]. Useful in tests.
|
||||||
|
//! - `rayon` — parallelises the within-slice sweep and the per-slice passes of
|
||||||
|
//! `learning_curves`/`log_evidence`. Opt-in; results stay bit-identical
|
||||||
|
//! regardless of worker count.
|
||||||
|
|
||||||
|
#![forbid(unsafe_code)]
|
||||||
|
|
||||||
|
/// Compiles every `rust` block in `README.md` as a doctest.
|
||||||
|
///
|
||||||
|
/// The README is not the crate's front page — the module docs above are — so it
|
||||||
|
/// is pulled in here rather than via a crate-level `#![doc = ...]`, purely so
|
||||||
|
/// its examples are type-checked. Without this nothing compiled them, and they
|
||||||
|
/// had drifted far enough that four blocks no longer built (#35). `cfg(doctest)`
|
||||||
|
/// means this type exists only while collecting doctests.
|
||||||
|
///
|
||||||
|
/// Blocks that are illustrative rather than runnable are fenced as `text`.
|
||||||
|
#[cfg(doctest)]
|
||||||
|
#[doc = include_str!("../README.md")]
|
||||||
|
pub struct ReadmeDoctests;
|
||||||
|
|
||||||
use std::{
|
use std::{
|
||||||
cmp::Reverse,
|
cmp::Reverse,
|
||||||
f64::consts::{FRAC_1_SQRT_2, FRAC_2_SQRT_PI, SQRT_2},
|
f64::consts::{FRAC_1_SQRT_2, FRAC_2_SQRT_PI, SQRT_2},
|
||||||
@@ -9,6 +110,7 @@ pub(crate) mod arena;
|
|||||||
mod time;
|
mod time;
|
||||||
mod time_slice;
|
mod time_slice;
|
||||||
pub use time_slice::{EventKind, TimeSlice};
|
pub use time_slice::{EventKind, TimeSlice};
|
||||||
|
mod acquisition;
|
||||||
mod color_group;
|
mod color_group;
|
||||||
mod competitor;
|
mod competitor;
|
||||||
mod convergence;
|
mod convergence;
|
||||||
@@ -17,18 +119,21 @@ 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;
|
||||||
mod observer;
|
mod observer;
|
||||||
mod outcome;
|
mod outcome;
|
||||||
|
mod predict;
|
||||||
|
pub(crate) mod quadrature;
|
||||||
mod rating;
|
mod rating;
|
||||||
pub(crate) mod schedule;
|
pub(crate) mod schedule;
|
||||||
pub mod storage;
|
pub mod storage;
|
||||||
|
|
||||||
|
pub use acquisition::expected_information_gain;
|
||||||
pub use competitor::Competitor;
|
pub use competitor::Competitor;
|
||||||
pub use convergence::{ConvergenceOptions, ConvergenceReport};
|
pub use convergence::{ConvergenceOptions, ConvergenceReport};
|
||||||
pub use drift::{ConstantDrift, Drift};
|
pub use drift::{ConstantDrift, Drift};
|
||||||
@@ -37,11 +142,12 @@ pub use event::{Event, Member, Team};
|
|||||||
pub use event_builder::EventBuilder;
|
pub use event_builder::EventBuilder;
|
||||||
pub use game::{Game, GameOptions, OwnedGame};
|
pub use game::{Game, GameOptions, OwnedGame};
|
||||||
pub use gaussian::Gaussian;
|
pub use gaussian::Gaussian;
|
||||||
pub use history::History;
|
pub use history::{History, HistoryBuilder};
|
||||||
pub use key_table::KeyTable;
|
pub use key_table::KeyTable;
|
||||||
use matrix::Matrix;
|
use matrix::Matrix;
|
||||||
pub use observer::{NullObserver, Observer};
|
pub use observer::{NullObserver, Observer};
|
||||||
pub use outcome::Outcome;
|
pub use outcome::Outcome;
|
||||||
|
pub use predict::Prediction;
|
||||||
pub use rating::Rating;
|
pub use rating::Rating;
|
||||||
pub use schedule::ScheduleReport;
|
pub use schedule::ScheduleReport;
|
||||||
pub use time::{Time, Untimed};
|
pub use time::{Time, Untimed};
|
||||||
@@ -54,7 +160,29 @@ pub const P_DRAW: f64 = 0.0;
|
|||||||
pub const EPSILON: f64 = 1e-6;
|
pub const EPSILON: f64 = 1e-6;
|
||||||
pub const ITERATIONS: usize = 30;
|
pub const ITERATIONS: usize = 30;
|
||||||
|
|
||||||
|
/// Largest team count `History::predict_outcome` will enumerate.
|
||||||
|
///
|
||||||
|
/// The outcome space holds `n! * 2^(n-1)` events, so it grows factorially:
|
||||||
|
/// 1_920 at five teams, 23_040 at six, 322_560 at seven. Six is where
|
||||||
|
/// enumerating on a caller's behalf stops being reasonable.
|
||||||
|
pub const MAX_PREDICTED_TEAMS: usize = predict::MAX_TEAMS_FOR_DISTRIBUTION;
|
||||||
|
|
||||||
const SQRT_TAU: f64 = 2.5066282746310002;
|
const SQRT_TAU: f64 = 2.5066282746310002;
|
||||||
|
/// `1 / sqrt(pi)`, the leading factor of the `erfcx` continued fraction.
|
||||||
|
const FRAC_1_SQRT_PI: f64 = 0.564_189_583_547_756_3;
|
||||||
|
/// `sqrt(2 / pi)`, the numerator of the inverse Mills ratio in scaled form.
|
||||||
|
const SQRT_2_OVER_PI: f64 = 0.797_884_560_802_865_4;
|
||||||
|
/// How many window widths into the tail before a tie window is treated as a
|
||||||
|
/// half-line. Beyond this the truncated mass is concentrated within `1/alpha`
|
||||||
|
/// of the near edge, so the far edge contributes nothing measurable.
|
||||||
|
const HALF_LINE_WINDOW: f64 = 10.0;
|
||||||
|
/// Where `v - alpha` switches from subtraction to its asymptotic series.
|
||||||
|
///
|
||||||
|
/// The subtraction loses roughly `eps * alpha^2` of relative precision, and the
|
||||||
|
/// four-term series is good to ~1e-10 by here, so the two are at their closest
|
||||||
|
/// agreement around this point. Below it the subtraction is exact; above it the
|
||||||
|
/// series is.
|
||||||
|
const ASYMPTOTIC_MILLS_ALPHA: f64 = 100.0;
|
||||||
|
|
||||||
pub const N01: Gaussian = Gaussian::from_ms(0.0, 1.0);
|
pub const N01: Gaussian = Gaussian::from_ms(0.0, 1.0);
|
||||||
pub const N00: Gaussian = Gaussian::from_ms(0.0, 0.0);
|
pub const N00: Gaussian = Gaussian::from_ms(0.0, 0.0);
|
||||||
@@ -63,30 +191,70 @@ pub const N_INF: Gaussian = Gaussian::from_ms(0.0, f64::INFINITY);
|
|||||||
#[derive(Copy, Clone, Default, PartialEq, PartialOrd, Eq, Ord, Hash, Debug)]
|
#[derive(Copy, Clone, Default, PartialEq, PartialOrd, Eq, Ord, Hash, Debug)]
|
||||||
pub struct Index(usize);
|
pub struct Index(usize);
|
||||||
|
|
||||||
|
impl Index {
|
||||||
|
/// The underlying slot number.
|
||||||
|
///
|
||||||
|
/// Indices are dense and assigned in interning order, so this is usable as
|
||||||
|
/// a key into a caller-side side table.
|
||||||
|
#[must_use]
|
||||||
|
pub fn get(self) -> usize {
|
||||||
|
self.0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl From<usize> for Index {
|
impl From<usize> for Index {
|
||||||
fn from(ix: usize) -> Self {
|
fn from(ix: usize) -> Self {
|
||||||
Self(ix)
|
Self(ix)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn erfc(x: f64) -> f64 {
|
impl From<Index> for usize {
|
||||||
let z = x.abs();
|
fn from(idx: Index) -> Self {
|
||||||
let t = 1.0 / (1.0 + z / 2.0);
|
idx.0
|
||||||
|
}
|
||||||
let a = -0.82215223 + t * 0.17087277;
|
|
||||||
let b = 1.48851587 + t * a;
|
|
||||||
let c = -1.13520398 + t * b;
|
|
||||||
let d = 0.27886807 + t * c;
|
|
||||||
let e = -0.18628806 + t * d;
|
|
||||||
let f = 0.09678418 + t * e;
|
|
||||||
let g = 0.37409196 + t * f;
|
|
||||||
let h = 1.00002368 + t * g;
|
|
||||||
|
|
||||||
let r = t * (-z * z - 1.26551223 + t * h).exp();
|
|
||||||
|
|
||||||
if x >= 0.0 { r } else { 2.0 - r }
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Complementary error function.
|
||||||
|
///
|
||||||
|
/// # Why every transcendental in this crate goes through `libm`
|
||||||
|
///
|
||||||
|
/// IEEE 754 specifies the basic operations and `sqrt` exactly, but says nothing
|
||||||
|
/// about `exp`, `log` or `erf`. `std`'s versions delegate to the *system* math
|
||||||
|
/// library, so they differ between platforms: measured here, `f64::exp` and
|
||||||
|
/// `libm::exp` disagree on 9.7% of inputs and `f64::ln` / `libm::log` on 5.0%,
|
||||||
|
/// each by one ULP.
|
||||||
|
///
|
||||||
|
/// Inference is an iterative fixed point, so a one-ULP difference can change an
|
||||||
|
/// iteration count and therefore the answer by more than one ULP. Routing every
|
||||||
|
/// transcendental through `libm` makes a fit reproducible across platforms, not
|
||||||
|
/// just across thread counts as `tests/determinism.rs` already checks.
|
||||||
|
///
|
||||||
|
/// **So: use `libm::exp` / `libm::log` in inference code, never `f64::exp` /
|
||||||
|
/// `f64::ln`.** `sqrt` is exempt — IEEE specifies it exactly, so `f64::sqrt` is
|
||||||
|
/// already portable. Test code may use whichever is clearer.
|
||||||
|
///
|
||||||
|
/// It costs nothing: `Batch::iteration` measured -2.7% [-5.7%, -0.3%] with the
|
||||||
|
/// whole set swapped.
|
||||||
|
///
|
||||||
|
/// Delegates to `libm`, which is the Rust port of FDLIBM and accurate to about
|
||||||
|
/// one ULP. This replaced a Numerical Recipes `erfcc` rational approximation
|
||||||
|
/// whose documented bound was 1.2e-7 *relative* — measured at ~1e-7 across the
|
||||||
|
/// whole range, and the binding accuracy constraint on the entire crate.
|
||||||
|
///
|
||||||
|
/// The swap is free. 98% of the arguments inference passes here have
|
||||||
|
/// `|x| < 0.84375`, which is exactly where FDLIBM skips the exponential
|
||||||
|
/// entirely, so the longer polynomial costs nothing on the distribution that
|
||||||
|
/// actually occurs: `Batch::iteration` moved -1.6% [-4.7%, +0.9%], p = 0.31.
|
||||||
|
///
|
||||||
|
/// What it bought: `compute_margin` went from 8.4e-8 to 1.7e-16 against exact
|
||||||
|
/// quantiles, `cdf(mu, mu, sigma)` is now exactly 0.5, and `sf + cdf` sums to
|
||||||
|
/// one within a single ULP where it was 3e-8 out.
|
||||||
|
fn erfc(x: f64) -> f64 {
|
||||||
|
libm::erfc(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The previous Numerical Recipes `erfcc`, kept only so the timing test can
|
||||||
|
/// compare both in one binary. Removed once the comparison is recorded.
|
||||||
fn erfc_inv(mut y: f64) -> f64 {
|
fn erfc_inv(mut y: f64) -> f64 {
|
||||||
if y >= 2.0 {
|
if y >= 2.0 {
|
||||||
return f64::NEG_INFINITY;
|
return f64::NEG_INFINITY;
|
||||||
@@ -102,14 +270,22 @@ fn erfc_inv(mut y: f64) -> f64 {
|
|||||||
y = 2.0 - y;
|
y = 2.0 - y;
|
||||||
}
|
}
|
||||||
|
|
||||||
let t = (-2.0 * (y / 2.0).ln()).sqrt();
|
let t = libm::sqrt(-2.0 * libm::log(y / 2.0));
|
||||||
|
|
||||||
let mut x = FRAC_1_SQRT_2 * ((2.30753 + t * 0.27061) / (1.0 + t * (0.99229 + t * 0.04481)) - t);
|
// The leading coefficient is NEGATIVE. `rational - t` is negative here, so
|
||||||
|
// a positive coefficient mirrors the starting point to `-x0` — the
|
||||||
|
// reflection of the root. Newton then has to cross the origin to get back,
|
||||||
|
// which a fixed iteration count does not manage: measured against the true
|
||||||
|
// value, `erfc_inv(0.1)` returned 1.044 instead of 1.16309, and the error
|
||||||
|
// grew as y shrank until `compute_margin` stopped being monotone in
|
||||||
|
// `p_draw` altogether.
|
||||||
|
let mut x =
|
||||||
|
-FRAC_1_SQRT_2 * ((2.30753 + t * 0.27061) / (1.0 + t * (0.99229 + t * 0.04481)) - t);
|
||||||
|
|
||||||
for _ in 0..3 {
|
for _ in 0..3 {
|
||||||
let err = erfc(x) - y;
|
let err = erfc(x) - y;
|
||||||
|
|
||||||
x += err / (FRAC_2_SQRT_PI * (-(x.powi(2))).exp() - x * err)
|
x += err / (FRAC_2_SQRT_PI * libm::exp(-(x * x)) - x * err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if y < 1.0 { x } else { -x }
|
if y < 1.0 { x } else { -x }
|
||||||
@@ -129,32 +305,216 @@ pub(crate) fn cdf(x: f64, mu: f64, sigma: f64) -> f64 {
|
|||||||
0.5 * erfc(z)
|
0.5 * erfc(z)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// `P(X > x)` for `X ~ N(mu, sigma^2)`.
|
||||||
|
///
|
||||||
|
/// The survival function, computed directly rather than as `1 - cdf(..)`.
|
||||||
|
///
|
||||||
|
/// The two are algebraically identical and numerically are not. `cdf` returns
|
||||||
|
/// a value approaching 1 for an upper tail, so subtracting it from 1 cancels
|
||||||
|
/// away every significant digit the tail had: measured against this function,
|
||||||
|
/// `1 - cdf` carries 7% error by four sigma past the mean and returns exactly
|
||||||
|
/// zero beyond about 8.3 sigma — where the true value is still 1e-19 and
|
||||||
|
/// perfectly representable. `erfc` holds *relative* accuracy all the way down
|
||||||
|
/// to 1e-296, so the precision is there to keep; only the subtraction threw it
|
||||||
|
/// away.
|
||||||
|
///
|
||||||
|
/// This matters most where evidence is smallest, which is exactly where an
|
||||||
|
/// upset makes it interesting: `ln` of a clamped zero is -708 regardless of
|
||||||
|
/// whether the truth was -43 or -600.
|
||||||
|
pub(crate) fn sf(x: f64, mu: f64, sigma: f64) -> f64 {
|
||||||
|
0.5 * erfc((x - mu) / (sigma * SQRT_2))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `e^(x^2) * erfc(x)`, the scaled complementary error function, for `x >= 0`.
|
||||||
|
///
|
||||||
|
/// Exists so the exponential factor common to a Gaussian density and its tail
|
||||||
|
/// integral can be cancelled *analytically* instead of being computed twice
|
||||||
|
/// and divided. Both underflow to zero past about 26 sigma, and their ratio is
|
||||||
|
/// then `0/0` — finite in the limit, `NaN` in floating point.
|
||||||
|
fn erfcx(x: f64) -> f64 {
|
||||||
|
if x < 2.0 {
|
||||||
|
// Below the crossover neither factor is extreme: erfc is O(1) and
|
||||||
|
// exp(x^2) is at most e^4, so the direct product is exact enough and
|
||||||
|
// cheaper than the continued fraction.
|
||||||
|
libm::exp(x * x) * erfc(x)
|
||||||
|
} else {
|
||||||
|
// erfcx(x) = 1/sqrt(pi) * 1/(x + (1/2)/(x + 1/(x + (3/2)/(x + ...)))),
|
||||||
|
// evaluated by backward recurrence. Converges quickly for x >= 2 and,
|
||||||
|
// unlike the product form, never touches an exponential.
|
||||||
|
let mut f = 0.0;
|
||||||
|
for n in (1..=60u32).rev() {
|
||||||
|
f = (f64::from(n) * 0.5) / (x + f);
|
||||||
|
}
|
||||||
|
FRAC_1_SQRT_PI / (x + f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `ln` of the normal density at `x`.
|
||||||
|
///
|
||||||
|
/// The density itself underflows to zero past about 38 sigma, and `ln` of a
|
||||||
|
/// clamped zero is -708 whatever the truth was. The log form is a polynomial:
|
||||||
|
/// it stays exact at any separation, and the values it produces (-5001 nats at
|
||||||
|
/// 100 sigma, -500001 at 1000) are perfectly representable.
|
||||||
|
pub(crate) fn ln_pdf(x: f64, mu: f64, sigma: f64) -> f64 {
|
||||||
|
let z = (x - mu) / sigma;
|
||||||
|
-libm::log(SQRT_TAU * sigma) - 0.5 * z * z
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `ln P(X > x)` for `X ~ N(mu, sigma^2)`.
|
||||||
|
///
|
||||||
|
/// In the upper tail the `exp(-z^2 / 2)` common to the tail integral is
|
||||||
|
/// factored out analytically via `erfcx`, so this never underflows — where
|
||||||
|
/// `sf(..).ln()` bottoms out at -708 once `erfc` itself reaches zero.
|
||||||
|
pub(crate) fn ln_sf(x: f64, mu: f64, sigma: f64) -> f64 {
|
||||||
|
let z = (x - mu) / sigma;
|
||||||
|
|
||||||
|
if z > 0.0 {
|
||||||
|
// ln(0.5 * erfc(z/sqrt2)) with erfc(y) = exp(-y^2) * erfcx(y).
|
||||||
|
-std::f64::consts::LN_2 - 0.5 * z * z + libm::log(erfcx(z / SQRT_2))
|
||||||
|
} else {
|
||||||
|
// The mass here is at least a half; nothing to lose.
|
||||||
|
libm::log(sf(x, mu, sigma))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `ln P(lo < X < hi)` for `X ~ N(mu, sigma^2)`.
|
||||||
|
///
|
||||||
|
/// When the interval sits in a tail both endpoint probabilities underflow
|
||||||
|
/// together, so their difference is taken in scaled form with the shared
|
||||||
|
/// exponential factored out. When it straddles the mean nothing is small and
|
||||||
|
/// the direct difference is exact.
|
||||||
|
pub(crate) fn ln_interval(lo: f64, hi: f64, mu: f64, sigma: f64) -> f64 {
|
||||||
|
let z_lo = (lo - mu) / sigma;
|
||||||
|
let z_hi = (hi - mu) / sigma;
|
||||||
|
|
||||||
|
if z_hi <= z_lo {
|
||||||
|
return f64::NEG_INFINITY;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fold a lower-tail interval onto the upper tail; the normal is symmetric.
|
||||||
|
let (near, far) = if z_lo >= 0.0 {
|
||||||
|
(z_lo, z_hi)
|
||||||
|
} else if z_hi <= 0.0 {
|
||||||
|
(-z_hi, -z_lo)
|
||||||
|
} else {
|
||||||
|
// Straddles the mean: the interval holds a non-negligible share of the
|
||||||
|
// mass, so neither endpoint is near enough to 1 to cancel.
|
||||||
|
return libm::log((cdf(hi, mu, sigma) - cdf(lo, mu, sigma)).max(f64::MIN_POSITIVE));
|
||||||
|
};
|
||||||
|
|
||||||
|
let (a, b) = (near / SQRT_2, far / SQRT_2);
|
||||||
|
// b > a >= 0, so this ratio of exponentials is at most 1 and cannot overflow.
|
||||||
|
let scale = libm::exp(a * a - b * b);
|
||||||
|
let bracket = erfcx(a) - scale * erfcx(b);
|
||||||
|
|
||||||
|
if bracket <= 0.0 {
|
||||||
|
return f64::NEG_INFINITY;
|
||||||
|
}
|
||||||
|
|
||||||
|
-std::f64::consts::LN_2 - a * a + libm::log(bracket)
|
||||||
|
}
|
||||||
|
|
||||||
fn pdf(x: f64, mu: f64, sigma: f64) -> f64 {
|
fn pdf(x: f64, mu: f64, sigma: f64) -> f64 {
|
||||||
let normalizer = (SQRT_TAU * sigma).powi(-1);
|
let normalizer = (SQRT_TAU * sigma).powi(-1);
|
||||||
let functional = (-((x - mu).powi(2)) / (2.0 * sigma.powi(2))).exp();
|
let functional = libm::exp(-((x - mu) * (x - mu)) / (2.0 * sigma * sigma));
|
||||||
|
|
||||||
normalizer * functional
|
normalizer * functional
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Truncated-Gaussian correction terms `(v, w)`.
|
||||||
|
///
|
||||||
|
/// `v` shifts the mean and `w` shrinks the variance. Both are ratios whose
|
||||||
|
/// numerator and denominator underflow together in the tails, so both are
|
||||||
|
/// computed in scaled form there: the shared `exp(-alpha^2 / 2)` is cancelled
|
||||||
|
/// analytically rather than evaluated and divided out. Without that, a
|
||||||
|
/// truncation point beyond about 39 sigma produced `0 / 0` and put `NaN`
|
||||||
|
/// straight into the posterior.
|
||||||
|
/// Truncation terms for a boundary `alpha` standard deviations into the upper
|
||||||
|
/// tail, from the asymptotic expansion of the inverse Mills ratio.
|
||||||
|
///
|
||||||
|
/// `v` tends to `alpha` out here, so the gap between them cannot be obtained by
|
||||||
|
/// subtracting one from the other — the series computes the gap directly, and
|
||||||
|
/// `w = v * gap` then never forms the difference of two large near-equal
|
||||||
|
/// numbers. A far-tail *window* behaves like a half-line once it is more than a
|
||||||
|
/// few multiples of its own width from the mean, so the tie branch shares this.
|
||||||
|
fn half_line_truncation(alpha: f64) -> (f64, f64) {
|
||||||
|
let inv = alpha.recip();
|
||||||
|
let inv_sq = inv * inv;
|
||||||
|
let gap = inv * (1.0 - inv_sq * (2.0 - inv_sq * (10.0 - 74.0 * inv_sq)));
|
||||||
|
let v = alpha + gap;
|
||||||
|
|
||||||
|
(v, v * gap)
|
||||||
|
}
|
||||||
|
|
||||||
fn v_w(mu: f64, sigma: f64, margin: f64, tie: bool) -> (f64, f64) {
|
fn v_w(mu: f64, sigma: f64, margin: f64, tie: bool) -> (f64, f64) {
|
||||||
if !tie {
|
if !tie {
|
||||||
let alpha = (margin - mu) / sigma;
|
let alpha = (margin - mu) / sigma;
|
||||||
|
|
||||||
let v = pdf(-alpha, 0.0, 1.0) / cdf(-alpha, 0.0, 1.0);
|
// v is the inverse Mills ratio, phi(alpha) / Phi(-alpha), and w needs
|
||||||
let w = v * (v + (-alpha));
|
// the gap `v - alpha` as well as v itself. Far into the tail v tends to
|
||||||
|
// alpha, so that gap is a subtraction of two nearly equal numbers and
|
||||||
|
// loses every digit it has: at alpha = 1e6 it drove w above 1 and made
|
||||||
|
// `sqrt(1 - w)` NaN. Past the crossover the gap comes from its
|
||||||
|
// asymptotic series instead, which has no subtraction in it.
|
||||||
|
if alpha >= ASYMPTOTIC_MILLS_ALPHA {
|
||||||
|
return half_line_truncation(alpha);
|
||||||
|
}
|
||||||
|
|
||||||
(v, w)
|
let (v, gap) = if alpha > 0.0 {
|
||||||
|
// Both terms carry exp(-alpha^2 / 2); in scaled form it cancels
|
||||||
|
// and the result stays exact however far into the tail alpha sits.
|
||||||
|
let v = SQRT_2_OVER_PI / erfcx(alpha / SQRT_2);
|
||||||
|
(v, v - alpha)
|
||||||
} else {
|
} else {
|
||||||
|
// Phi(-alpha) >= 1/2 here, so the direct ratio loses nothing.
|
||||||
|
let v = pdf(-alpha, 0.0, 1.0) / cdf(-alpha, 0.0, 1.0);
|
||||||
|
(v, v - alpha)
|
||||||
|
};
|
||||||
|
|
||||||
|
(v, v * gap)
|
||||||
|
} else {
|
||||||
|
// v is odd in mu and w is even, so fold to mu <= 0. Both truncation
|
||||||
|
// points then sit in the upper tail, where the scaled form applies.
|
||||||
|
let flipped = mu > 0.0;
|
||||||
|
let mu = if flipped { -mu } else { mu };
|
||||||
|
|
||||||
let alpha = (-margin - mu) / sigma;
|
let alpha = (-margin - mu) / sigma;
|
||||||
let beta = (margin - mu) / sigma;
|
let beta = (margin - mu) / sigma;
|
||||||
|
|
||||||
let v = (pdf(alpha, 0.0, 1.0) - pdf(beta, 0.0, 1.0))
|
// `w` comes out of `v * v - u`, and both terms grow as alpha^2 while
|
||||||
/ (cdf(beta, 0.0, 1.0) - cdf(alpha, 0.0, 1.0));
|
// their difference stays O(1) — at alpha = 1e9 that subtraction had no
|
||||||
let u = (alpha * pdf(alpha, 0.0, 1.0) - beta * pdf(beta, 0.0, 1.0))
|
// digits left and returned w = -128, making `sqrt(1 - w)` nonsense.
|
||||||
/ (cdf(beta, 0.0, 1.0) - cdf(alpha, 0.0, 1.0));
|
// Once the window sits many of its own widths into the tail it is
|
||||||
|
// indistinguishable from a half-line, so the asymptotic covers it with
|
||||||
|
// no subtraction at all.
|
||||||
|
if alpha >= ASYMPTOTIC_MILLS_ALPHA && alpha * (beta - alpha) >= HALF_LINE_WINDOW {
|
||||||
|
let (v, w) = half_line_truncation(alpha);
|
||||||
|
return (if flipped { -v } else { v }, w);
|
||||||
|
}
|
||||||
|
|
||||||
|
let (v, u) = if alpha > 0.0 {
|
||||||
|
// beta > alpha > 0, so this ratio of exponentials is at most 1 and
|
||||||
|
// cannot overflow.
|
||||||
|
let scale = libm::exp(0.5 * (alpha * alpha - beta * beta));
|
||||||
|
let denominator = 0.5 * (erfcx(alpha / SQRT_2) - scale * erfcx(beta / SQRT_2));
|
||||||
|
|
||||||
|
(
|
||||||
|
(1.0 - scale) / SQRT_TAU / denominator,
|
||||||
|
(alpha - beta * scale) / SQRT_TAU / denominator,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
// The interval straddles the mean, so nothing here is small.
|
||||||
|
let denominator = cdf(beta, 0.0, 1.0) - cdf(alpha, 0.0, 1.0);
|
||||||
|
|
||||||
|
(
|
||||||
|
(pdf(alpha, 0.0, 1.0) - pdf(beta, 0.0, 1.0)) / denominator,
|
||||||
|
(alpha * pdf(alpha, 0.0, 1.0) - beta * pdf(beta, 0.0, 1.0)) / denominator,
|
||||||
|
)
|
||||||
|
};
|
||||||
|
|
||||||
let w = -(u - v.powi(2));
|
let w = -(u - v.powi(2));
|
||||||
|
|
||||||
(v, w)
|
(if flipped { -v } else { v }, w)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -184,6 +544,56 @@ pub(crate) fn tuple_gt(t: (f64, f64), e: f64) -> bool {
|
|||||||
t.0 > e || t.1 > e
|
t.0 > e || t.1 > e
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Whether a convergence step is finite in both components.
|
||||||
|
///
|
||||||
|
/// A NaN step means EP broke down numerically. Because every comparison
|
||||||
|
/// against NaN is false, `tuple_gt` reads NaN as "below epsilon" — so
|
||||||
|
/// convergence checks must test finiteness explicitly rather than inferring
|
||||||
|
/// success from `!tuple_gt(..)`.
|
||||||
|
pub(crate) fn step_is_finite(t: (f64, f64)) -> bool {
|
||||||
|
t.0.is_finite() && t.1.is_finite()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Whether a step counts as converged: finite *and* within `epsilon`.
|
||||||
|
pub(crate) fn step_converged(t: (f64, f64), epsilon: f64) -> bool {
|
||||||
|
step_is_finite(t) && !tuple_gt(t, epsilon)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Indices of the first pair of teams sharing a rank, if any.
|
||||||
|
///
|
||||||
|
/// A tie is only representable when the draw probability is positive: with
|
||||||
|
/// `p_draw == 0.0` the truncation margin collapses to zero and the two-sided
|
||||||
|
/// tie update evaluates `0/0`. Callers use this to reject such events before
|
||||||
|
/// they reach inference.
|
||||||
|
pub(crate) fn first_tied_pair(ranks: &[u32]) -> Option<(usize, usize)> {
|
||||||
|
for (i, a) in ranks.iter().enumerate() {
|
||||||
|
for (j, b) in ranks.iter().enumerate().skip(i + 1) {
|
||||||
|
if a == b {
|
||||||
|
return Some((i, j));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
|
/// As `first_tied_pair`, but over the engine's internal `f64` outputs.
|
||||||
|
///
|
||||||
|
/// Ranks reach the engine already converted to descending `f64` outputs, and
|
||||||
|
/// `Game` decides a tie by exact equality of those values — so this mirrors
|
||||||
|
/// the comparison inference itself performs.
|
||||||
|
pub(crate) fn first_tied_output(outputs: &[f64]) -> Option<(usize, usize)> {
|
||||||
|
for (i, a) in outputs.iter().enumerate() {
|
||||||
|
for (j, b) in outputs.iter().enumerate().skip(i + 1) {
|
||||||
|
if a == b {
|
||||||
|
return Some((i, j));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn sort_time<T: Copy + Ord>(xs: &[T], reverse: bool) -> Vec<usize> {
|
pub(crate) fn sort_time<T: Copy + Ord>(xs: &[T], reverse: bool) -> Vec<usize> {
|
||||||
let mut x: Vec<(usize, T)> = xs.iter().enumerate().map(|(i, &t)| (i, t)).collect();
|
let mut x: Vec<(usize, T)> = xs.iter().enumerate().map(|(i, &t)| (i, t)).collect();
|
||||||
|
|
||||||
@@ -197,7 +607,27 @@ pub(crate) fn sort_time<T: Copy + Ord>(xs: &[T], reverse: bool) -> Vec<usize> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Calculates the match quality of the given rating groups. A result is the draw probability in the association
|
/// Calculates the match quality of the given rating groups. A result is the draw probability in the association
|
||||||
|
///
|
||||||
|
/// Supports any number of groups. Values range roughly `[0, 1]`; 1 means a
|
||||||
|
/// perfectly balanced match.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if fewer than two rating groups are supplied, or if any group is
|
||||||
|
/// empty — match quality is a property of a contest between at least two
|
||||||
|
/// non-empty sides.
|
||||||
|
#[must_use]
|
||||||
pub fn quality(rating_groups: &[&[Gaussian]], beta: f64) -> f64 {
|
pub fn quality(rating_groups: &[&[Gaussian]], beta: f64) -> f64 {
|
||||||
|
assert!(
|
||||||
|
rating_groups.len() >= 2,
|
||||||
|
"quality() requires at least 2 rating groups, got {}",
|
||||||
|
rating_groups.len()
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
rating_groups.iter().all(|group| !group.is_empty()),
|
||||||
|
"quality() requires every rating group to be non-empty"
|
||||||
|
);
|
||||||
|
|
||||||
let flatten_ratings = rating_groups
|
let flatten_ratings = rating_groups
|
||||||
.iter()
|
.iter()
|
||||||
.flat_map(|group| group.iter())
|
.flat_map(|group| group.iter())
|
||||||
@@ -221,8 +651,10 @@ pub fn quality(rating_groups: &[&[Gaussian]], beta: f64) -> f64 {
|
|||||||
|
|
||||||
let mut rotated_a_matrix = Matrix::new(rating_groups.len() - 1, length);
|
let mut rotated_a_matrix = Matrix::new(rating_groups.len() - 1, length);
|
||||||
|
|
||||||
|
// Row `row` contrasts group `row` (+weight) against group `row + 1`
|
||||||
|
// (-weight). `t` is the column where the current group's players start;
|
||||||
|
// the negative block begins immediately after it.
|
||||||
let mut t = 0;
|
let mut t = 0;
|
||||||
let mut x = 0;
|
|
||||||
|
|
||||||
for (row, group) in rating_groups.windows(2).enumerate() {
|
for (row, group) in rating_groups.windows(2).enumerate() {
|
||||||
let current = group[0];
|
let current = group[0];
|
||||||
@@ -230,17 +662,13 @@ pub fn quality(rating_groups: &[&[Gaussian]], beta: f64) -> f64 {
|
|||||||
|
|
||||||
for n in t..t + current.len() {
|
for n in t..t + current.len() {
|
||||||
rotated_a_matrix[(row, n)] = flatten_weights[n];
|
rotated_a_matrix[(row, n)] = flatten_weights[n];
|
||||||
|
|
||||||
x += 1;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
t += current.len();
|
t += current.len();
|
||||||
|
|
||||||
for n in x..x + next.len() {
|
for n in t..t + next.len() {
|
||||||
rotated_a_matrix[(row, n)] = -flatten_weights[n];
|
rotated_a_matrix[(row, n)] = -flatten_weights[n];
|
||||||
}
|
}
|
||||||
|
|
||||||
x += next.len();
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let a_matrix = rotated_a_matrix.transpose();
|
let a_matrix = rotated_a_matrix.transpose();
|
||||||
@@ -255,7 +683,7 @@ pub fn quality(rating_groups: &[&[Gaussian]], beta: f64) -> f64 {
|
|||||||
let e_arg = (-0.5 * &start * &middle.inverse() * &end).determinant();
|
let e_arg = (-0.5 * &start * &middle.inverse() * &end).determinant();
|
||||||
let s_arg = ata.determinant() / middle.determinant();
|
let s_arg = ata.determinant() / middle.determinant();
|
||||||
|
|
||||||
e_arg.exp() * s_arg.sqrt()
|
libm::exp(e_arg) * s_arg.sqrt()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -269,6 +697,313 @@ mod tests {
|
|||||||
assert_eq!(sort_time(&[0i64, 1, 2, 0], true), vec![2, 1, 0, 3]);
|
assert_eq!(sort_time(&[0i64, 1, 2, 0], true), vec![2, 1, 0, 3]);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Upper-tail values of the standard normal, from published tables. The
|
||||||
|
/// point is not the digits — these are 7-digit table values — but that a
|
||||||
|
/// number comes back at all: `1 - cdf` returned exactly zero for every one
|
||||||
|
/// of these.
|
||||||
|
#[test]
|
||||||
|
fn survival_function_survives_the_far_tail() {
|
||||||
|
for (z, expected) in [
|
||||||
|
(9.0f64, 1.128_588e-19),
|
||||||
|
(12.0, 1.776_482e-33),
|
||||||
|
(20.0, 2.753_624e-89),
|
||||||
|
(37.0, 5.725_571e-300),
|
||||||
|
] {
|
||||||
|
let got = sf(z, 0.0, 1.0);
|
||||||
|
assert!(got > 0.0, "sf({z}) collapsed to zero");
|
||||||
|
assert!(
|
||||||
|
(got - expected).abs() / expected < 1e-6, // published table values, 7 digits
|
||||||
|
"sf({z}) = {got}, expected ~{expected}"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
1.0 - cdf(z, 0.0, 1.0),
|
||||||
|
0.0,
|
||||||
|
"the naive form should still be zero here"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Where no cancellation happens the two forms must agree exactly enough
|
||||||
|
/// that nothing else in the crate shifts.
|
||||||
|
#[test]
|
||||||
|
fn survival_function_matches_the_naive_form_where_that_form_works() {
|
||||||
|
for z in [-4.0f64, -1.0, 0.0, 0.5, 1.0, 2.0, 3.0, 4.0] {
|
||||||
|
let naive = 1.0 - cdf(z, 0.0, 1.0);
|
||||||
|
let direct = sf(z, 0.0, 1.0);
|
||||||
|
assert!(
|
||||||
|
(naive - direct).abs() < 1e-15,
|
||||||
|
"z={z}: naive {naive} vs direct {direct}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn survival_and_cdf_partition_the_mass() {
|
||||||
|
for z in [-3.0f64, -0.5, 0.0, 1.0, 2.5] {
|
||||||
|
let total = sf(z, 1.0, 2.0) + cdf(z, 1.0, 2.0);
|
||||||
|
assert!((total - 1.0).abs() < 1e-15, "z={z}: {total}");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `erfcx` switches formulation at x = 2; the two sides must meet.
|
||||||
|
#[test]
|
||||||
|
fn erfcx_is_continuous_across_its_crossover() {
|
||||||
|
for x in [1.90f64, 1.99, 1.999, 2.0, 2.001, 2.01, 2.10] {
|
||||||
|
let direct = (x * x).exp() * erfc(x);
|
||||||
|
let scaled = erfcx(x);
|
||||||
|
assert!(
|
||||||
|
(direct - scaled).abs() / scaled < 1e-14,
|
||||||
|
"x={x}: direct {direct} vs erfcx {scaled}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The whole reason `erfcx` exists: it stays finite and O(1/x) exactly
|
||||||
|
/// where `exp(x^2)` overflows and `erfc(x)` underflows.
|
||||||
|
#[test]
|
||||||
|
fn erfcx_stays_finite_where_its_factors_do_not() {
|
||||||
|
for x in [27.0f64, 50.0, 1.0e3, 1.0e8] {
|
||||||
|
let scaled = erfcx(x);
|
||||||
|
assert!(scaled.is_finite() && scaled > 0.0, "erfcx({x}) = {scaled}");
|
||||||
|
// Asymptotically erfcx(x) -> 1 / (x * sqrt(pi)).
|
||||||
|
let asymptote = 1.0 / (x * std::f64::consts::PI.sqrt());
|
||||||
|
assert!(
|
||||||
|
(scaled - asymptote).abs() / asymptote < 1e-2,
|
||||||
|
"erfcx({x}) = {scaled} strays from its asymptote {asymptote}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(x * x).exp().is_infinite(),
|
||||||
|
"x={x} should overflow the direct form"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Truncation must never produce a non-finite posterior. Before the scaled
|
||||||
|
/// formulation these returned NaN from `0 / 0` past about 39 sigma.
|
||||||
|
#[test]
|
||||||
|
fn truncation_stays_finite_arbitrarily_far_into_the_tail() {
|
||||||
|
for alpha in [0.0f64, 8.0, 38.0, 40.0, 100.0, 1.0e3, 1.0e6, 1.0e9, 1.0e15] {
|
||||||
|
for tie in [false, true] {
|
||||||
|
let (v, w) = v_w(-alpha, 1.0, if tie { 1.0 } else { 0.0 }, tie);
|
||||||
|
assert!(v.is_finite(), "alpha={alpha} tie={tie}: v = {v}");
|
||||||
|
assert!(w.is_finite(), "alpha={alpha} tie={tie}: w = {w}");
|
||||||
|
// sigma_trunc = sigma * sqrt(1 - w) must stay real.
|
||||||
|
assert!(
|
||||||
|
(0.0..=1.0).contains(&w),
|
||||||
|
"alpha={alpha} tie={tie}: w = {w} leaves sqrt(1 - w) imaginary"
|
||||||
|
);
|
||||||
|
|
||||||
|
let (mu_t, sigma_t) = trunc(-alpha, 1.0, if tie { 1.0 } else { 0.0 }, tie);
|
||||||
|
assert!(
|
||||||
|
mu_t.is_finite() && sigma_t.is_finite(),
|
||||||
|
"alpha={alpha} tie={tie}: trunc = ({mu_t}, {sigma_t})"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The Mills gap switches from subtraction to series at alpha = 100. Both
|
||||||
|
/// are supposed to be right there; if they disagree, the crossover is in
|
||||||
|
/// the wrong place.
|
||||||
|
#[test]
|
||||||
|
fn the_mills_gap_series_meets_the_scaled_form() {
|
||||||
|
for alpha in [50.0f64, 99.0, 100.0, 101.0, 200.0] {
|
||||||
|
let scaled = SQRT_2_OVER_PI / erfcx(alpha / SQRT_2) - alpha;
|
||||||
|
let inv = alpha.recip();
|
||||||
|
let inv_sq = inv * inv;
|
||||||
|
let series = inv * (1.0 - inv_sq * (2.0 - inv_sq * (10.0 - 74.0 * inv_sq)));
|
||||||
|
assert!(
|
||||||
|
(scaled - series).abs() / series < 1e-9,
|
||||||
|
"alpha={alpha}: scaled {scaled} vs series {series}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Folding the tie branch to `mu <= 0` is only valid if v is odd in mu and
|
||||||
|
/// w is even. Assert the symmetry the implementation relies on.
|
||||||
|
#[test]
|
||||||
|
fn tie_truncation_is_odd_in_v_and_even_in_w() {
|
||||||
|
for mu in [0.5f64, 3.0, 20.0, 40.0, 100.0, 1.0e3] {
|
||||||
|
let (v_pos, w_pos) = v_w(mu, 1.0, 1.0, true);
|
||||||
|
let (v_neg, w_neg) = v_w(-mu, 1.0, 1.0, true);
|
||||||
|
assert!(
|
||||||
|
(v_pos + v_neg).abs() < 1e-9,
|
||||||
|
"mu={mu}: v should be odd, got {v_pos} and {v_neg}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(w_pos - w_neg).abs() < 1e-9,
|
||||||
|
"mu={mu}: w should be even, got {w_pos} and {w_neg}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `erfc_inv`'s initial guess had the wrong sign, putting Newton on the
|
||||||
|
/// mirror image of the root. Three fixed iterations could not cross back,
|
||||||
|
/// so the error grew as the argument shrank: at `p_draw = 0.99` the margin
|
||||||
|
/// came out 0.503 where the answer is 2.576.
|
||||||
|
#[test]
|
||||||
|
fn erfc_inv_matches_known_quantiles() {
|
||||||
|
// sqrt(2) * erfc_inv(1 - p) is the standard normal quantile
|
||||||
|
// Phi^-1((1 + p) / 2).
|
||||||
|
for (p, exact) in [
|
||||||
|
(0.5f64, 0.674_489_750_196_081_7f64),
|
||||||
|
(0.9, 1.644_853_626_951_472_7),
|
||||||
|
(0.95, 1.959_963_984_540_054_2),
|
||||||
|
(0.99, 2.575_829_303_548_9),
|
||||||
|
(0.999, 3.290_526_731_491_896_4),
|
||||||
|
] {
|
||||||
|
let got = SQRT_2 * erfc_inv(1.0 - p);
|
||||||
|
assert!(
|
||||||
|
(got - exact).abs() / exact < 1e-14,
|
||||||
|
"p={p}: got {got}, exact {exact}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The draw margin must grow with the draw probability. It did not: it ran
|
||||||
|
/// 0.674 -> 1.476 -> 0.503 -> 0.982 as `p_draw` went 0.5 -> 0.9 -> 0.99 ->
|
||||||
|
/// 0.999, which is not a rounding error but a broken function.
|
||||||
|
/// Deep in the tail the accuracy limit is the *caller's* argument, not this
|
||||||
|
/// function.
|
||||||
|
///
|
||||||
|
/// `compute_margin(0.999999, ..)` computes `1.0 - p_draw`, and 0.999999 is
|
||||||
|
/// not representable: the subtraction cancels and leaves 2.9e-11 of
|
||||||
|
/// relative error in the argument before `erfc_inv` is even entered. Given
|
||||||
|
/// an exactly-representable argument the result is good to 1.8e-16, so this
|
||||||
|
/// is inherent to taking `p_draw` near one rather than something to fix
|
||||||
|
/// here. At `p_draw = 0.999` the whole path is still accurate to 4e-16.
|
||||||
|
///
|
||||||
|
/// Worth pinning: measured against a 70-digit reference, `puruspe`'s
|
||||||
|
/// `inverfc` returns the identical wrong value for the identical reason,
|
||||||
|
/// which is what makes it clear the fault is upstream of both.
|
||||||
|
#[test]
|
||||||
|
fn erfc_inv_is_exact_given_an_exactly_representable_argument() {
|
||||||
|
// erfc(z / sqrt2) = 1e-6 exactly, so z = Phi^-1(0.9999995).
|
||||||
|
let got = SQRT_2 * erfc_inv(1e-6);
|
||||||
|
let exact = 4.891_638_475_698_59;
|
||||||
|
assert!(
|
||||||
|
(got - exact).abs() / exact < 1e-14,
|
||||||
|
"got {got}, exact {exact}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn compute_margin_is_monotone_in_the_draw_probability() {
|
||||||
|
let mut previous = 0.0;
|
||||||
|
for p_draw in [
|
||||||
|
0.001f64, 0.01, 0.1, 0.25, 0.5, 0.75, 0.9, 0.99, 0.999, 0.9999,
|
||||||
|
] {
|
||||||
|
let margin = compute_margin(p_draw, 1.0);
|
||||||
|
assert!(
|
||||||
|
margin > previous,
|
||||||
|
"p_draw={p_draw}: margin {margin} did not exceed {previous}"
|
||||||
|
);
|
||||||
|
previous = margin;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Round-tripping the margin back through the model's own CDF must recover
|
||||||
|
/// the draw probability it was built from.
|
||||||
|
#[test]
|
||||||
|
fn compute_margin_round_trips_through_the_cdf() {
|
||||||
|
for p_draw in [0.001f64, 0.1, 0.5, 0.9, 0.99, 0.999] {
|
||||||
|
for sd in [0.5f64, 1.0, 5.892_557] {
|
||||||
|
let margin = compute_margin(p_draw, sd);
|
||||||
|
// P(|X| < margin) for X ~ N(0, sd^2).
|
||||||
|
let recovered = 1.0 - 2.0 * cdf(-margin, 0.0, sd);
|
||||||
|
assert!(
|
||||||
|
(recovered - p_draw).abs() < 1e-14,
|
||||||
|
"p_draw={p_draw} sd={sd}: recovered {recovered}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `ln_pdf`, `ln_sf` and `ln_interval` exist so evidence stays exact where
|
||||||
|
/// the linear forms underflow. Past ~38 sigma the linear value is zero and
|
||||||
|
/// its log is whatever floor it was clamped to.
|
||||||
|
#[test]
|
||||||
|
fn log_space_helpers_stay_exact_where_the_linear_forms_underflow() {
|
||||||
|
for z in [40.0f64, 60.0, 100.0, 1000.0] {
|
||||||
|
assert_eq!(pdf(z, 0.0, 1.0), 0.0, "pdf should underflow at {z}");
|
||||||
|
assert_eq!(sf(z, 0.0, 1.0), 0.0, "sf should underflow at {z}");
|
||||||
|
|
||||||
|
let lp = ln_pdf(z, 0.0, 1.0);
|
||||||
|
let expected_lp = -(SQRT_TAU).ln() - 0.5 * z * z;
|
||||||
|
assert!(
|
||||||
|
(lp - expected_lp).abs() < 1e-9,
|
||||||
|
"ln_pdf({z}) = {lp}, expected {expected_lp}"
|
||||||
|
);
|
||||||
|
|
||||||
|
let ls = ln_sf(z, 0.0, 1.0);
|
||||||
|
// ln Phi(-z) ~ -z^2/2 - ln(z) - ln(sqrt(2 pi)) for large z.
|
||||||
|
let approx = -0.5 * z * z - z.ln() - SQRT_TAU.ln();
|
||||||
|
assert!(
|
||||||
|
(ls - approx).abs() / approx.abs() < 1e-3,
|
||||||
|
"ln_sf({z}) = {ls}, asymptote {approx}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
ls < f64::MIN_POSITIVE.ln(),
|
||||||
|
"ln_sf({z}) still on the clamp floor"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Where nothing underflows, the log helpers must agree with the direct
|
||||||
|
/// forms exactly enough that nothing else in the crate shifts.
|
||||||
|
#[test]
|
||||||
|
fn log_space_helpers_agree_with_the_linear_forms_in_range() {
|
||||||
|
for z in [-3.0f64, -1.0, 0.0, 1.0, 2.0, 5.0, 10.0, 20.0] {
|
||||||
|
let lp = ln_pdf(z, 0.5, 2.0);
|
||||||
|
let direct_pdf = pdf(z, 0.5, 2.0);
|
||||||
|
assert!(
|
||||||
|
(lp.exp() - direct_pdf).abs() <= 1e-12 * direct_pdf,
|
||||||
|
"ln_pdf at {z}: {} vs {direct_pdf}",
|
||||||
|
lp.exp()
|
||||||
|
);
|
||||||
|
|
||||||
|
let ls = ln_sf(z, 0.5, 2.0);
|
||||||
|
let direct = sf(z, 0.5, 2.0);
|
||||||
|
assert!(
|
||||||
|
(ls.exp() - direct).abs() <= 1e-13 * direct.max(1e-300),
|
||||||
|
"ln_sf at {z}: {} vs {direct}",
|
||||||
|
ls.exp()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn ln_interval_matches_the_direct_difference_when_nothing_is_small() {
|
||||||
|
for mu in [-2.0f64, 0.0, 0.5, 2.0] {
|
||||||
|
let direct = cdf(1.0, mu, 1.0) - cdf(-1.0, mu, 1.0);
|
||||||
|
let logged = ln_interval(-1.0, 1.0, mu, 1.0).exp();
|
||||||
|
assert!(
|
||||||
|
(logged - direct).abs() <= 1e-13 * direct,
|
||||||
|
"mu={mu}: {logged} vs {direct}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A window far out in the tail: both endpoints underflow together, so the
|
||||||
|
/// difference has to be taken in scaled form.
|
||||||
|
#[test]
|
||||||
|
fn ln_interval_survives_a_window_deep_in_the_tail() {
|
||||||
|
for mu in [-50.0f64, -100.0, -1000.0] {
|
||||||
|
let logged = ln_interval(-1.0, 1.0, mu, 1.0);
|
||||||
|
assert!(logged.is_finite(), "mu={mu}: {logged}");
|
||||||
|
assert!(
|
||||||
|
logged < f64::MIN_POSITIVE.ln(),
|
||||||
|
"mu={mu}: {logged} is stuck on the clamp floor"
|
||||||
|
);
|
||||||
|
// Dominated by the near edge: ln P ~ ln Phi(-(|mu| - 1)).
|
||||||
|
let near = ln_sf(-1.0, mu, 1.0);
|
||||||
|
assert!(
|
||||||
|
(logged - near).abs() < 5.0,
|
||||||
|
"mu={mu}: {logged} strays from the near-edge tail {near}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_quality() {
|
fn test_quality() {
|
||||||
let a = Gaussian::from_ms(25.0, 3.0);
|
let a = Gaussian::from_ms(25.0, 3.0);
|
||||||
|
|||||||
+316
-119
@@ -1,29 +1,13 @@
|
|||||||
|
//! Minimal dense matrix used by `quality()`.
|
||||||
|
//!
|
||||||
|
//! `determinant` and `inverse` go through one LU decomposition with partial
|
||||||
|
//! pivoting — O(n³) and numerically stable. The previous implementation
|
||||||
|
//! expanded cofactors recursively (O(n!), allocating a `Vec` per minor) and
|
||||||
|
//! only implemented `inverse` for the 1×1 case, which limited `quality()` to
|
||||||
|
//! exactly two rating groups.
|
||||||
|
|
||||||
use std::ops;
|
use std::ops;
|
||||||
|
|
||||||
fn det(m: &[f64], x: usize) -> f64 {
|
|
||||||
if x == 1 {
|
|
||||||
m[0]
|
|
||||||
} else if x == 2 {
|
|
||||||
m[0] * m[3] - m[1] * m[2]
|
|
||||||
} else {
|
|
||||||
let mut d = 0.0;
|
|
||||||
|
|
||||||
for n in 0..x {
|
|
||||||
let ms = m
|
|
||||||
.iter()
|
|
||||||
.enumerate()
|
|
||||||
.skip(x)
|
|
||||||
.filter(|(i, _)| (i % x) != n)
|
|
||||||
.map(|(_, v)| *v)
|
|
||||||
.collect::<Vec<_>>();
|
|
||||||
|
|
||||||
d += (-1.0f64).powi(n as i32) * m[n] * det(&ms, x - 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
d
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct Matrix {
|
pub struct Matrix {
|
||||||
data: Box<[f64]>,
|
data: Box<[f64]>,
|
||||||
@@ -31,6 +15,107 @@ pub struct Matrix {
|
|||||||
width: usize,
|
width: usize,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// LU decomposition with partial pivoting: `PA = LU`, stored compactly.
|
||||||
|
///
|
||||||
|
/// `lu` holds `L` below the diagonal (unit diagonal implied) and `U` on and
|
||||||
|
/// above it. `sign` is the determinant sign contributed by row swaps, or 0.0
|
||||||
|
/// when the matrix is singular.
|
||||||
|
struct Lu {
|
||||||
|
lu: Vec<f64>,
|
||||||
|
perm: Vec<usize>,
|
||||||
|
n: usize,
|
||||||
|
sign: f64,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Lu {
|
||||||
|
fn decompose(m: &Matrix) -> Self {
|
||||||
|
debug_assert_eq!(m.width, m.height, "LU requires a square matrix");
|
||||||
|
|
||||||
|
let n = m.width;
|
||||||
|
let mut lu = m.data.to_vec();
|
||||||
|
let mut perm: Vec<usize> = (0..n).collect();
|
||||||
|
let mut sign = 1.0;
|
||||||
|
|
||||||
|
for col in 0..n {
|
||||||
|
// Partial pivot: take the largest-magnitude candidate to limit
|
||||||
|
// growth of round-off in the elimination below.
|
||||||
|
let mut pivot_row = col;
|
||||||
|
let mut pivot_max = lu[col * n + col].abs();
|
||||||
|
|
||||||
|
for row in (col + 1)..n {
|
||||||
|
let candidate = lu[row * n + col].abs();
|
||||||
|
if candidate > pivot_max {
|
||||||
|
pivot_max = candidate;
|
||||||
|
pivot_row = row;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if pivot_max == 0.0 {
|
||||||
|
sign = 0.0;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if pivot_row != col {
|
||||||
|
for k in 0..n {
|
||||||
|
lu.swap(col * n + k, pivot_row * n + k);
|
||||||
|
}
|
||||||
|
perm.swap(col, pivot_row);
|
||||||
|
sign = -sign;
|
||||||
|
}
|
||||||
|
|
||||||
|
let pivot = lu[col * n + col];
|
||||||
|
|
||||||
|
for row in (col + 1)..n {
|
||||||
|
let factor = lu[row * n + col] / pivot;
|
||||||
|
lu[row * n + col] = factor;
|
||||||
|
|
||||||
|
for k in (col + 1)..n {
|
||||||
|
lu[row * n + k] -= factor * lu[col * n + k];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Self { lu, perm, n, sign }
|
||||||
|
}
|
||||||
|
|
||||||
|
fn determinant(&self) -> f64 {
|
||||||
|
if self.sign == 0.0 {
|
||||||
|
return 0.0;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut det = self.sign;
|
||||||
|
for i in 0..self.n {
|
||||||
|
det *= self.lu[i * self.n + i];
|
||||||
|
}
|
||||||
|
|
||||||
|
det
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Solve `Ax = b` for a single column of the identity, giving one column
|
||||||
|
/// of the inverse.
|
||||||
|
fn solve_column(&self, col: usize, out: &mut [f64]) {
|
||||||
|
let n = self.n;
|
||||||
|
|
||||||
|
// Forward substitution through L, applying the row permutation.
|
||||||
|
for i in 0..n {
|
||||||
|
let mut sum = if self.perm[i] == col { 1.0 } else { 0.0 };
|
||||||
|
for (k, &solved) in out.iter().enumerate().take(i) {
|
||||||
|
sum -= self.lu[i * n + k] * solved;
|
||||||
|
}
|
||||||
|
out[i] = sum;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Back substitution through U.
|
||||||
|
for i in (0..n).rev() {
|
||||||
|
let mut sum = out[i];
|
||||||
|
for (k, &solved) in out.iter().enumerate().skip(i + 1) {
|
||||||
|
sum -= self.lu[i * n + k] * solved;
|
||||||
|
}
|
||||||
|
out[i] = sum / self.lu[i * n + i];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl Matrix {
|
impl Matrix {
|
||||||
pub fn new(height: usize, width: usize) -> Matrix {
|
pub fn new(height: usize, width: usize) -> Matrix {
|
||||||
Matrix {
|
Matrix {
|
||||||
@@ -52,73 +137,59 @@ impl Matrix {
|
|||||||
matrix
|
matrix
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn minor(&self, row_n: usize, col_n: usize) -> Matrix {
|
/// Determinant of a square matrix. The 0×0 determinant is 1 by convention
|
||||||
let mut matrix = Matrix::new(self.height - 1, self.width - 1);
|
/// (the empty product).
|
||||||
|
///
|
||||||
let mut nr = 0;
|
/// # Panics
|
||||||
|
///
|
||||||
for r in 0..self.height {
|
/// Panics if the matrix is not square.
|
||||||
if r == row_n {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut nc = 0;
|
|
||||||
|
|
||||||
for c in 0..self.width {
|
|
||||||
if c == col_n {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
matrix[(nr, nc)] = self[(r, c)];
|
|
||||||
|
|
||||||
nc += 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
nr += 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
matrix
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn determinant(&self) -> f64 {
|
pub fn determinant(&self) -> f64 {
|
||||||
debug_assert!(self.width == self.height);
|
assert_eq!(
|
||||||
|
self.width, self.height,
|
||||||
|
"determinant requires a square matrix, got {}x{}",
|
||||||
|
self.height, self.width
|
||||||
|
);
|
||||||
|
|
||||||
det(&self.data, self.width)
|
if self.width == 0 {
|
||||||
|
return 1.0;
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn adjugate(&self) -> Matrix {
|
Lu::decompose(self).determinant()
|
||||||
debug_assert!(self.width == self.height);
|
|
||||||
|
|
||||||
let mut matrix = Matrix::new(self.height, self.width);
|
|
||||||
|
|
||||||
if matrix.height == 2 {
|
|
||||||
matrix[(0, 0)] = self[(1, 1)];
|
|
||||||
matrix[(0, 1)] = -self[(0, 1)];
|
|
||||||
matrix[(1, 0)] = -self[(1, 0)];
|
|
||||||
matrix[(1, 1)] = self[(0, 0)];
|
|
||||||
} else {
|
|
||||||
for r in 0..matrix.height {
|
|
||||||
for c in 0..matrix.width {
|
|
||||||
let sign = if (r + c) % 2 == 0 { 1.0 } else { -1.0 };
|
|
||||||
|
|
||||||
matrix[(r, c)] = self.minor(r, c).determinant() * sign;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
matrix
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Matrix inverse via LU decomposition.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if the matrix is not square or is singular.
|
||||||
pub fn inverse(&self) -> Matrix {
|
pub fn inverse(&self) -> Matrix {
|
||||||
let mut matrix = Matrix::new(self.width, self.height);
|
assert_eq!(
|
||||||
|
self.width, self.height,
|
||||||
|
"inverse requires a square matrix, got {}x{}",
|
||||||
|
self.height, self.width
|
||||||
|
);
|
||||||
|
|
||||||
if self.height == self.width && self.height == 1 {
|
let n = self.width;
|
||||||
matrix[(0, 0)] = 1.0 / self[(0, 0)];
|
let mut inverse = Matrix::new(n, n);
|
||||||
} else {
|
|
||||||
panic!("eh, okey")
|
if n == 0 {
|
||||||
|
return inverse;
|
||||||
}
|
}
|
||||||
|
|
||||||
matrix
|
let lu = Lu::decompose(self);
|
||||||
|
assert!(lu.sign != 0.0, "cannot invert a singular matrix");
|
||||||
|
|
||||||
|
let mut column = vec![0.0; n];
|
||||||
|
|
||||||
|
for c in 0..n {
|
||||||
|
lu.solve_column(c, &mut column);
|
||||||
|
|
||||||
|
for (r, &value) in column.iter().enumerate() {
|
||||||
|
inverse[(r, c)] = value;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
inverse
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -126,20 +197,62 @@ impl ops::Index<(usize, usize)> for Matrix {
|
|||||||
type Output = f64;
|
type Output = f64;
|
||||||
|
|
||||||
fn index(&self, pos: (usize, usize)) -> &Self::Output {
|
fn index(&self, pos: (usize, usize)) -> &Self::Output {
|
||||||
|
debug_assert!(
|
||||||
|
pos.0 < self.height && pos.1 < self.width,
|
||||||
|
"index ({}, {}) out of bounds for {}x{} matrix",
|
||||||
|
pos.0,
|
||||||
|
pos.1,
|
||||||
|
self.height,
|
||||||
|
self.width
|
||||||
|
);
|
||||||
|
|
||||||
&self.data[(self.width * pos.0) + pos.1]
|
&self.data[(self.width * pos.0) + pos.1]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ops::IndexMut<(usize, usize)> for Matrix {
|
impl ops::IndexMut<(usize, usize)> for Matrix {
|
||||||
fn index_mut(&mut self, pos: (usize, usize)) -> &mut Self::Output {
|
fn index_mut(&mut self, pos: (usize, usize)) -> &mut Self::Output {
|
||||||
|
debug_assert!(
|
||||||
|
pos.0 < self.height && pos.1 < self.width,
|
||||||
|
"index ({}, {}) out of bounds for {}x{} matrix",
|
||||||
|
pos.0,
|
||||||
|
pos.1,
|
||||||
|
self.height,
|
||||||
|
self.width
|
||||||
|
);
|
||||||
|
|
||||||
&mut self.data[(self.width * pos.0) + pos.1]
|
&mut self.data[(self.width * pos.0) + pos.1]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> ops::Mul<&'a Matrix> for f64 {
|
fn multiply(lhs: &Matrix, rhs: &Matrix) -> Matrix {
|
||||||
|
assert_eq!(
|
||||||
|
lhs.width, rhs.height,
|
||||||
|
"cannot multiply {}x{} by {}x{}",
|
||||||
|
lhs.height, lhs.width, rhs.height, rhs.width
|
||||||
|
);
|
||||||
|
|
||||||
|
let mut matrix = Matrix::new(lhs.height, rhs.width);
|
||||||
|
|
||||||
|
for r in 0..matrix.height {
|
||||||
|
for c in 0..matrix.width {
|
||||||
|
let mut value = 0.0;
|
||||||
|
|
||||||
|
for x in 0..lhs.width {
|
||||||
|
value += lhs[(r, x)] * rhs[(x, c)];
|
||||||
|
}
|
||||||
|
|
||||||
|
matrix[(r, c)] = value;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
matrix
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ops::Mul<&Matrix> for f64 {
|
||||||
type Output = Matrix;
|
type Output = Matrix;
|
||||||
|
|
||||||
fn mul(self, rhs: &'a Matrix) -> Matrix {
|
fn mul(self, rhs: &Matrix) -> Matrix {
|
||||||
let mut matrix = Matrix::new(rhs.height, rhs.width);
|
let mut matrix = Matrix::new(rhs.height, rhs.width);
|
||||||
|
|
||||||
for r in 0..rhs.height {
|
for r in 0..rhs.height {
|
||||||
@@ -152,54 +265,35 @@ impl<'a> ops::Mul<&'a Matrix> for f64 {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> ops::Mul<&'a Matrix> for Matrix {
|
impl ops::Mul<&Matrix> for Matrix {
|
||||||
type Output = Matrix;
|
type Output = Matrix;
|
||||||
|
|
||||||
fn mul(self, rhs: &'a Matrix) -> Matrix {
|
fn mul(self, rhs: &Matrix) -> Matrix {
|
||||||
let mut matrix = Matrix::new(self.height, rhs.width);
|
multiply(&self, rhs)
|
||||||
|
|
||||||
for r in 0..matrix.height {
|
|
||||||
for c in 0..matrix.width {
|
|
||||||
let mut value = 0.0;
|
|
||||||
|
|
||||||
for x in 0..self.width {
|
|
||||||
value += self[(r, x)] * rhs[(x, c)];
|
|
||||||
}
|
|
||||||
|
|
||||||
matrix[(r, c)] = value;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
matrix
|
impl ops::Mul<&Matrix> for &Matrix {
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<'a> ops::Mul<&'a Matrix> for &'a Matrix {
|
|
||||||
type Output = Matrix;
|
type Output = Matrix;
|
||||||
|
|
||||||
fn mul(self, rhs: &'a Matrix) -> Matrix {
|
fn mul(self, rhs: &Matrix) -> Matrix {
|
||||||
let mut matrix = Matrix::new(self.height, rhs.width);
|
multiply(self, rhs)
|
||||||
|
|
||||||
for r in 0..matrix.height {
|
|
||||||
for c in 0..matrix.width {
|
|
||||||
let mut value = 0.0;
|
|
||||||
|
|
||||||
for x in 0..self.width {
|
|
||||||
value += self[(r, x)] * rhs[(x, c)];
|
|
||||||
}
|
|
||||||
|
|
||||||
matrix[(r, c)] = value;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
matrix
|
impl ops::Add<&Matrix> for &Matrix {
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<'a> ops::Add<&'a Matrix> for &'a Matrix {
|
|
||||||
type Output = Matrix;
|
type Output = Matrix;
|
||||||
|
|
||||||
fn add(self, rhs: &'a Matrix) -> Matrix {
|
fn add(self, rhs: &Matrix) -> Matrix {
|
||||||
|
assert!(
|
||||||
|
self.height == rhs.height && self.width == rhs.width,
|
||||||
|
"cannot add {}x{} to {}x{}",
|
||||||
|
self.height,
|
||||||
|
self.width,
|
||||||
|
rhs.height,
|
||||||
|
rhs.width
|
||||||
|
);
|
||||||
|
|
||||||
let mut matrix = Matrix::new(self.height, self.width);
|
let mut matrix = Matrix::new(self.height, self.width);
|
||||||
|
|
||||||
for r in 0..matrix.height {
|
for r in 0..matrix.height {
|
||||||
@@ -211,3 +305,106 @@ impl<'a> ops::Add<&'a Matrix> for &'a Matrix {
|
|||||||
matrix
|
matrix
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
fn from_rows(rows: &[&[f64]]) -> Matrix {
|
||||||
|
let mut m = Matrix::new(rows.len(), rows[0].len());
|
||||||
|
for (r, row) in rows.iter().enumerate() {
|
||||||
|
for (c, &v) in row.iter().enumerate() {
|
||||||
|
m[(r, c)] = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
m
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn determinant_1x1() {
|
||||||
|
assert!((from_rows(&[&[3.0]]).determinant() - 3.0).abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn determinant_2x2() {
|
||||||
|
let m = from_rows(&[&[1.0, 2.0], &[3.0, 4.0]]);
|
||||||
|
assert!((m.determinant() - (-2.0)).abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn determinant_3x3() {
|
||||||
|
let m = from_rows(&[&[6.0, 1.0, 1.0], &[4.0, -2.0, 5.0], &[2.0, 8.0, 7.0]]);
|
||||||
|
assert!((m.determinant() - (-306.0)).abs() < 1e-10);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn determinant_requires_no_pivot_at_origin() {
|
||||||
|
// A zero in the top-left forces a row swap; the sign must follow.
|
||||||
|
let m = from_rows(&[&[0.0, 1.0], &[1.0, 0.0]]);
|
||||||
|
assert!((m.determinant() - (-1.0)).abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn determinant_of_singular_is_zero() {
|
||||||
|
let m = from_rows(&[&[1.0, 2.0], &[2.0, 4.0]]);
|
||||||
|
assert!(m.determinant().abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn inverse_1x1() {
|
||||||
|
let inv = from_rows(&[&[4.0]]).inverse();
|
||||||
|
assert!((inv[(0, 0)] - 0.25).abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn inverse_times_original_is_identity() {
|
||||||
|
for rows in [
|
||||||
|
vec![vec![1.0, 2.0], vec![3.0, 4.0]],
|
||||||
|
vec![
|
||||||
|
vec![6.0, 1.0, 1.0],
|
||||||
|
vec![4.0, -2.0, 5.0],
|
||||||
|
vec![2.0, 8.0, 7.0],
|
||||||
|
],
|
||||||
|
vec![
|
||||||
|
vec![2.0, 0.0, 1.0, 3.0],
|
||||||
|
vec![1.0, 5.0, 2.0, 0.0],
|
||||||
|
vec![0.0, 1.0, 4.0, 1.0],
|
||||||
|
vec![3.0, 2.0, 0.0, 6.0],
|
||||||
|
],
|
||||||
|
] {
|
||||||
|
let refs: Vec<&[f64]> = rows.iter().map(|r| r.as_slice()).collect();
|
||||||
|
let m = from_rows(&refs);
|
||||||
|
let product = &m * &m.inverse();
|
||||||
|
|
||||||
|
for r in 0..product.height {
|
||||||
|
for c in 0..product.width {
|
||||||
|
let expected = if r == c { 1.0 } else { 0.0 };
|
||||||
|
assert!(
|
||||||
|
(product[(r, c)] - expected).abs() < 1e-9,
|
||||||
|
"({r},{c}) = {} expected {expected}",
|
||||||
|
product[(r, c)]
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "singular")]
|
||||||
|
fn inverse_of_singular_panics() {
|
||||||
|
let _ = from_rows(&[&[1.0, 2.0], &[2.0, 4.0]]).inverse();
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_determinant_is_one() {
|
||||||
|
assert!((Matrix::new(0, 0).determinant() - 1.0).abs() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn transpose_round_trips() {
|
||||||
|
let m = from_rows(&[&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
|
||||||
|
let t = m.transpose();
|
||||||
|
assert_eq!((t.height, t.width), (3, 2));
|
||||||
|
assert_eq!(t.transpose()[(1, 2)], m[(1, 2)]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+85
-2
@@ -14,13 +14,95 @@ 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) {}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Shared and boxed observers forward to what they point at.
|
||||||
|
///
|
||||||
|
/// `History` takes its observer by value, so a caller who wants to *read* what
|
||||||
|
/// an observer recorded has to keep a handle to it. Without these impls the
|
||||||
|
/// natural spelling does not compile:
|
||||||
|
///
|
||||||
|
/// ```
|
||||||
|
/// # use std::sync::{Arc, Mutex};
|
||||||
|
/// # use trueskill_tt::{History, Observer};
|
||||||
|
/// #[derive(Default)]
|
||||||
|
/// struct Recorder {
|
||||||
|
/// iterations: Mutex<Vec<usize>>,
|
||||||
|
/// }
|
||||||
|
///
|
||||||
|
/// impl Observer<i64> for Recorder {
|
||||||
|
/// fn on_iteration_end(&self, iter: usize, _step: (f64, f64)) {
|
||||||
|
/// self.iterations.lock().unwrap().push(iter);
|
||||||
|
/// }
|
||||||
|
/// }
|
||||||
|
///
|
||||||
|
/// let recorder = Arc::new(Recorder::default());
|
||||||
|
/// let mut h = History::builder().observer(Arc::clone(&recorder)).build();
|
||||||
|
/// h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
/// h.converge().unwrap();
|
||||||
|
///
|
||||||
|
/// // The caller's handle sees what the history's copy recorded.
|
||||||
|
/// assert!(!recorder.iterations.lock().unwrap().is_empty());
|
||||||
|
/// ```
|
||||||
|
///
|
||||||
|
/// The alternative was for every observer to wrap each of its own fields in an
|
||||||
|
/// `Arc` and derive `Clone` — one allocation and one lock per field, and a
|
||||||
|
/// pattern each implementor had to rediscover.
|
||||||
|
///
|
||||||
|
/// `?Sized` is deliberate: it makes `Arc<dyn Observer<T>>` and
|
||||||
|
/// `Box<dyn Observer<T>>` work, so observers can be chosen at runtime.
|
||||||
|
impl<T: Time, O: Observer<T> + ?Sized> Observer<T> for std::sync::Arc<O> {
|
||||||
|
fn on_iteration_end(&self, iter: usize, max_step: (f64, f64)) {
|
||||||
|
(**self).on_iteration_end(iter, max_step);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_slice_processed(&self, time: &T, slice_idx: usize, n_events: usize) {
|
||||||
|
(**self).on_slice_processed(time, slice_idx, n_events);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_converged(&self, iters: usize, final_step: (f64, f64), converged: bool) {
|
||||||
|
(**self).on_converged(iters, final_step, converged);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T: Time, O: Observer<T> + ?Sized> Observer<T> for Box<O> {
|
||||||
|
fn on_iteration_end(&self, iter: usize, max_step: (f64, f64)) {
|
||||||
|
(**self).on_iteration_end(iter, max_step);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_slice_processed(&self, time: &T, slice_idx: usize, n_events: usize) {
|
||||||
|
(**self).on_slice_processed(time, slice_idx, n_events);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_converged(&self, iters: usize, final_step: (f64, f64), converged: bool) {
|
||||||
|
(**self).on_converged(iters, final_step, converged);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T: Time, O: Observer<T> + ?Sized> Observer<T> for &O {
|
||||||
|
fn on_iteration_end(&self, iter: usize, max_step: (f64, f64)) {
|
||||||
|
(**self).on_iteration_end(iter, max_step);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_slice_processed(&self, time: &T, slice_idx: usize, n_events: usize) {
|
||||||
|
(**self).on_slice_processed(time, slice_idx, n_events);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_converged(&self, iters: usize, final_step: (f64, f64), converged: bool) {
|
||||||
|
(**self).on_converged(iters, final_step, converged);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// ZST no-op observer; the default when none is configured.
|
/// ZST no-op observer; the default when none is configured.
|
||||||
#[derive(Copy, Clone, Debug, Default)]
|
#[derive(Copy, Clone, Debug, Default)]
|
||||||
pub struct NullObserver;
|
pub struct NullObserver;
|
||||||
@@ -35,6 +117,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);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+21
-5
@@ -29,7 +29,13 @@ pub enum Outcome {
|
|||||||
impl Outcome {
|
impl Outcome {
|
||||||
/// `n`-team outcome where team `winner` won and everyone else tied for last.
|
/// `n`-team outcome where team `winner` won and everyone else tied for last.
|
||||||
///
|
///
|
||||||
|
/// Note this ties every loser, so for `n >= 3` it needs a positive
|
||||||
|
/// `p_draw` — see `InferenceError::TieWithoutDrawProbability`.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
/// Panics if `winner >= n`.
|
/// Panics if `winner >= n`.
|
||||||
|
#[must_use]
|
||||||
pub fn winner(winner: u32, n: u32) -> Self {
|
pub fn winner(winner: u32, n: u32) -> Self {
|
||||||
assert!(winner < n, "winner index {winner} out of range 0..{n}");
|
assert!(winner < n, "winner index {winner} out of range 0..{n}");
|
||||||
let ranks: SmallVec<[u32; 4]> = (0..n).map(|i| if i == winner { 0 } else { 1 }).collect();
|
let ranks: SmallVec<[u32; 4]> = (0..n).map(|i| if i == winner { 0 } else { 1 }).collect();
|
||||||
@@ -37,6 +43,7 @@ impl Outcome {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// All `n` teams tied.
|
/// All `n` teams tied.
|
||||||
|
#[must_use]
|
||||||
pub fn draw(n: u32) -> Self {
|
pub fn draw(n: u32) -> Self {
|
||||||
Self::Ranked(SmallVec::from_vec(vec![0; n as usize]))
|
Self::Ranked(SmallVec::from_vec(vec![0; n as usize]))
|
||||||
}
|
}
|
||||||
@@ -57,15 +64,18 @@ impl Outcome {
|
|||||||
|
|
||||||
/// Explicit per-team continuous scores with a per-event noise override.
|
/// Explicit per-team continuous scores with a per-event noise override.
|
||||||
///
|
///
|
||||||
/// `sigma` must be `> 0.0`; debug-asserts otherwise.
|
/// `sigma` must be `> 0.0`. Constructing an `Outcome` with a non-positive
|
||||||
|
/// or NaN sigma is allowed; the value is rejected with
|
||||||
|
/// `InferenceError::InvalidParameter` when the event is ingested, so
|
||||||
|
/// callers get an error rather than a panic.
|
||||||
pub fn scores_with_sigma<I: IntoIterator<Item = f64>>(scores: I, sigma: f64) -> Self {
|
pub fn scores_with_sigma<I: IntoIterator<Item = f64>>(scores: I, sigma: f64) -> Self {
|
||||||
debug_assert!(sigma > 0.0, "score_sigma must be > 0.0 (got {sigma})");
|
|
||||||
Self::Scored {
|
Self::Scored {
|
||||||
scores: scores.into_iter().collect(),
|
scores: scores.into_iter().collect(),
|
||||||
sigma: Some(sigma),
|
sigma: Some(sigma),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn team_count(&self) -> usize {
|
pub fn team_count(&self) -> usize {
|
||||||
match self {
|
match self {
|
||||||
Self::Ranked(r) => r.len(),
|
Self::Ranked(r) => r.len(),
|
||||||
@@ -169,9 +179,15 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Construction accepts any sigma; the value is validated at ingestion so
|
||||||
|
/// callers receive an `InferenceError` rather than a panic. See
|
||||||
|
/// `tests/degenerate_inputs.rs::scored_event_rejects_non_positive_sigma`.
|
||||||
#[test]
|
#[test]
|
||||||
#[should_panic(expected = "score_sigma must be > 0.0")]
|
fn scores_with_sigma_defers_validation_to_ingestion() {
|
||||||
fn scores_with_sigma_rejects_zero() {
|
let o = Outcome::scores_with_sigma([3.0, 1.0], 0.0);
|
||||||
let _ = Outcome::scores_with_sigma([3.0, 1.0], 0.0);
|
match o {
|
||||||
|
Outcome::Scored { sigma, .. } => assert_eq!(sigma, Some(0.0)),
|
||||||
|
Outcome::Ranked(_) => panic!("expected Scored variant"),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+729
@@ -0,0 +1,729 @@
|
|||||||
|
//! Outcome prediction: who wins, and how likely is a given finishing order.
|
||||||
|
//!
|
||||||
|
//! Prediction runs on *performances*, not skills. A competitor's skill is
|
||||||
|
//! inflated by their performance noise `beta` before any comparison, which is
|
||||||
|
//! what separates "how good are they" from "how will they do today".
|
||||||
|
//!
|
||||||
|
//! Two questions, two algorithms:
|
||||||
|
//!
|
||||||
|
//! - **Who finishes first.** Because performances are independent Gaussians,
|
||||||
|
//! the probability that team `i` beats every other team separates into a
|
||||||
|
//! *one-dimensional* integral — no multivariate orthant integral is
|
||||||
|
//! involved. [`quadrature::integrate`] evaluates it to near machine
|
||||||
|
//! precision for a few hundred `cdf` calls.
|
||||||
|
//! - **A specific finishing order.** The factor graph only ever constrains
|
||||||
|
//! rank-*adjacent* teams (see `Game::run_chain`), so the joint probability
|
||||||
|
//! of a full order is a chain of local constraints rather than a general
|
||||||
|
//! orthant probability. That chain collapses into a sequential recursion:
|
||||||
|
//! one cumulative integral per adjacent pair, `O(teams * grid)` overall.
|
||||||
|
//!
|
||||||
|
//! Both are deterministic. A sampler would have been easier to write and
|
||||||
|
//! would have made every `predict_*` call return a slightly different number,
|
||||||
|
//! which is not a property a rating library should have.
|
||||||
|
|
||||||
|
use crate::{Gaussian, quadrature};
|
||||||
|
|
||||||
|
/// Teams beyond this count make the outcome enumeration impractical.
|
||||||
|
///
|
||||||
|
/// Each realisation sorts into exactly one (permutation, tie-pattern) event,
|
||||||
|
/// so the space has `n! * 2^(n-1)` members: 24 at 3 teams, 192 at 4, 1_920 at
|
||||||
|
/// 5, 23_040 at 6. The jump to 322_560 at 7 is where enumerating stops being
|
||||||
|
/// a reasonable thing to do on a caller's behalf.
|
||||||
|
pub(crate) const MAX_TEAMS_FOR_DISTRIBUTION: usize = 6;
|
||||||
|
|
||||||
|
/// Relative tolerance for the first-place integrals.
|
||||||
|
///
|
||||||
|
/// The adaptive integrator reaches the exact two-team closed form to ~1e-15 at
|
||||||
|
/// this tolerance, which is round-off for a probability. `cdf` is no longer the
|
||||||
|
/// limit — it went to ~1 ULP when `erfc` moved to `libm` — so this is the
|
||||||
|
/// integrator's own floor.
|
||||||
|
const WIN_TOLERANCE: f64 = 1e-8;
|
||||||
|
|
||||||
|
/// Nodes for the ranking grid, and the floor below which a grid is pointless.
|
||||||
|
///
|
||||||
|
/// The recursion converges as O(h^2), so this trades nodes against accuracy
|
||||||
|
/// directly. Measured against the exact two-team closed form, 2_048 nodes leave
|
||||||
|
/// ~1.2e-6 of discretisation error and 8_192 reach ~1e-7.
|
||||||
|
///
|
||||||
|
/// Unlike the adaptive path there is no approximation floor underneath this any
|
||||||
|
/// more — `cdf` is accurate to ~1 ULP since `erfc` moved to `libm` — so the
|
||||||
|
/// error here is purely the grid, and a caller who needs more can only get it
|
||||||
|
/// by paying for more nodes. 8_192 is the accuracy/cost point chosen, not a
|
||||||
|
/// point where refining stops helping.
|
||||||
|
const MIN_GRID_POINTS: usize = 8_192;
|
||||||
|
const MAX_GRID_POINTS: usize = 262_144;
|
||||||
|
|
||||||
|
/// How many standard deviations of support the grid and integrals cover.
|
||||||
|
///
|
||||||
|
/// The normal density is below 1e-18 of its peak past nine sigma, far under
|
||||||
|
/// the precision of everything else here.
|
||||||
|
const SUPPORT_SIGMAS: f64 = 9.0;
|
||||||
|
|
||||||
|
/// Standard normal CDF at `z`.
|
||||||
|
fn phi(z: f64) -> f64 {
|
||||||
|
crate::cdf(z, 0.0, 1.0)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Normal density of `x` under `g`.
|
||||||
|
fn density(g: Gaussian, x: f64) -> f64 {
|
||||||
|
let sigma = g.sigma();
|
||||||
|
let z = (x - g.mu()) / sigma;
|
||||||
|
libm::exp(-0.5 * z * z) / (sigma * (2.0 * std::f64::consts::PI).sqrt())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Per-pair draw margins.
|
||||||
|
///
|
||||||
|
/// The margin is *not* a single number for the whole game: inference derives
|
||||||
|
/// it per rank-adjacent pair from those two teams' betas (`Game::likelihoods`).
|
||||||
|
/// Prediction has to use the same per-pair values or it answers a question
|
||||||
|
/// about a different model than the one that will actually be fitted.
|
||||||
|
pub(crate) struct Margins {
|
||||||
|
n: usize,
|
||||||
|
values: Vec<f64>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Margins {
|
||||||
|
/// Build from a per-pair margin function.
|
||||||
|
pub(crate) fn new<F: Fn(usize, usize) -> f64>(n: usize, f: F) -> Self {
|
||||||
|
let mut values = vec![0.0; n * n];
|
||||||
|
for i in 0..n {
|
||||||
|
for j in 0..n {
|
||||||
|
if i != j {
|
||||||
|
values[i * n + j] = f(i, j);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Self { n, values }
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get(&self, i: usize, j: usize) -> f64 {
|
||||||
|
self.values[i * self.n + j]
|
||||||
|
}
|
||||||
|
|
||||||
|
/// True when no pair can draw, so every tie has probability zero.
|
||||||
|
fn all_zero(&self) -> bool {
|
||||||
|
self.values.iter().all(|&v| v == 0.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `P(team i finishes strictly first)` for every team.
|
||||||
|
///
|
||||||
|
/// Strictly means beating each rival by more than that pair's draw margin, so
|
||||||
|
/// with a non-zero margin these sum to less than one; the shortfall is the
|
||||||
|
/// probability that the top place is shared.
|
||||||
|
pub(crate) fn win_probabilities(perf: &[Gaussian], margins: &Margins) -> Vec<f64> {
|
||||||
|
(0..perf.len())
|
||||||
|
.map(|i| {
|
||||||
|
let (mu, sigma) = (perf[i].mu(), perf[i].sigma());
|
||||||
|
let (lo, hi) = (mu - SUPPORT_SIGMAS * sigma, mu + SUPPORT_SIGMAS * sigma);
|
||||||
|
|
||||||
|
// Each rival's CDF turns over near its own mean plus the margin.
|
||||||
|
// Seeding there is what keeps a rival with a tiny sigma — a step
|
||||||
|
// function in disguise — from being stepped over.
|
||||||
|
let mut seeds = Vec::with_capacity(3 * perf.len());
|
||||||
|
for (j, rival) in perf.iter().enumerate().filter(|&(j, _)| j != i) {
|
||||||
|
let centre = rival.mu() + margins.get(i, j);
|
||||||
|
seeds.extend_from_slice(&[centre - rival.sigma(), centre, centre + rival.sigma()]);
|
||||||
|
}
|
||||||
|
|
||||||
|
quadrature::integrate(
|
||||||
|
|x| {
|
||||||
|
let d = density(perf[i], x);
|
||||||
|
if d == 0.0 {
|
||||||
|
return 0.0;
|
||||||
|
}
|
||||||
|
let beaten: f64 = (0..perf.len())
|
||||||
|
.filter(|&j| j != i)
|
||||||
|
.map(|j| phi((x - margins.get(i, j) - perf[j].mu()) / perf[j].sigma()))
|
||||||
|
.product();
|
||||||
|
d * beaten
|
||||||
|
},
|
||||||
|
lo,
|
||||||
|
hi,
|
||||||
|
&seeds,
|
||||||
|
WIN_TOLERANCE,
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Grid bounds and resolution covering every team's support.
|
||||||
|
///
|
||||||
|
/// Resolution is set by the *smallest* feature in play — the narrowest sigma,
|
||||||
|
/// or a draw margin narrower still — because that is what the recursion has to
|
||||||
|
/// resolve. A grid sized off the widest team would step over the narrow one.
|
||||||
|
fn grid_shape(perf: &[Gaussian], margins: &Margins) -> (f64, f64, usize) {
|
||||||
|
let lo = perf
|
||||||
|
.iter()
|
||||||
|
.map(|g| g.mu() - SUPPORT_SIGMAS * g.sigma())
|
||||||
|
.fold(f64::INFINITY, f64::min);
|
||||||
|
let hi = perf
|
||||||
|
.iter()
|
||||||
|
.map(|g| g.mu() + SUPPORT_SIGMAS * g.sigma())
|
||||||
|
.fold(f64::NEG_INFINITY, f64::max);
|
||||||
|
|
||||||
|
let narrowest = perf
|
||||||
|
.iter()
|
||||||
|
.map(Gaussian::sigma)
|
||||||
|
.fold(f64::INFINITY, f64::min);
|
||||||
|
let smallest_margin = margins
|
||||||
|
.values
|
||||||
|
.iter()
|
||||||
|
.copied()
|
||||||
|
.filter(|&m| m > 0.0)
|
||||||
|
.fold(f64::INFINITY, f64::min);
|
||||||
|
|
||||||
|
let feature = narrowest.min(smallest_margin);
|
||||||
|
let wanted = if feature.is_finite() && feature > 0.0 {
|
||||||
|
((hi - lo) / (feature / 12.0)).ceil()
|
||||||
|
} else {
|
||||||
|
MIN_GRID_POINTS as f64
|
||||||
|
};
|
||||||
|
|
||||||
|
let points = if wanted.is_finite() {
|
||||||
|
(wanted as usize).clamp(MIN_GRID_POINTS, MAX_GRID_POINTS)
|
||||||
|
} else {
|
||||||
|
MIN_GRID_POINTS
|
||||||
|
};
|
||||||
|
|
||||||
|
(lo, hi, points)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Densities of each team sampled on the shared grid.
|
||||||
|
struct Sampled {
|
||||||
|
lo: f64,
|
||||||
|
step: f64,
|
||||||
|
points: usize,
|
||||||
|
density: Vec<Vec<f64>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Sampled {
|
||||||
|
fn new(perf: &[Gaussian], margins: &Margins) -> Self {
|
||||||
|
let (lo, hi, points) = grid_shape(perf, margins);
|
||||||
|
let step = (hi - lo) / (points - 1) as f64;
|
||||||
|
let density = perf
|
||||||
|
.iter()
|
||||||
|
.map(|&g| {
|
||||||
|
(0..points)
|
||||||
|
.map(|i| density(g, lo + i as f64 * step))
|
||||||
|
.collect()
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
Self {
|
||||||
|
lo,
|
||||||
|
step,
|
||||||
|
points,
|
||||||
|
density,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn node(&self, i: usize) -> f64 {
|
||||||
|
self.lo + i as f64 * self.step
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `P(order[0] >= order[1] >= ... )` with the given adjacency pattern.
|
||||||
|
///
|
||||||
|
/// `tied[k]` says whether `order[k]` and `order[k + 1]` finish within that
|
||||||
|
/// pair's draw margin. The recursion runs bottom-up: `carry` holds, for each
|
||||||
|
/// grid node, the probability that everything *below* the current team holds
|
||||||
|
/// given that team landed on that node. A strict gap reads a cumulative
|
||||||
|
/// integral; a tie reads a window. Both are O(1) against one prefix array,
|
||||||
|
/// so each level costs O(grid) and the whole order costs O(teams * grid).
|
||||||
|
fn order_probability(margins: &Margins, sampled: &Sampled, order: &[usize], tied: &[bool]) -> f64 {
|
||||||
|
let mut carry = vec![1.0; sampled.points];
|
||||||
|
|
||||||
|
for k in (0..order.len() - 1).rev() {
|
||||||
|
let below = order[k + 1];
|
||||||
|
let above = order[k];
|
||||||
|
let margin = margins.get(above, below);
|
||||||
|
|
||||||
|
let integrand: Vec<f64> = (0..sampled.points)
|
||||||
|
.map(|i| sampled.density[below][i] * carry[i])
|
||||||
|
.collect();
|
||||||
|
let cumulative = quadrature::Grid::from_values(sampled.lo, sampled.step, integrand);
|
||||||
|
|
||||||
|
carry = (0..sampled.points)
|
||||||
|
.map(|i| {
|
||||||
|
let x = sampled.node(i);
|
||||||
|
if tied[k] {
|
||||||
|
// Sorted order already implies `below <= above`, so the
|
||||||
|
// tie window is one-sided: [x - margin, x].
|
||||||
|
cumulative.integral_between(x - margin, x)
|
||||||
|
} else {
|
||||||
|
cumulative.integral_to(x - margin)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
}
|
||||||
|
|
||||||
|
let top = order[0];
|
||||||
|
let integrand: Vec<f64> = (0..sampled.points)
|
||||||
|
.map(|i| sampled.density[top][i] * carry[i])
|
||||||
|
.collect();
|
||||||
|
quadrature::Grid::from_values(sampled.lo, sampled.step, integrand).total()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Dense ranks implied by a sorted order and its tie pattern.
|
||||||
|
fn ranks_of(order: &[usize], tied: &[bool], n: usize) -> Vec<u32> {
|
||||||
|
let mut ranks = vec![0u32; n];
|
||||||
|
let mut rank = 0u32;
|
||||||
|
ranks[order[0]] = 0;
|
||||||
|
for k in 0..order.len() - 1 {
|
||||||
|
if !tied[k] {
|
||||||
|
rank += 1;
|
||||||
|
}
|
||||||
|
ranks[order[k + 1]] = rank;
|
||||||
|
}
|
||||||
|
ranks
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Every (order, tie-pattern) event, or only the strict ones when no pair can
|
||||||
|
/// draw — a tie then has probability exactly zero and is not worth integrating.
|
||||||
|
fn events(n: usize, strict_only: bool) -> Vec<(Vec<usize>, Vec<bool>)> {
|
||||||
|
fn permute(current: &mut Vec<usize>, k: usize, out: &mut Vec<Vec<usize>>) {
|
||||||
|
if k == current.len() {
|
||||||
|
out.push(current.clone());
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
for i in k..current.len() {
|
||||||
|
current.swap(k, i);
|
||||||
|
permute(current, k + 1, out);
|
||||||
|
current.swap(k, i);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut orders = Vec::new();
|
||||||
|
permute(&mut (0..n).collect(), 0, &mut orders);
|
||||||
|
|
||||||
|
let patterns: Vec<Vec<bool>> = if strict_only {
|
||||||
|
vec![vec![false; n - 1]]
|
||||||
|
} else {
|
||||||
|
(0..(1u32 << (n - 1)))
|
||||||
|
.map(|mask| (0..n - 1).map(|i| mask >> i & 1 == 1).collect())
|
||||||
|
.collect()
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut out = Vec::with_capacity(orders.len() * patterns.len());
|
||||||
|
for order in orders {
|
||||||
|
for pattern in &patterns {
|
||||||
|
out.push((order.clone(), pattern.clone()));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The full distribution over finishing orders, aggregated by rank vector.
|
||||||
|
///
|
||||||
|
/// Orders that differ only *within* a tied group describe the same finishing
|
||||||
|
/// order, so their probabilities are summed into one entry.
|
||||||
|
pub(crate) fn outcome_distribution(perf: &[Gaussian], margins: &Margins) -> Vec<(Vec<u32>, f64)> {
|
||||||
|
let n = perf.len();
|
||||||
|
let sampled = Sampled::new(perf, margins);
|
||||||
|
|
||||||
|
let mut aggregated: Vec<(Vec<u32>, f64)> = Vec::new();
|
||||||
|
for (order, tied) in events(n, margins.all_zero()) {
|
||||||
|
let p = order_probability(margins, &sampled, &order, &tied);
|
||||||
|
let ranks = ranks_of(&order, &tied, n);
|
||||||
|
match aggregated.iter_mut().find(|(r, _)| *r == ranks) {
|
||||||
|
Some((_, acc)) => *acc += p,
|
||||||
|
None => aggregated.push((ranks, p)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
aggregated.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
|
aggregated
|
||||||
|
}
|
||||||
|
|
||||||
|
/// All permutations of `items`.
|
||||||
|
fn permutations(items: &[usize]) -> Vec<Vec<usize>> {
|
||||||
|
fn go(current: &mut Vec<usize>, k: usize, out: &mut Vec<Vec<usize>>) {
|
||||||
|
if k == current.len() {
|
||||||
|
out.push(current.clone());
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
for i in k..current.len() {
|
||||||
|
current.swap(k, i);
|
||||||
|
go(current, k + 1, out);
|
||||||
|
current.swap(k, i);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let mut out = Vec::new();
|
||||||
|
go(&mut items.to_vec(), 0, &mut out);
|
||||||
|
out
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Every (order, tie-pattern) event consistent with a grouping by rank.
|
||||||
|
///
|
||||||
|
/// Teams sharing a rank may finish in any internal order, so this is the
|
||||||
|
/// product of each group's permutations. Adjacencies inside a group are ties;
|
||||||
|
/// the adjacency joining one group to the next is not.
|
||||||
|
fn orders_for_groups(groups: &[Vec<usize>]) -> Vec<(Vec<usize>, Vec<bool>)> {
|
||||||
|
let per_group: Vec<Vec<Vec<usize>>> = groups.iter().map(|g| permutations(g)).collect();
|
||||||
|
|
||||||
|
let mut out = Vec::new();
|
||||||
|
let mut choice = vec![0usize; groups.len()];
|
||||||
|
|
||||||
|
loop {
|
||||||
|
let mut order = Vec::new();
|
||||||
|
let mut tied = Vec::new();
|
||||||
|
for (gi, group) in per_group.iter().enumerate() {
|
||||||
|
for (offset, &member) in group[choice[gi]].iter().enumerate() {
|
||||||
|
if !order.is_empty() {
|
||||||
|
tied.push(offset != 0);
|
||||||
|
}
|
||||||
|
order.push(member);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out.push((order, tied));
|
||||||
|
|
||||||
|
let mut k = 0;
|
||||||
|
loop {
|
||||||
|
if k == choice.len() {
|
||||||
|
return out;
|
||||||
|
}
|
||||||
|
choice[k] += 1;
|
||||||
|
if choice[k] < per_group[k].len() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
choice[k] = 0;
|
||||||
|
k += 1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Probability of one specific rank vector.
|
||||||
|
///
|
||||||
|
/// Ties in `ranks` mean the tied teams may finish in any internal order, so
|
||||||
|
/// this sums the orders consistent with the requested ranking rather than
|
||||||
|
/// picking one.
|
||||||
|
pub(crate) fn ranking_probability(perf: &[Gaussian], margins: &Margins, ranks: &[u32]) -> f64 {
|
||||||
|
let n = perf.len();
|
||||||
|
let sampled = Sampled::new(perf, margins);
|
||||||
|
|
||||||
|
let mut distinct: Vec<u32> = ranks.to_vec();
|
||||||
|
distinct.sort_unstable();
|
||||||
|
distinct.dedup();
|
||||||
|
|
||||||
|
let groups: Vec<Vec<usize>> = distinct
|
||||||
|
.iter()
|
||||||
|
.map(|&r| (0..n).filter(|&i| ranks[i] == r).collect())
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
orders_for_groups(&groups)
|
||||||
|
.iter()
|
||||||
|
.map(|(order, tied)| order_probability(margins, &sampled, order, tied))
|
||||||
|
.sum()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A distribution over the ways a contest could finish.
|
||||||
|
///
|
||||||
|
/// Each entry pairs a rank vector — the same shape [`crate::Outcome::ranking`]
|
||||||
|
/// takes, with equal ranks meaning a tie — against its probability. Entries
|
||||||
|
/// are ordered most likely first, and cover the whole outcome space, so the
|
||||||
|
/// probabilities sum to one.
|
||||||
|
///
|
||||||
|
/// The rank vectors compose directly with inference: feeding one to
|
||||||
|
/// `Game::ranked` asks "what would we believe if *this* happened", which is
|
||||||
|
/// what an expected-information-gain calculation needs alongside the weight.
|
||||||
|
#[derive(Clone, Debug, PartialEq)]
|
||||||
|
pub struct Prediction {
|
||||||
|
outcomes: Vec<(Vec<u32>, f64)>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Prediction {
|
||||||
|
pub(crate) fn new(outcomes: Vec<(Vec<u32>, f64)>) -> Self {
|
||||||
|
Self { outcomes }
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Every possible finishing order and its probability, most likely first.
|
||||||
|
pub fn outcomes(&self) -> impl ExactSizeIterator<Item = (&[u32], f64)> {
|
||||||
|
self.outcomes.iter().map(|(r, p)| (r.as_slice(), *p))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The single most likely finishing order.
|
||||||
|
#[must_use]
|
||||||
|
pub fn most_likely(&self) -> Option<(&[u32], f64)> {
|
||||||
|
self.outcomes.first().map(|(r, p)| (r.as_slice(), *p))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Probability of one specific finishing order, or zero if it cannot occur.
|
||||||
|
#[must_use]
|
||||||
|
pub fn probability_of(&self, ranks: &[u32]) -> f64 {
|
||||||
|
self.outcomes
|
||||||
|
.iter()
|
||||||
|
.find(|(r, _)| r.as_slice() == ranks)
|
||||||
|
.map_or(0.0, |(_, p)| *p)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `P(team i finishes strictly first)`, for each team.
|
||||||
|
///
|
||||||
|
/// Sums to less than one exactly when the top place can be shared; the
|
||||||
|
/// shortfall is [`Prediction::shared_first_place`].
|
||||||
|
#[must_use]
|
||||||
|
pub fn win_probabilities(&self) -> Vec<f64> {
|
||||||
|
let n = self.outcomes.first().map_or(0, |(r, _)| r.len());
|
||||||
|
let mut wins = vec![0.0; n];
|
||||||
|
for (ranks, p) in &self.outcomes {
|
||||||
|
let leaders = ranks.iter().filter(|&&r| r == 0).count();
|
||||||
|
if leaders == 1 {
|
||||||
|
let winner = ranks.iter().position(|&r| r == 0).expect("a rank-0 team");
|
||||||
|
wins[winner] += p;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wins
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Probability that two or more teams share first place.
|
||||||
|
#[must_use]
|
||||||
|
pub fn shared_first_place(&self) -> f64 {
|
||||||
|
self.outcomes
|
||||||
|
.iter()
|
||||||
|
.filter(|(r, _)| r.iter().filter(|&&x| x == 0).count() > 1)
|
||||||
|
.map(|(_, p)| p)
|
||||||
|
.sum()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Total probability mass, which should be one.
|
||||||
|
///
|
||||||
|
/// Exposed because it is a genuine check on the numerics rather than a
|
||||||
|
/// formality: the outcome space is exhaustive and disjoint by construction,
|
||||||
|
/// so any drift from one is integration error and nothing else.
|
||||||
|
#[must_use]
|
||||||
|
pub fn total(&self) -> f64 {
|
||||||
|
self.outcomes.iter().map(|(_, p)| p).sum()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
fn g(mu: f64, sigma: f64) -> Gaussian {
|
||||||
|
Gaussian::from_ms(mu, sigma)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn flat(n: usize, eps: f64) -> Margins {
|
||||||
|
Margins::new(n, |_, _| eps)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Exact two-team result: `P(a first) = Phi((mu_a - mu_b - eps) / sd)`.
|
||||||
|
fn closed_form_two(a: Gaussian, b: Gaussian, eps: f64) -> (f64, f64) {
|
||||||
|
let sd = a.sigma().hypot(b.sigma());
|
||||||
|
(
|
||||||
|
phi((a.mu() - b.mu() - eps) / sd),
|
||||||
|
phi((b.mu() - a.mu() - eps) / sd),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn two_team_win_probabilities_match_the_closed_form() {
|
||||||
|
for (ma, sa, mb, sb, eps) in [
|
||||||
|
(0.0, 6.0, 0.0, 6.0, 0.0),
|
||||||
|
(3.0, 6.0, -2.0, 1.0, 0.0),
|
||||||
|
(0.0, 6.0, 0.0, 6.0, 2.0),
|
||||||
|
(3.0, 6.0, -2.0, 1.0, 1.5),
|
||||||
|
(40.0, 1.0, 0.0, 1.0, 0.0),
|
||||||
|
] {
|
||||||
|
let perf = [g(ma, sa), g(mb, sb)];
|
||||||
|
let got = win_probabilities(&perf, &flat(2, eps));
|
||||||
|
let (wa, wb) = closed_form_two(perf[0], perf[1], eps);
|
||||||
|
assert!(
|
||||||
|
(got[0] - wa).abs() < 1e-12 && (got[1] - wb).abs() < 1e-12,
|
||||||
|
"mu=({ma},{mb}) sigma=({sa},{sb}) eps={eps}: got {got:?}, want [{wa}, {wb}]"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The identity that a wrong-but-plausible implementation cannot fake:
|
||||||
|
/// with no draw margin, exactly one team finishes first.
|
||||||
|
#[test]
|
||||||
|
fn win_probabilities_sum_to_one_without_a_draw_margin() {
|
||||||
|
for perf in [
|
||||||
|
vec![g(0.0, 6.0), g(0.0, 6.0)],
|
||||||
|
vec![g(5.0, 6.0), g(0.0, 3.0), g(-5.0, 1.0)],
|
||||||
|
vec![
|
||||||
|
g(8.0, 2.0),
|
||||||
|
g(3.0, 6.0),
|
||||||
|
g(0.0, 1.0),
|
||||||
|
g(-3.0, 4.0),
|
||||||
|
g(-8.0, 6.0),
|
||||||
|
],
|
||||||
|
] {
|
||||||
|
let sum: f64 = win_probabilities(&perf, &flat(perf.len(), 0.0))
|
||||||
|
.iter()
|
||||||
|
.sum();
|
||||||
|
assert!(
|
||||||
|
(sum - 1.0).abs() < 1e-7,
|
||||||
|
"{} teams: sum = {sum}",
|
||||||
|
perf.len()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A rival with a tiny sigma is a step function in disguise. Fixed-node
|
||||||
|
/// quadrature steps over it and lands ~1e-2 out while still looking like a
|
||||||
|
/// probability; this is the case that rules that approach out.
|
||||||
|
#[test]
|
||||||
|
fn win_probabilities_survive_a_rival_with_a_tiny_sigma() {
|
||||||
|
let perf = [g(0.0, 0.001), g(0.5, 6.0), g(-0.5, 6.0)];
|
||||||
|
let got = win_probabilities(&perf, &flat(3, 0.0));
|
||||||
|
let sum: f64 = got.iter().sum();
|
||||||
|
assert!((sum - 1.0).abs() < 1e-6, "sum = {sum}, probs = {got:?}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn a_stronger_team_is_more_likely_to_win() {
|
||||||
|
let perf = [g(10.0, 3.0), g(0.0, 3.0), g(-10.0, 3.0)];
|
||||||
|
let p = win_probabilities(&perf, &flat(3, 0.0));
|
||||||
|
assert!(p[0] > p[1] && p[1] > p[2], "not monotone: {p:?}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn identical_teams_are_equally_likely_to_win() {
|
||||||
|
let perf = [g(1.0, 4.0), g(1.0, 4.0), g(1.0, 4.0)];
|
||||||
|
let p = win_probabilities(&perf, &flat(3, 0.0));
|
||||||
|
for probs in p.windows(2) {
|
||||||
|
assert!((probs[0] - probs[1]).abs() < 1e-9, "asymmetric: {p:?}");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Every realisation sorts into exactly one finishing order, so the whole
|
||||||
|
/// distribution must sum to one — with or without a draw margin.
|
||||||
|
#[test]
|
||||||
|
fn outcome_distribution_sums_to_one() {
|
||||||
|
for (perf, eps) in [
|
||||||
|
(vec![g(0.0, 6.0), g(0.0, 6.0)], 0.0),
|
||||||
|
(vec![g(0.0, 6.0), g(0.0, 6.0)], 2.0),
|
||||||
|
(vec![g(0.0, 6.0), g(0.0, 6.0), g(0.0, 6.0)], 0.0),
|
||||||
|
(vec![g(5.0, 6.0), g(0.0, 3.0), g(-5.0, 1.0)], 1.5),
|
||||||
|
(vec![g(0.0, 0.05), g(0.5, 6.0), g(-0.5, 6.0)], 1.0),
|
||||||
|
(
|
||||||
|
vec![g(6.0, 2.0), g(2.0, 6.0), g(-2.0, 1.0), g(-6.0, 4.0)],
|
||||||
|
1.0,
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
let n = perf.len();
|
||||||
|
let dist = outcome_distribution(&perf, &flat(n, eps));
|
||||||
|
let sum: f64 = dist.iter().map(|(_, p)| p).sum();
|
||||||
|
assert!(
|
||||||
|
(sum - 1.0).abs() < 1e-6,
|
||||||
|
"{n} teams, eps={eps}: sum = {sum} over {} outcomes",
|
||||||
|
dist.len()
|
||||||
|
);
|
||||||
|
assert!(dist.iter().all(|(_, p)| *p >= 0.0), "negative probability");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// With two teams the distribution is the exact win/draw/loss triple.
|
||||||
|
#[test]
|
||||||
|
fn two_team_distribution_matches_the_closed_form() {
|
||||||
|
let perf = [g(3.0, 6.0), g(-2.0, 1.0)];
|
||||||
|
let eps = 1.5;
|
||||||
|
let dist = outcome_distribution(&perf, &flat(2, eps));
|
||||||
|
let (wa, wb) = closed_form_two(perf[0], perf[1], eps);
|
||||||
|
|
||||||
|
let find = |ranks: &[u32]| {
|
||||||
|
dist.iter()
|
||||||
|
.find(|(r, _)| r == ranks)
|
||||||
|
.map_or(0.0, |(_, p)| *p)
|
||||||
|
};
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(find(&[0, 1]) - wa).abs() < 1e-6,
|
||||||
|
"a wins: {}",
|
||||||
|
find(&[0, 1])
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(find(&[1, 0]) - wb).abs() < 1e-6,
|
||||||
|
"b wins: {}",
|
||||||
|
find(&[1, 0])
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(find(&[0, 0]) - (1.0 - wa - wb)).abs() < 1e-6,
|
||||||
|
"draw: {}",
|
||||||
|
find(&[0, 0])
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Asking for one ranking must agree with that ranking's entry in the
|
||||||
|
/// full distribution — the two use different code paths to the same value.
|
||||||
|
#[test]
|
||||||
|
fn ranking_probability_agrees_with_the_distribution() {
|
||||||
|
let perf = [g(5.0, 6.0), g(0.0, 3.0), g(-5.0, 1.0)];
|
||||||
|
let eps = 1.5;
|
||||||
|
let margins = flat(3, eps);
|
||||||
|
let dist = outcome_distribution(&perf, &margins);
|
||||||
|
|
||||||
|
for (ranks, expected) in &dist {
|
||||||
|
let direct = ranking_probability(&perf, &margins, ranks);
|
||||||
|
assert!(
|
||||||
|
(direct - expected).abs() < 1e-9,
|
||||||
|
"ranks {ranks:?}: direct {direct} vs distribution {expected}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Tie mass is controlled by the draw margin. Only the *all-tied* outcome
|
||||||
|
/// is monotone in it: every one of its constraints is a window that widens
|
||||||
|
/// with the margin. A partially-tied outcome like `[0, 0, 1]` is not, and
|
||||||
|
/// must not be asserted to be — widening the margin makes its tie easier
|
||||||
|
/// but its "and the last team is strictly behind by more than the margin"
|
||||||
|
/// clause harder, so it peaks and then falls.
|
||||||
|
#[test]
|
||||||
|
fn all_tied_probability_grows_with_the_draw_margin() {
|
||||||
|
let perf = [g(0.0, 4.0), g(0.0, 4.0), g(-8.0, 2.0)];
|
||||||
|
let mut previous = 0.0;
|
||||||
|
for eps in [0.0, 0.5, 1.0, 2.0, 4.0, 8.0, 24.0] {
|
||||||
|
let p = ranking_probability(&perf, &flat(3, eps), &[0, 0, 0]);
|
||||||
|
assert!(p >= previous, "eps={eps}: {p} < {previous}");
|
||||||
|
if eps == 0.0 {
|
||||||
|
assert!(p < 1e-12, "a tie needs a margin, got {p}");
|
||||||
|
}
|
||||||
|
previous = p;
|
||||||
|
}
|
||||||
|
assert!(
|
||||||
|
previous > 0.9,
|
||||||
|
"a very wide margin ties everyone: {previous}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The converse, stated as the non-property it is: a partially-tied
|
||||||
|
/// outcome is non-monotone in the margin. Pinning this down stops a future
|
||||||
|
/// change from "fixing" it into monotonicity and quietly breaking the model.
|
||||||
|
#[test]
|
||||||
|
fn a_partially_tied_outcome_peaks_in_the_middle() {
|
||||||
|
let perf = [g(0.0, 4.0), g(0.0, 4.0), g(-8.0, 2.0)];
|
||||||
|
let sweep: Vec<f64> = [0.5, 2.0, 4.0, 8.0, 16.0]
|
||||||
|
.iter()
|
||||||
|
.map(|&eps| ranking_probability(&perf, &flat(3, eps), &[0, 0, 1]))
|
||||||
|
.collect();
|
||||||
|
let peak = sweep
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.fold(
|
||||||
|
(0, 0.0),
|
||||||
|
|(bi, bv), (i, &v)| if v > bv { (i, v) } else { (bi, bv) },
|
||||||
|
)
|
||||||
|
.0;
|
||||||
|
assert!(
|
||||||
|
peak > 0 && peak < sweep.len() - 1,
|
||||||
|
"expected an interior peak: {sweep:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// With no draw margin a tie has probability exactly zero, and the
|
||||||
|
/// enumeration must not waste work pretending otherwise.
|
||||||
|
#[test]
|
||||||
|
fn ties_are_impossible_without_a_draw_margin() {
|
||||||
|
let perf = [g(0.0, 4.0), g(0.0, 4.0), g(0.0, 4.0)];
|
||||||
|
let dist = outcome_distribution(&perf, &flat(3, 0.0));
|
||||||
|
assert_eq!(dist.len(), 6, "expected only the 6 strict orders: {dist:?}");
|
||||||
|
assert!(dist.iter().all(|(r, _)| {
|
||||||
|
let mut seen = r.clone();
|
||||||
|
seen.sort_unstable();
|
||||||
|
seen.dedup();
|
||||||
|
seen.len() == r.len()
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,322 @@
|
|||||||
|
//! Deterministic numerical integration for the prediction paths.
|
||||||
|
//!
|
||||||
|
//! Prediction asks two questions that have no closed form beyond two teams:
|
||||||
|
//! "who finishes first" and "how likely is this exact finishing order". Both
|
||||||
|
//! reduce to integrals over a single performance variable, so neither needs a
|
||||||
|
//! sampler — and that matters, because a Monte Carlo predictor would make
|
||||||
|
//! `predict_*` non-reproducible and would answer a slightly different question
|
||||||
|
//! on every call.
|
||||||
|
//!
|
||||||
|
//! Two routines live here:
|
||||||
|
//!
|
||||||
|
//! - [`integrate`], adaptive Gauss-Kronrod G7-K15, for the first-place
|
||||||
|
//! marginals. It carries its own error estimate, so it can refine where the
|
||||||
|
//! integrand actually bends instead of guessing a node count up front.
|
||||||
|
//! - [`Grid`], a uniform grid with trapezoid prefix sums, for the ranking
|
||||||
|
//! chain recursion, where each level needs the *running* integral of the
|
||||||
|
//! level below at arbitrary points rather than one definite integral.
|
||||||
|
//!
|
||||||
|
//! Fixed-node Gauss-Hermite is the obvious tool for the first of these and is
|
||||||
|
//! a trap: the integrand is a product of normal CDFs, and when one team's
|
||||||
|
//! sigma is much smaller than the integrating team's, that product turns into
|
||||||
|
//! a near-step function narrower than the node spacing. The nodes step over
|
||||||
|
//! it and the result is wrong by ~1e-2 while still looking like a probability.
|
||||||
|
//! Adaptive refinement is what makes the small-sigma case safe.
|
||||||
|
|
||||||
|
/// Kronrod 15-point abscissae, non-negative half, descending.
|
||||||
|
const XGK: [f64; 8] = [
|
||||||
|
0.991_455_371_120_813,
|
||||||
|
0.949_107_912_342_759,
|
||||||
|
0.864_864_423_359_769,
|
||||||
|
0.741_531_185_599_394,
|
||||||
|
0.586_087_235_467_691,
|
||||||
|
0.405_845_151_377_397,
|
||||||
|
0.207_784_955_007_898,
|
||||||
|
0.0,
|
||||||
|
];
|
||||||
|
|
||||||
|
/// Kronrod 15-point weights, matching [`XGK`].
|
||||||
|
const WGK: [f64; 8] = [
|
||||||
|
0.022_935_322_010_529,
|
||||||
|
0.063_092_092_629_979,
|
||||||
|
0.104_790_010_322_250,
|
||||||
|
0.140_653_259_715_525,
|
||||||
|
0.169_004_726_639_267,
|
||||||
|
0.190_350_578_064_785,
|
||||||
|
0.204_432_940_075_298,
|
||||||
|
0.209_482_141_084_728,
|
||||||
|
];
|
||||||
|
|
||||||
|
/// Gauss 7-point weights, applying to the odd-indexed [`XGK`] entries.
|
||||||
|
const WG: [f64; 4] = [
|
||||||
|
0.129_484_966_168_870,
|
||||||
|
0.279_705_391_489_277,
|
||||||
|
0.381_830_050_505_119,
|
||||||
|
0.417_959_183_673_469,
|
||||||
|
];
|
||||||
|
|
||||||
|
/// Panels are bisected worst-first; this bounds the work on a pathological
|
||||||
|
/// integrand rather than letting it spin.
|
||||||
|
const MAX_SUBDIVISIONS: usize = 200;
|
||||||
|
|
||||||
|
/// One G7-K15 panel over `[a, b]`: `(integral, absolute error estimate)`.
|
||||||
|
///
|
||||||
|
/// The error estimate is the gap between the embedded 7-point Gauss rule and
|
||||||
|
/// the 15-point Kronrod extension. It is the only reason this is preferable
|
||||||
|
/// to a fixed rule: it tells the caller *where* the integrand is hard.
|
||||||
|
fn gk15<F: Fn(f64) -> f64>(f: &F, a: f64, b: f64) -> (f64, f64) {
|
||||||
|
let centre = 0.5 * (a + b);
|
||||||
|
let half = 0.5 * (b - a);
|
||||||
|
|
||||||
|
let mut kronrod = 0.0;
|
||||||
|
let mut gauss = 0.0;
|
||||||
|
|
||||||
|
for i in 0..8 {
|
||||||
|
let offset = XGK[i] * half;
|
||||||
|
// XGK[7] is the centre node and must not be counted twice.
|
||||||
|
let sum = if i == 7 {
|
||||||
|
f(centre)
|
||||||
|
} else {
|
||||||
|
f(centre - offset) + f(centre + offset)
|
||||||
|
};
|
||||||
|
kronrod += WGK[i] * sum;
|
||||||
|
if i % 2 == 1 {
|
||||||
|
gauss += WG[i / 2] * sum;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
(kronrod * half, ((kronrod - gauss) * half).abs())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Adaptively integrate `f` over `[a, b]` to relative tolerance `tol`.
|
||||||
|
///
|
||||||
|
/// `seeds` are interior points where the integrand is known to bend sharply —
|
||||||
|
/// for a product of normal CDFs, each rival's transition centre. Splitting
|
||||||
|
/// there up front costs nothing and saves the adaptive loop from having to
|
||||||
|
/// discover a step by bisection.
|
||||||
|
///
|
||||||
|
/// Returns the integral. The error estimate is consumed internally rather
|
||||||
|
/// than returned: callers here integrate probability densities, where the
|
||||||
|
/// meaningful check is the sum-to-one identity over a whole outcome space,
|
||||||
|
/// not a per-integral residual.
|
||||||
|
pub(crate) fn integrate<F: Fn(f64) -> f64>(f: F, a: f64, b: f64, seeds: &[f64], tol: f64) -> f64 {
|
||||||
|
// Explicit rather than `!(b > a)`: a NaN bound must fall through to zero
|
||||||
|
// rather than being read as a valid ordering.
|
||||||
|
if a.partial_cmp(&b) != Some(std::cmp::Ordering::Less) {
|
||||||
|
return 0.0;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut edges: Vec<f64> = Vec::with_capacity(seeds.len() + 2);
|
||||||
|
edges.push(a);
|
||||||
|
edges.push(b);
|
||||||
|
for &s in seeds {
|
||||||
|
if s > a && s < b {
|
||||||
|
edges.push(s);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
edges.sort_by(|p, q| p.partial_cmp(q).expect("integration bounds are finite"));
|
||||||
|
edges.dedup();
|
||||||
|
|
||||||
|
// (lo, hi, integral, error)
|
||||||
|
let mut panels: Vec<(f64, f64, f64, f64)> = edges
|
||||||
|
.windows(2)
|
||||||
|
.map(|w| {
|
||||||
|
let (v, e) = gk15(&f, w[0], w[1]);
|
||||||
|
(w[0], w[1], v, e)
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
for _ in 0..MAX_SUBDIVISIONS {
|
||||||
|
let total: f64 = panels.iter().map(|p| p.2).sum();
|
||||||
|
let error: f64 = panels.iter().map(|p| p.3).sum();
|
||||||
|
|
||||||
|
// Absolute floor as well as relative: these integrands are
|
||||||
|
// probabilities, so an absolute 1e-15 is already past the useful
|
||||||
|
// precision of the underlying `cdf`.
|
||||||
|
if error <= tol * total.abs().max(1e-12) || error < 1e-15 {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
let worst = panels
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.fold((0usize, f64::NEG_INFINITY), |(bi, be), (i, p)| {
|
||||||
|
if p.3 > be { (i, p.3) } else { (bi, be) }
|
||||||
|
})
|
||||||
|
.0;
|
||||||
|
|
||||||
|
let (lo, hi, _, _) = panels[worst];
|
||||||
|
let mid = 0.5 * (lo + hi);
|
||||||
|
// Bisection has hit the floating-point floor; refining further would
|
||||||
|
// loop without reducing the error.
|
||||||
|
if !(mid > lo && mid < hi) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
let (v1, e1) = gk15(&f, lo, mid);
|
||||||
|
let (v2, e2) = gk15(&f, mid, hi);
|
||||||
|
panels[worst] = (lo, mid, v1, e1);
|
||||||
|
panels.push((mid, hi, v2, e2));
|
||||||
|
}
|
||||||
|
|
||||||
|
panels.iter().map(|p| p.2).sum()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A uniform grid carrying trapezoid prefix sums of one integrand.
|
||||||
|
///
|
||||||
|
/// The ranking recursion needs, at every level, the running integral of the
|
||||||
|
/// level below evaluated at arbitrary points — a cumulative integral, not a
|
||||||
|
/// definite one. Prefix sums give that in O(1) per query after an O(G) build,
|
||||||
|
/// which is what keeps a full ranking probability linear in the team count.
|
||||||
|
pub(crate) struct Grid {
|
||||||
|
lo: f64,
|
||||||
|
step: f64,
|
||||||
|
/// Integrand sampled at each node.
|
||||||
|
values: Vec<f64>,
|
||||||
|
/// `prefix[i]` is the integral from `lo` to node `i`.
|
||||||
|
prefix: Vec<f64>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Grid {
|
||||||
|
/// Build directly from already-sampled values.
|
||||||
|
///
|
||||||
|
/// The ranking recursion evaluates every level on the same nodes, so the
|
||||||
|
/// per-team densities are sampled once and reused; re-evaluating `exp`
|
||||||
|
/// per level would dominate the cost.
|
||||||
|
pub(crate) fn from_values(lo: f64, step: f64, values: Vec<f64>) -> Self {
|
||||||
|
let mut prefix = vec![0.0; values.len()];
|
||||||
|
for i in 1..values.len() {
|
||||||
|
prefix[i] = prefix[i - 1] + 0.5 * step * (values[i - 1] + values[i]);
|
||||||
|
}
|
||||||
|
Self {
|
||||||
|
lo,
|
||||||
|
step,
|
||||||
|
values,
|
||||||
|
prefix,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Integral from the grid's lower bound up to `x`.
|
||||||
|
///
|
||||||
|
/// Clamped at both ends: the caller sizes the grid to cover the whole
|
||||||
|
/// support, so a query outside it is asking for a tail that is zero (below)
|
||||||
|
/// or the whole mass (above).
|
||||||
|
pub(crate) fn integral_to(&self, x: f64) -> f64 {
|
||||||
|
let last = self.values.len() - 1;
|
||||||
|
if x <= self.lo {
|
||||||
|
return 0.0;
|
||||||
|
}
|
||||||
|
if x >= self.lo + last as f64 * self.step {
|
||||||
|
return self.prefix[last];
|
||||||
|
}
|
||||||
|
|
||||||
|
let scaled = (x - self.lo) / self.step;
|
||||||
|
let i = scaled.floor() as usize;
|
||||||
|
let frac = scaled - i as f64;
|
||||||
|
|
||||||
|
// Whole cells, plus the trapezoid over the partial cell. The integrand
|
||||||
|
// is linear within a cell under the trapezoid rule, so the partial
|
||||||
|
// piece is exact with respect to that same approximation.
|
||||||
|
self.prefix[i]
|
||||||
|
+ frac
|
||||||
|
* self.step
|
||||||
|
* (self.values[i] + 0.5 * frac * (self.values[i + 1] - self.values[i]))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Integral over `[from, to]`.
|
||||||
|
pub(crate) fn integral_between(&self, from: f64, to: f64) -> f64 {
|
||||||
|
(self.integral_to(to) - self.integral_to(from)).max(0.0)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Total integral over the whole grid.
|
||||||
|
pub(crate) fn total(&self) -> f64 {
|
||||||
|
self.prefix[self.values.len() - 1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
const TOL: f64 = 1e-10;
|
||||||
|
|
||||||
|
/// Sample `f` over `[lo, hi]` at `points` nodes.
|
||||||
|
fn sample<F: FnMut(f64) -> f64>(lo: f64, hi: f64, points: usize, mut f: F) -> Grid {
|
||||||
|
let step = (hi - lo) / (points - 1) as f64;
|
||||||
|
Grid::from_values(
|
||||||
|
lo,
|
||||||
|
step,
|
||||||
|
(0..points).map(|i| f(lo + i as f64 * step)).collect(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn integrates_a_polynomial_exactly() {
|
||||||
|
// G7-K15 is exact for polynomials well past cubic, so a single panel
|
||||||
|
// should already be at round-off.
|
||||||
|
let v = integrate(|x| 3.0 * x * x + 2.0 * x + 1.0, 0.0, 2.0, &[], TOL);
|
||||||
|
assert!((v - 14.0).abs() < 1e-12, "got {v}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn integrates_a_gaussian_density_to_one() {
|
||||||
|
let f = |x: f64| (-0.5 * x * x).exp() / (2.0 * std::f64::consts::PI).sqrt();
|
||||||
|
let v = integrate(f, -10.0, 10.0, &[], TOL);
|
||||||
|
assert!((v - 1.0).abs() < 1e-12, "got {v}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn resolves_a_step_far_narrower_than_the_initial_panel() {
|
||||||
|
// The failure mode that rules out fixed-node quadrature: a transition
|
||||||
|
// 1e-4 wide inside a range of 20. A fixed rule steps over it.
|
||||||
|
let f = |x: f64| if x < 0.5 { 0.0 } else { 1.0 };
|
||||||
|
let v = integrate(f, -10.0, 10.0, &[0.5], TOL);
|
||||||
|
assert!((v - 9.5).abs() < 1e-6, "got {v}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn seeds_do_not_change_the_value_of_a_smooth_integrand() {
|
||||||
|
let f = |x: f64| (-0.5 * x * x).exp();
|
||||||
|
let plain = integrate(f, -8.0, 8.0, &[], TOL);
|
||||||
|
let seeded = integrate(f, -8.0, 8.0, &[-3.0, 0.25, 5.5], TOL);
|
||||||
|
assert!((plain - seeded).abs() < 1e-12, "{plain} vs {seeded}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_or_inverted_range_integrates_to_zero() {
|
||||||
|
assert_eq!(integrate(|_| 1.0, 1.0, 1.0, &[], TOL), 0.0);
|
||||||
|
assert_eq!(integrate(|_| 1.0, 2.0, 1.0, &[], TOL), 0.0);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn grid_prefix_matches_a_known_cumulative_integral() {
|
||||||
|
// f(x) = x over [0, 4]; integral to x is x^2/2.
|
||||||
|
let g = sample(0.0, 4.0, 4001, |x| x);
|
||||||
|
for probe in [0.0, 0.5, 1.0, 2.5, 3.75, 4.0] {
|
||||||
|
let want = probe * probe / 2.0;
|
||||||
|
let got = g.integral_to(probe);
|
||||||
|
assert!(
|
||||||
|
(got - want).abs() < 1e-9,
|
||||||
|
"at {probe}: got {got}, want {want}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
assert!((g.total() - 8.0).abs() < 1e-9);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn grid_between_is_the_difference_of_two_prefixes() {
|
||||||
|
let g = sample(-5.0, 5.0, 8001, |x| (-0.5 * x * x).exp());
|
||||||
|
let whole = g.integral_between(-5.0, 5.0);
|
||||||
|
let split = g.integral_between(-5.0, 0.3) + g.integral_between(0.3, 5.0);
|
||||||
|
assert!((whole - split).abs() < 1e-12, "{whole} vs {split}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn grid_clamps_queries_outside_its_support() {
|
||||||
|
let g = sample(0.0, 1.0, 101, |_| 1.0);
|
||||||
|
assert_eq!(g.integral_to(-3.0), 0.0);
|
||||||
|
assert!((g.integral_to(9.0) - 1.0).abs() < 1e-12);
|
||||||
|
// Reversed bounds must not produce negative probability mass.
|
||||||
|
assert_eq!(g.integral_between(0.8, 0.2), 0.0);
|
||||||
|
}
|
||||||
|
}
|
||||||
+57
-2
@@ -9,13 +9,16 @@ use crate::{
|
|||||||
|
|
||||||
/// Static rating configuration: prior skill, performance noise `beta`, drift.
|
/// Static rating configuration: prior skill, performance noise `beta`, drift.
|
||||||
///
|
///
|
||||||
/// Renamed from `Player` in T2; `Rating` better describes the data
|
/// A configuration rather than a person: the per-history temporal state
|
||||||
/// (a configuration) vs. a person (who's a `Competitor` with state).
|
/// (messages, last appearance) lives on `Competitor`.
|
||||||
#[derive(Clone, Copy, Debug)]
|
#[derive(Clone, Copy, Debug)]
|
||||||
pub struct Rating<T: Time = i64, D: Drift<T> = ConstantDrift> {
|
pub struct Rating<T: Time = i64, D: Drift<T> = ConstantDrift> {
|
||||||
pub(crate) prior: Gaussian,
|
pub(crate) prior: Gaussian,
|
||||||
pub(crate) beta: f64,
|
pub(crate) beta: f64,
|
||||||
pub(crate) drift: D,
|
pub(crate) drift: D,
|
||||||
|
/// Multiplier on the drift *variance* this competitor accumulates; 1.0 is
|
||||||
|
/// the neutral default. Set per competitor via `Member::with_drift_scale`.
|
||||||
|
pub(crate) drift_scale: f64,
|
||||||
pub(crate) _time: PhantomData<T>,
|
pub(crate) _time: PhantomData<T>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -25,10 +28,61 @@ impl<T: Time, D: Drift<T>> Rating<T, D> {
|
|||||||
prior,
|
prior,
|
||||||
beta,
|
beta,
|
||||||
drift,
|
drift,
|
||||||
|
drift_scale: 1.0,
|
||||||
_time: PhantomData,
|
_time: PhantomData,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Scale how fast this competitor drifts, relative to `drift`.
|
||||||
|
///
|
||||||
|
/// Multiplies the drift *variance*, so the scale is in the same units as
|
||||||
|
/// `gamma`. `0.0` pins the competitor still.
|
||||||
|
#[must_use]
|
||||||
|
pub fn with_drift_scale(mut self, drift_scale: f64) -> Self {
|
||||||
|
self.drift_scale = drift_scale;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The configured prior skill estimate.
|
||||||
|
#[must_use]
|
||||||
|
pub fn prior(&self) -> Gaussian {
|
||||||
|
self.prior
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Performance noise: how much a single showing varies around the skill.
|
||||||
|
#[must_use]
|
||||||
|
pub fn beta(&self) -> f64 {
|
||||||
|
self.beta
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The drift model governing how skill may move between events.
|
||||||
|
#[must_use]
|
||||||
|
pub fn drift(&self) -> D {
|
||||||
|
self.drift
|
||||||
|
}
|
||||||
|
|
||||||
|
/// This competitor's multiplier on the drift variance; 1.0 is neutral.
|
||||||
|
#[must_use]
|
||||||
|
pub fn drift_scale(&self) -> f64 {
|
||||||
|
self.drift_scale
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Drift variance accumulated over `from -> to`, scaled for this competitor.
|
||||||
|
///
|
||||||
|
/// The single place the scale is applied for a `Time`-typed span. Callers
|
||||||
|
/// must go through this rather than `self.drift` directly, so a competitor's
|
||||||
|
/// scale cannot be silently skipped.
|
||||||
|
pub(crate) fn drift_variance_delta(&self, from: &T, to: &T) -> f64 {
|
||||||
|
self.drift.variance_delta(from, to) * self.drift_scale * self.drift_scale
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Drift variance for a cached elapsed count, scaled for this competitor.
|
||||||
|
///
|
||||||
|
/// The counterpart of `drift_variance_delta` for the cached-elapsed paths.
|
||||||
|
pub(crate) fn drift_variance_for_elapsed(&self, elapsed: i64) -> f64 {
|
||||||
|
self.drift.variance_for_elapsed(elapsed) * self.drift_scale * self.drift_scale
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn performance(&self) -> Gaussian {
|
pub(crate) fn performance(&self) -> Gaussian {
|
||||||
self.prior.forget(self.beta.powi(2))
|
self.prior.forget(self.beta.powi(2))
|
||||||
}
|
}
|
||||||
@@ -40,6 +94,7 @@ impl Default for Rating<i64, ConstantDrift> {
|
|||||||
prior: Gaussian::default(),
|
prior: Gaussian::default(),
|
||||||
beta: BETA,
|
beta: BETA,
|
||||||
drift: ConstantDrift(GAMMA),
|
drift: ConstantDrift(GAMMA),
|
||||||
|
drift_scale: 1.0,
|
||||||
_time: PhantomData,
|
_time: PhantomData,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+33
-7
@@ -1,7 +1,7 @@
|
|||||||
//! Schedule trait and built-in implementations.
|
//! Schedule trait and built-in implementations.
|
||||||
//!
|
//!
|
||||||
//! A schedule drives factor propagation to convergence. The default
|
//! A schedule drives factor propagation to convergence. The default
|
||||||
//! `EpsilonOrMax` performs one TeamSum sweep (setup) then alternating
|
//! `EpsilonOrMax` performs one `TeamSum` sweep (setup) then alternating
|
||||||
//! forward/backward sweeps over the iterating factors until the max
|
//! forward/backward sweeps over the iterating factors until the max
|
||||||
//! delta drops below epsilon or `max` iterations is reached.
|
//! delta drops below epsilon or `max` iterations is reached.
|
||||||
|
|
||||||
@@ -23,7 +23,7 @@ pub trait Schedule: Send + Sync {
|
|||||||
/// Default schedule: sweep forward then backward until step ≤ eps or iter == max.
|
/// Default schedule: sweep forward then backward until step ≤ eps or iter == max.
|
||||||
///
|
///
|
||||||
/// Matches the existing `Game::likelihoods` loop bit-for-bit when given the
|
/// Matches the existing `Game::likelihoods` loop bit-for-bit when given the
|
||||||
/// same factor layout (TeamSums first, then alternating RankDiff/Trunc pairs).
|
/// same factor layout (`TeamSums` first, then alternating RankDiff/Trunc pairs).
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
pub struct EpsilonOrMax {
|
pub struct EpsilonOrMax {
|
||||||
pub eps: f64,
|
pub eps: f64,
|
||||||
@@ -32,8 +32,17 @@ pub struct EpsilonOrMax {
|
|||||||
|
|
||||||
impl Default for EpsilonOrMax {
|
impl Default for EpsilonOrMax {
|
||||||
fn default() -> Self {
|
fn default() -> Self {
|
||||||
// Matches today's hard-coded tolerance and iteration cap.
|
// Derived from `ConvergenceOptions` so there is one source of truth for
|
||||||
Self { eps: 1e-6, max: 10 }
|
// the tolerance and iteration cap. These previously disagreed: this
|
||||||
|
// default capped at 10 iterations while `ConvergenceOptions` allowed 30,
|
||||||
|
// and which applied depended on whether inference went through
|
||||||
|
// `run_chain` or a `Schedule`.
|
||||||
|
let defaults = crate::ConvergenceOptions::default();
|
||||||
|
|
||||||
|
Self {
|
||||||
|
eps: defaults.epsilon,
|
||||||
|
max: defaults.max_iter,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -50,10 +59,16 @@ impl Schedule for EpsilonOrMax {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let mut iterations = 0;
|
let mut iterations = 0;
|
||||||
let mut final_step = (f64::INFINITY, f64::INFINITY);
|
// With no iterating factors the graph is already at its fixed point:
|
||||||
let mut converged = false;
|
// the setup pass above is all there is to do. Reporting `converged:
|
||||||
|
// false` with an infinite step for that case gave callers a false
|
||||||
|
// negative.
|
||||||
|
let mut final_step = (0.0, 0.0);
|
||||||
|
let mut converged = true;
|
||||||
|
|
||||||
if n_setup < factors.len() {
|
if n_setup < factors.len() {
|
||||||
|
final_step = (f64::INFINITY, f64::INFINITY);
|
||||||
|
converged = false;
|
||||||
for _ in 0..self.max {
|
for _ in 0..self.max {
|
||||||
let mut step = (0.0_f64, 0.0_f64);
|
let mut step = (0.0_f64, 0.0_f64);
|
||||||
|
|
||||||
@@ -113,7 +128,8 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn report_marks_converged_when_no_iterating_factors() {
|
fn report_marks_converged_when_no_iterating_factors() {
|
||||||
// No iterating factors → 0 iterations, converged stays false (loop never ran).
|
// A graph of only setup factors has nothing to iterate, so it is at its
|
||||||
|
// fixed point after the setup pass: 0 iterations, and converged.
|
||||||
let mut vars = VarStore::new();
|
let mut vars = VarStore::new();
|
||||||
let out = vars.alloc(N_INF);
|
let out = vars.alloc(N_INF);
|
||||||
let mut factors = vec![BuiltinFactor::TeamSum(TeamSumFactor {
|
let mut factors = vec![BuiltinFactor::TeamSum(TeamSumFactor {
|
||||||
@@ -122,5 +138,15 @@ mod tests {
|
|||||||
})];
|
})];
|
||||||
let report = EpsilonOrMax::default().run(&mut factors, &mut vars);
|
let report = EpsilonOrMax::default().run(&mut factors, &mut vars);
|
||||||
assert_eq!(report.iterations, 0);
|
assert_eq!(report.iterations, 0);
|
||||||
|
assert!(report.converged);
|
||||||
|
assert_eq!(report.final_step, (0.0, 0.0));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn default_matches_convergence_options() {
|
||||||
|
let schedule = EpsilonOrMax::default();
|
||||||
|
let options = crate::ConvergenceOptions::default();
|
||||||
|
assert_eq!(schedule.max, options.max_iter);
|
||||||
|
assert_eq!(schedule.eps, options.epsilon);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ use crate::{Index, competitor::Competitor, drift::Drift, time::Time};
|
|||||||
|
|
||||||
/// Dense Vec-backed store for competitor state in History.
|
/// Dense Vec-backed store for competitor state in History.
|
||||||
///
|
///
|
||||||
/// Indexed directly by Index.0, eliminating HashMap hashing in the
|
/// Indexed directly by Index.0, eliminating `HashMap` hashing in the
|
||||||
/// forward/backward sweep. Uses `Vec<Option<Competitor<T, D>>>` so slots can be
|
/// forward/backward sweep. Uses `Vec<Option<Competitor<T, D>>>` so slots can be
|
||||||
/// absent without an explicit present mask.
|
/// absent without an explicit present mask.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
@@ -21,6 +21,7 @@ impl<T: Time, D: Drift<T>> Default for CompetitorStore<T, D> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl<T: Time, D: Drift<T>> CompetitorStore<T, D> {
|
impl<T: Time, D: Drift<T>> CompetitorStore<T, D> {
|
||||||
|
#[must_use]
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self::default()
|
Self::default()
|
||||||
}
|
}
|
||||||
@@ -39,6 +40,7 @@ impl<T: Time, D: Drift<T>> CompetitorStore<T, D> {
|
|||||||
self.competitors[idx.0] = Some(competitor);
|
self.competitors[idx.0] = Some(competitor);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn get(&self, idx: Index) -> Option<&Competitor<T, D>> {
|
pub fn get(&self, idx: Index) -> Option<&Competitor<T, D>> {
|
||||||
self.competitors.get(idx.0).and_then(|slot| slot.as_ref())
|
self.competitors.get(idx.0).and_then(|slot| slot.as_ref())
|
||||||
}
|
}
|
||||||
@@ -49,14 +51,17 @@ impl<T: Time, D: Drift<T>> CompetitorStore<T, D> {
|
|||||||
.and_then(|slot| slot.as_mut())
|
.and_then(|slot| slot.as_mut())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn contains(&self, idx: Index) -> bool {
|
pub fn contains(&self, idx: Index) -> bool {
|
||||||
self.get(idx).is_some()
|
self.get(idx).is_some()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.n_present
|
self.n_present
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[must_use]
|
||||||
pub fn is_empty(&self) -> bool {
|
pub fn is_empty(&self) -> bool {
|
||||||
self.n_present == 0
|
self.n_present == 0
|
||||||
}
|
}
|
||||||
|
|||||||
+110
-51
@@ -1,15 +1,27 @@
|
|||||||
|
use std::collections::HashMap;
|
||||||
|
|
||||||
use crate::{Index, time_slice::Skill};
|
use crate::{Index, time_slice::Skill};
|
||||||
|
|
||||||
/// Dense Vec-backed store for per-agent skill state within a TimeSlice.
|
/// Compact per-slice store for skill state, addressed by a slice-local slot.
|
||||||
///
|
///
|
||||||
/// Indexed directly by Index.0, eliminating HashMap hashing in the inner
|
/// `skills` holds one entry per competitor **in this slice**, so memory is
|
||||||
/// convergence loop. Uses a parallel `present` mask so iteration skips
|
/// O(competitors in the slice). It used to be a dense `Vec<Skill>` indexed by
|
||||||
/// absent slots without incurring per-slot Option overhead in the hot path.
|
/// the global `Index.0`, which made a slice's footprint O(largest index it
|
||||||
|
/// touches): a single 1v1 game between competitors 19998 and 19999 reserved
|
||||||
|
/// 20,000 slots.
|
||||||
|
///
|
||||||
|
/// The dense layout existed to keep `HashMap` hashing out of the inner
|
||||||
|
/// convergence loop, and that property is preserved. `slots` is consulted only
|
||||||
|
/// while building a slice; every hot-path access goes through
|
||||||
|
/// [`SkillStore::at`] / [`SkillStore::at_mut`] with a slot resolved once at
|
||||||
|
/// ingestion and cached on the event's `Item`.
|
||||||
#[derive(Debug, Default)]
|
#[derive(Debug, Default)]
|
||||||
pub struct SkillStore {
|
pub struct SkillStore {
|
||||||
skills: Vec<Skill>,
|
skills: Vec<Skill>,
|
||||||
present: Vec<bool>,
|
/// Slot -> global index, parallel to `skills`, so iteration can report the
|
||||||
n_present: usize,
|
/// global index without a reverse lookup.
|
||||||
|
indices: Vec<Index>,
|
||||||
|
slots: HashMap<Index, u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl SkillStore {
|
impl SkillStore {
|
||||||
@@ -17,76 +29,99 @@ impl SkillStore {
|
|||||||
Self::default()
|
Self::default()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn ensure_capacity(&mut self, idx: usize) {
|
/// Resolve a global index to this slice's slot, if the competitor is here.
|
||||||
if idx >= self.skills.len() {
|
///
|
||||||
self.skills.resize_with(idx + 1, Skill::default);
|
/// This hashes. Call it at ingestion and cache the result; do not call it
|
||||||
self.present.resize(idx + 1, false);
|
/// from the convergence loop.
|
||||||
}
|
pub fn slot_of(&self, idx: Index) -> Option<u32> {
|
||||||
|
self.slots.get(&idx).copied()
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn insert(&mut self, idx: Index, skill: Skill) {
|
/// Skill at a slot resolved earlier by [`SkillStore::slot_of`].
|
||||||
self.ensure_capacity(idx.0);
|
///
|
||||||
if !self.present[idx.0] {
|
/// # Panics
|
||||||
self.n_present += 1;
|
///
|
||||||
|
/// Panics if `slot` is out of range, which means it came from a different
|
||||||
|
/// slice's store.
|
||||||
|
pub fn at(&self, slot: u32) -> &Skill {
|
||||||
|
&self.skills[slot as usize]
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Mutable counterpart to [`SkillStore::at`].
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if `slot` is out of range.
|
||||||
|
pub fn at_mut(&mut self, slot: u32) -> &mut Skill {
|
||||||
|
&mut self.skills[slot as usize]
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Insert or overwrite a competitor's skill, returning its slot.
|
||||||
|
pub fn insert(&mut self, idx: Index, skill: Skill) -> u32 {
|
||||||
|
match self.slots.get(&idx) {
|
||||||
|
Some(&slot) => {
|
||||||
|
self.skills[slot as usize] = skill;
|
||||||
|
slot
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
let slot = u32::try_from(self.skills.len())
|
||||||
|
.expect("a time slice cannot hold more than u32::MAX competitors");
|
||||||
|
|
||||||
|
self.skills.push(skill);
|
||||||
|
self.indices.push(idx);
|
||||||
|
self.slots.insert(idx, slot);
|
||||||
|
|
||||||
|
slot
|
||||||
|
}
|
||||||
}
|
}
|
||||||
self.skills[idx.0] = skill;
|
|
||||||
self.present[idx.0] = true;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get(&self, idx: Index) -> Option<&Skill> {
|
pub fn get(&self, idx: Index) -> Option<&Skill> {
|
||||||
if idx.0 < self.present.len() && self.present[idx.0] {
|
self.slot_of(idx).map(|slot| self.at(slot))
|
||||||
Some(&self.skills[idx.0])
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_mut(&mut self, idx: Index) -> Option<&mut Skill> {
|
pub fn get_mut(&mut self, idx: Index) -> Option<&mut Skill> {
|
||||||
if idx.0 < self.present.len() && self.present[idx.0] {
|
self.slot_of(idx)
|
||||||
Some(&mut self.skills[idx.0])
|
.map(|slot| &mut self.skills[slot as usize])
|
||||||
} else {
|
|
||||||
None
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)]
|
/// Whether a competitor is present in this slice. Test-only.
|
||||||
|
#[cfg(test)]
|
||||||
pub fn contains(&self, idx: Index) -> bool {
|
pub fn contains(&self, idx: Index) -> bool {
|
||||||
idx.0 < self.present.len() && self.present[idx.0]
|
self.slots.contains_key(&idx)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)]
|
/// Number of competitors in this slice. Test-only.
|
||||||
|
#[cfg(test)]
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.n_present
|
self.skills.len()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)]
|
/// Slots actually allocated — the quantity #17 is about, and NOT the same
|
||||||
pub fn is_empty(&self) -> bool {
|
/// as `len` for every possible implementation.
|
||||||
self.n_present == 0
|
///
|
||||||
|
/// A store indexed by the global `Index` must report `max_index + 1` here
|
||||||
|
/// while reporting the true competitor count from `len`, which is exactly
|
||||||
|
/// how the original defect hid. Tests that mean to pin the footprint must
|
||||||
|
/// assert on this.
|
||||||
|
#[cfg(test)]
|
||||||
|
pub fn allocated_slots(&self) -> usize {
|
||||||
|
self.skills.len()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Iterate in slot order — the order competitors were first seen in this
|
||||||
|
/// slice. Deterministic for a given event order, which is what the
|
||||||
|
/// cross-thread determinism test relies on.
|
||||||
pub fn iter(&self) -> impl Iterator<Item = (Index, &Skill)> {
|
pub fn iter(&self) -> impl Iterator<Item = (Index, &Skill)> {
|
||||||
self.present.iter().enumerate().filter_map(|(i, &p)| {
|
self.indices.iter().copied().zip(self.skills.iter())
|
||||||
if p {
|
|
||||||
Some((Index(i), &self.skills[i]))
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn iter_mut(&mut self) -> impl Iterator<Item = (Index, &mut Skill)> {
|
pub fn iter_mut(&mut self) -> impl Iterator<Item = (Index, &mut Skill)> {
|
||||||
self.skills
|
self.indices.iter().copied().zip(self.skills.iter_mut())
|
||||||
.iter_mut()
|
|
||||||
.zip(self.present.iter())
|
|
||||||
.enumerate()
|
|
||||||
.filter_map(|(i, (s, &p))| if p { Some((Index(i), s)) } else { None })
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn keys(&self) -> impl Iterator<Item = Index> + '_ {
|
pub fn keys(&self) -> impl Iterator<Item = Index> + '_ {
|
||||||
self.present
|
self.indices.iter().copied()
|
||||||
.iter()
|
|
||||||
.enumerate()
|
|
||||||
.filter_map(|(i, &p)| if p { Some(Index(i)) } else { None })
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -112,7 +147,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn iter_skips_absent_slots() {
|
fn iter_reports_global_indices() {
|
||||||
let mut store = SkillStore::new();
|
let mut store = SkillStore::new();
|
||||||
store.insert(Index(0), Skill::default());
|
store.insert(Index(0), Skill::default());
|
||||||
store.insert(Index(5), Skill::default());
|
store.insert(Index(5), Skill::default());
|
||||||
@@ -127,4 +162,28 @@ mod tests {
|
|||||||
store.insert(Index(2), Skill::default());
|
store.insert(Index(2), Skill::default());
|
||||||
assert_eq!(store.len(), 1);
|
assert_eq!(store.len(), 1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// The defect in #17: a slice holding two competitors must cost the same
|
||||||
|
/// whether their indices are small or large.
|
||||||
|
#[test]
|
||||||
|
fn footprint_is_independent_of_index_magnitude() {
|
||||||
|
let mut low = SkillStore::new();
|
||||||
|
low.insert(Index(0), Skill::default());
|
||||||
|
low.insert(Index(1), Skill::default());
|
||||||
|
|
||||||
|
let mut high = SkillStore::new();
|
||||||
|
high.insert(Index(19_998), Skill::default());
|
||||||
|
high.insert(Index(19_999), Skill::default());
|
||||||
|
|
||||||
|
assert_eq!(low.len(), high.len());
|
||||||
|
assert_eq!(low.skills.capacity(), high.skills.capacity());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slot_survives_reinsert() {
|
||||||
|
let mut store = SkillStore::new();
|
||||||
|
let first = store.insert(Index(7), Skill::default());
|
||||||
|
let again = store.insert(Index(7), Skill::default());
|
||||||
|
assert_eq!(first, again);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+303
-109
@@ -14,7 +14,6 @@ use crate::{
|
|||||||
rating::Rating,
|
rating::Rating,
|
||||||
storage::{CompetitorStore, SkillStore},
|
storage::{CompetitorStore, SkillStore},
|
||||||
time::Time,
|
time::Time,
|
||||||
tuple_gt, tuple_max,
|
|
||||||
};
|
};
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
@@ -23,7 +22,6 @@ pub(crate) struct Skill {
|
|||||||
backward: Gaussian,
|
backward: Gaussian,
|
||||||
likelihood: Gaussian,
|
likelihood: Gaussian,
|
||||||
pub(crate) elapsed: i64,
|
pub(crate) elapsed: i64,
|
||||||
pub(crate) online: Gaussian,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Skill {
|
impl Skill {
|
||||||
@@ -39,7 +37,6 @@ impl Default for Skill {
|
|||||||
backward: N_INF,
|
backward: N_INF,
|
||||||
likelihood: N_INF,
|
likelihood: N_INF,
|
||||||
elapsed: 0,
|
elapsed: 0,
|
||||||
online: N_INF,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -51,43 +48,48 @@ pub enum EventKind {
|
|||||||
Scored { score_sigma: f64 },
|
Scored { score_sigma: f64 },
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Clone, Debug)]
|
||||||
struct Item {
|
struct Item {
|
||||||
agent: Index,
|
agent: Index,
|
||||||
|
/// This competitor's slot in the owning slice's `SkillStore`, resolved
|
||||||
|
/// once at ingestion.
|
||||||
|
///
|
||||||
|
/// The convergence loop reaches skills through this rather than through
|
||||||
|
/// `agent`, which is what keeps `HashMap` hashing out of the hot path now
|
||||||
|
/// that the store is compact rather than indexed by the global `Index`.
|
||||||
|
slot: u32,
|
||||||
likelihood: Gaussian,
|
likelihood: Gaussian,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Item {
|
impl Item {
|
||||||
fn within_prior<T: Time, D: Drift<T>>(
|
fn within_prior<T: Time, D: Drift<T>>(
|
||||||
&self,
|
&self,
|
||||||
online: bool,
|
|
||||||
forward: bool,
|
forward: bool,
|
||||||
skills: &SkillStore,
|
skills: &SkillStore,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
) -> Rating<T, D> {
|
) -> Rating<T, D> {
|
||||||
let r = &agents[self.agent].rating;
|
let r = &agents[self.agent].rating;
|
||||||
let skill = skills.get(self.agent).unwrap();
|
let skill = skills.at(self.slot);
|
||||||
|
|
||||||
if online {
|
if forward {
|
||||||
Rating::new(skill.online, r.beta, r.drift)
|
Rating::new(skill.forward, r.beta, r.drift).with_drift_scale(r.drift_scale)
|
||||||
} else if forward {
|
|
||||||
Rating::new(skill.forward, r.beta, r.drift)
|
|
||||||
} else {
|
} else {
|
||||||
Rating::new(skill.posterior() / self.likelihood, r.beta, r.drift)
|
Rating::new(skill.posterior() / self.likelihood, r.beta, r.drift)
|
||||||
|
.with_drift_scale(r.drift_scale)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Clone, Debug)]
|
||||||
struct Team {
|
struct Team {
|
||||||
items: Vec<Item>,
|
items: Vec<Item>,
|
||||||
output: f64,
|
output: f64,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub(crate) struct Event {
|
pub(crate) struct Event {
|
||||||
teams: Vec<Team>,
|
teams: Vec<Team>,
|
||||||
evidence: f64,
|
log_evidence: f64,
|
||||||
weights: Vec<Vec<f64>>,
|
weights: Vec<Vec<f64>>,
|
||||||
kind: EventKind,
|
kind: EventKind,
|
||||||
}
|
}
|
||||||
@@ -108,7 +110,6 @@ impl Event {
|
|||||||
|
|
||||||
pub(crate) fn within_priors<T: Time, D: Drift<T>>(
|
pub(crate) fn within_priors<T: Time, D: Drift<T>>(
|
||||||
&self,
|
&self,
|
||||||
online: bool,
|
|
||||||
forward: bool,
|
forward: bool,
|
||||||
skills: &SkillStore,
|
skills: &SkillStore,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
@@ -118,25 +119,26 @@ impl Event {
|
|||||||
.map(|team| {
|
.map(|team| {
|
||||||
team.items
|
team.items
|
||||||
.iter()
|
.iter()
|
||||||
.map(|item| item.within_prior(online, forward, skills, agents))
|
.map(|item| item.within_prior(forward, skills, agents))
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Direct in-loop update: mutates self and `skills` inline with no
|
/// Run inference for this event and return its per-item likelihoods.
|
||||||
/// intermediate allocation. Used by both the sequential sweep path and,
|
///
|
||||||
/// via unsafe, by the parallel rayon path for events in the same color
|
/// Reads `skills` immutably and does not touch `self`, so every event in
|
||||||
/// group (which have disjoint agent sets — see `sweep_color_groups`).
|
/// a color group can run concurrently without any aliasing question —
|
||||||
fn iteration_direct<T: Time, D: Drift<T>>(
|
/// the mutation is deferred to `apply`.
|
||||||
&mut self,
|
fn compute<T: Time, D: Drift<T>>(
|
||||||
skills: &mut SkillStore,
|
&self,
|
||||||
|
skills: &SkillStore,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
p_draw: f64,
|
p_draw: f64,
|
||||||
convergence: crate::ConvergenceOptions,
|
convergence: crate::ConvergenceOptions,
|
||||||
arena: &mut ScratchArena,
|
arena: &mut ScratchArena,
|
||||||
) {
|
) -> EventUpdate {
|
||||||
let teams = self.within_priors(false, false, skills, agents);
|
let teams = self.within_priors(false, skills, agents);
|
||||||
let result = self.outputs();
|
let result = self.outputs();
|
||||||
let g = match self.kind {
|
let g = match self.kind {
|
||||||
EventKind::Ranked => {
|
EventKind::Ranked => {
|
||||||
@@ -152,19 +154,60 @@ impl Event {
|
|||||||
),
|
),
|
||||||
};
|
};
|
||||||
|
|
||||||
for (t, team) in self.teams.iter_mut().enumerate() {
|
EventUpdate {
|
||||||
for (i, item) in team.items.iter_mut().enumerate() {
|
log_evidence: g.log_evidence,
|
||||||
let old_likelihood = skills.get(item.agent).unwrap().likelihood;
|
likelihoods: g.likelihoods,
|
||||||
let new_likelihood = (old_likelihood / item.likelihood) * g.likelihoods[t][i];
|
|
||||||
skills.get_mut(item.agent).unwrap().likelihood = new_likelihood;
|
|
||||||
item.likelihood = g.likelihoods[t][i];
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
self.evidence = g.evidence;
|
/// Fold a computed update into the skill store and cache it on the items.
|
||||||
|
fn apply(&mut self, skills: &mut SkillStore, update: EventUpdate) {
|
||||||
|
for (t, team) in self.teams.iter_mut().enumerate() {
|
||||||
|
for (i, item) in team.items.iter_mut().enumerate() {
|
||||||
|
let fresh = update.likelihoods[t][i];
|
||||||
|
let old_likelihood = skills.at(item.slot).likelihood;
|
||||||
|
let new_likelihood = (old_likelihood / item.likelihood) * fresh;
|
||||||
|
skills.at_mut(item.slot).likelihood = new_likelihood;
|
||||||
|
item.likelihood = fresh;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
self.log_evidence = update.log_evidence;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Compute and apply in one step — the sequential sweep.
|
||||||
|
fn iteration_direct<T: Time, D: Drift<T>>(
|
||||||
|
&mut self,
|
||||||
|
skills: &mut SkillStore,
|
||||||
|
agents: &CompetitorStore<T, D>,
|
||||||
|
p_draw: f64,
|
||||||
|
convergence: crate::ConvergenceOptions,
|
||||||
|
arena: &mut ScratchArena,
|
||||||
|
) {
|
||||||
|
let update = self.compute(skills, agents, p_draw, convergence, arena);
|
||||||
|
self.apply(skills, update);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The result of running inference for one event, before it is folded back
|
||||||
|
/// into the shared skill store.
|
||||||
|
#[derive(Debug)]
|
||||||
|
struct EventUpdate {
|
||||||
|
log_evidence: f64,
|
||||||
|
likelihoods: Vec<Vec<Gaussian>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// One slice's worth of forward-only inference.
|
||||||
|
///
|
||||||
|
/// `posteriors` doubles as the outgoing forward message: the scratch sweep
|
||||||
|
/// never writes `backward`, so it stays `N_INF`, and `Skill::posterior()`
|
||||||
|
/// and `forward_prior_out` are then the same product.
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub(crate) struct FilteredStep {
|
||||||
|
pub(crate) log_evidence: f64,
|
||||||
|
pub(crate) posteriors: Vec<(Index, Gaussian)>,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct TimeSlice<T: Time = i64> {
|
pub struct TimeSlice<T: Time = i64> {
|
||||||
pub(crate) events: Vec<Event>,
|
pub(crate) events: Vec<Event>,
|
||||||
@@ -174,6 +217,14 @@ pub struct TimeSlice<T: Time = i64> {
|
|||||||
pub(crate) convergence: crate::ConvergenceOptions,
|
pub(crate) convergence: crate::ConvergenceOptions,
|
||||||
arena: ScratchArena,
|
arena: ScratchArena,
|
||||||
pub(crate) color_groups: ColorGroups,
|
pub(crate) color_groups: ColorGroups,
|
||||||
|
/// Whether `color_groups` still reflects `events`.
|
||||||
|
///
|
||||||
|
/// Coloring is rebuilt lazily, on the first full sweep after an append,
|
||||||
|
/// rather than eagerly per append: the partition is thrown away and
|
||||||
|
/// recomputed wholesale either way, so doing it per append made ingesting
|
||||||
|
/// n events O(n^2) with no benefit — nothing reads the partition between
|
||||||
|
/// an append and the next full sweep.
|
||||||
|
color_groups_dirty: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T: Time> TimeSlice<T> {
|
impl<T: Time> TimeSlice<T> {
|
||||||
@@ -186,6 +237,7 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
convergence,
|
convergence,
|
||||||
arena: ScratchArena::new(),
|
arena: ScratchArena::new(),
|
||||||
color_groups: ColorGroups::new(),
|
color_groups: ColorGroups::new(),
|
||||||
|
color_groups_dirty: false,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -198,6 +250,7 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
let n = self.events.len();
|
let n = self.events.len();
|
||||||
if n == 0 {
|
if n == 0 {
|
||||||
self.color_groups = ColorGroups::new();
|
self.color_groups = ColorGroups::new();
|
||||||
|
self.color_groups_dirty = false;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -221,13 +274,19 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
|
|
||||||
self.events = reordered;
|
self.events = reordered;
|
||||||
self.color_groups = ColorGroups { groups: new_groups };
|
self.color_groups = ColorGroups { groups: new_groups };
|
||||||
|
self.color_groups_dirty = false;
|
||||||
|
|
||||||
|
debug_assert!(
|
||||||
|
self.color_groups.groups_are_contiguous(),
|
||||||
|
"color groups must occupy contiguous event ranges"
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn add_events<D: Drift<T>>(
|
pub fn add_events<D: Drift<T>>(
|
||||||
&mut self,
|
&mut self,
|
||||||
composition: Vec<Vec<Vec<Index>>>,
|
composition: Vec<Vec<Vec<Index>>>,
|
||||||
results: Vec<Vec<f64>>,
|
results: Option<Vec<Vec<f64>>>,
|
||||||
weights: Vec<Vec<Vec<f64>>>,
|
weights: Option<Vec<Vec<Vec<f64>>>>,
|
||||||
kinds: Vec<EventKind>,
|
kinds: Vec<EventKind>,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
) {
|
) {
|
||||||
@@ -246,21 +305,26 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
for idx in this_agent {
|
for idx in this_agent {
|
||||||
let elapsed = compute_elapsed(agents[*idx].last_time.as_ref(), &self.time);
|
let elapsed = compute_elapsed(agents[*idx].last_time.as_ref(), &self.time);
|
||||||
|
|
||||||
|
let forward = agents[*idx].receive(&self.time);
|
||||||
|
|
||||||
if let Some(skill) = self.skills.get_mut(*idx) {
|
if let Some(skill) = self.skills.get_mut(*idx) {
|
||||||
skill.elapsed = elapsed;
|
skill.elapsed = elapsed;
|
||||||
skill.forward = agents[*idx].receive(&self.time);
|
skill.forward = forward;
|
||||||
} else {
|
} else {
|
||||||
self.skills.insert(
|
self.skills.insert(
|
||||||
*idx,
|
*idx,
|
||||||
Skill {
|
Skill {
|
||||||
forward: agents[*idx].receive(&self.time),
|
forward,
|
||||||
|
backward: N_INF,
|
||||||
|
likelihood: N_INF,
|
||||||
elapsed,
|
elapsed,
|
||||||
..Default::default()
|
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let skills = &self.skills;
|
||||||
|
|
||||||
let events = composition.iter().enumerate().map(|(e, event)| {
|
let events = composition.iter().enumerate().map(|(e, event)| {
|
||||||
let teams = event
|
let teams = event
|
||||||
.iter()
|
.iter()
|
||||||
@@ -270,33 +334,37 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|&agent| Item {
|
.map(|&agent| Item {
|
||||||
agent,
|
agent,
|
||||||
|
// Every participant was inserted into `skills`
|
||||||
|
// just above, so the slot always resolves.
|
||||||
|
slot: skills
|
||||||
|
.slot_of(agent)
|
||||||
|
.expect("participant must be present in the slice store"),
|
||||||
likelihood: N_INF,
|
likelihood: N_INF,
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
|
|
||||||
Team {
|
Team {
|
||||||
items,
|
items,
|
||||||
output: if results.is_empty() {
|
output: match &results {
|
||||||
(event.len() - (t + 1)) as f64
|
Some(results) => results[e][t],
|
||||||
} else {
|
// No explicit result: rank by position, first team best.
|
||||||
results[e][t]
|
None => (event.len() - (t + 1)) as f64,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
|
|
||||||
let weights = if weights.is_empty() {
|
let weights = match &weights {
|
||||||
teams
|
Some(weights) => weights[e].clone(),
|
||||||
|
None => teams
|
||||||
.iter()
|
.iter()
|
||||||
.map(|team| vec![1.0; team.items.len()])
|
.map(|team| vec![1.0; team.items.len()])
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>(),
|
||||||
} else {
|
|
||||||
weights[e].clone()
|
|
||||||
};
|
};
|
||||||
|
|
||||||
Event {
|
Event {
|
||||||
teams,
|
teams,
|
||||||
evidence: 0.0,
|
log_evidence: 0.0,
|
||||||
weights,
|
weights,
|
||||||
kind: kinds[e],
|
kind: kinds[e],
|
||||||
}
|
}
|
||||||
@@ -306,8 +374,9 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
|
|
||||||
self.events.extend(events);
|
self.events.extend(events);
|
||||||
|
|
||||||
|
self.color_groups_dirty = true;
|
||||||
|
|
||||||
self.iteration(from, agents);
|
self.iteration(from, agents);
|
||||||
self.recompute_color_groups();
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn posteriors(&self) -> HashMap<Index, Gaussian> {
|
pub(crate) fn posteriors(&self) -> HashMap<Index, Gaussian> {
|
||||||
@@ -317,11 +386,22 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
.collect::<HashMap<_, _>>()
|
.collect::<HashMap<_, _>>()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Sweep this slice's events once, starting at index `from`.
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// Panics if an event references a competitor with no entry in this
|
||||||
|
/// slice's skill store. `add_events` inserts one for every participant, so
|
||||||
|
/// this cannot happen for slices built through the public API.
|
||||||
pub fn iteration<D: Drift<T>>(&mut self, from: usize, agents: &CompetitorStore<T, D>) {
|
pub fn iteration<D: Drift<T>>(&mut self, from: usize, agents: &CompetitorStore<T, D>) {
|
||||||
|
if from == 0 && self.color_groups_dirty {
|
||||||
|
self.recompute_color_groups();
|
||||||
|
}
|
||||||
|
|
||||||
if from > 0 || self.color_groups.is_empty() {
|
if from > 0 || self.color_groups.is_empty() {
|
||||||
// Initial pass (add_events) or no color groups yet: simple sequential sweep.
|
// Initial pass (add_events) or no color groups yet: simple sequential sweep.
|
||||||
for event in self.events.iter_mut().skip(from) {
|
for event in self.events.iter_mut().skip(from) {
|
||||||
let teams = event.within_priors(false, false, &self.skills, agents);
|
let teams = event.within_priors(false, &self.skills, agents);
|
||||||
let result = event.outputs();
|
let result = event.outputs();
|
||||||
|
|
||||||
let g = match event.kind {
|
let g = match event.kind {
|
||||||
@@ -345,15 +425,15 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
|
|
||||||
for (t, team) in event.teams.iter_mut().enumerate() {
|
for (t, team) in event.teams.iter_mut().enumerate() {
|
||||||
for (i, item) in team.items.iter_mut().enumerate() {
|
for (i, item) in team.items.iter_mut().enumerate() {
|
||||||
let old_likelihood = self.skills.get(item.agent).unwrap().likelihood;
|
let old_likelihood = self.skills.at(item.slot).likelihood;
|
||||||
let new_likelihood =
|
let new_likelihood =
|
||||||
(old_likelihood / item.likelihood) * g.likelihoods[t][i];
|
(old_likelihood / item.likelihood) * g.likelihoods[t][i];
|
||||||
self.skills.get_mut(item.agent).unwrap().likelihood = new_likelihood;
|
self.skills.at_mut(item.slot).likelihood = new_likelihood;
|
||||||
item.likelihood = g.likelihoods[t][i];
|
item.likelihood = g.likelihoods[t][i];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
event.evidence = g.evidence;
|
event.log_evidence = g.log_evidence;
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
self.sweep_color_groups(agents);
|
self.sweep_color_groups(agents);
|
||||||
@@ -363,14 +443,13 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
/// Full event sweep using the color-group partition. Colors are processed
|
/// Full event sweep using the color-group partition. Colors are processed
|
||||||
/// sequentially; within each color the inner loop is parallel under rayon.
|
/// sequentially; within each color the inner loop is parallel under rayon.
|
||||||
///
|
///
|
||||||
/// Events within each color group touch disjoint agent sets (guaranteed by
|
/// Events in one color group touch disjoint agent sets, so none of them
|
||||||
/// the greedy coloring). This lets each rayon thread write directly to its
|
/// can observe another's writes. That makes the sweep separable: inference
|
||||||
/// events' skill likelihoods without a deferred-apply step, matching the
|
/// runs concurrently over shared `&self.skills`, and the resulting updates
|
||||||
/// sequential path's allocation profile. The unsafe block is sound because:
|
/// are folded in afterwards in index order. Splitting it this way needs no
|
||||||
/// 1. `self.events[range]` and `self.skills` are separate fields → disjoint.
|
/// `unsafe` and no aliasing argument, and it keeps results bit-identical
|
||||||
/// 2. Events in the same color group access disjoint `Index` values in
|
/// across thread counts because the apply order does not depend on which
|
||||||
/// `self.skills`, so concurrent writes land on different memory locations.
|
/// worker finished first.
|
||||||
/// 3. Each event only writes to its own items' likelihoods (no sharing).
|
|
||||||
#[cfg(feature = "rayon")]
|
#[cfg(feature = "rayon")]
|
||||||
fn sweep_color_groups<D: Drift<T>>(&mut self, agents: &CompetitorStore<T, D>) {
|
fn sweep_color_groups<D: Drift<T>>(&mut self, agents: &CompetitorStore<T, D>) {
|
||||||
use rayon::prelude::*;
|
use rayon::prelude::*;
|
||||||
@@ -390,29 +469,28 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
if group_len == 0 {
|
if group_len == 0 {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
let range = self.color_groups.color_range(color_idx);
|
let range = self.color_groups.color_range(color_idx);
|
||||||
let p_draw = self.p_draw;
|
let p_draw = self.p_draw;
|
||||||
let convergence = self.convergence;
|
let convergence = self.convergence;
|
||||||
|
|
||||||
if group_len >= RAYON_THRESHOLD {
|
if group_len >= RAYON_THRESHOLD {
|
||||||
// Obtain a raw pointer from the unique `&mut self.skills` reference.
|
let skills = &self.skills;
|
||||||
// Casting back to `&mut` inside the closure is sound because:
|
let updates: Vec<EventUpdate> = self.events[range.clone()]
|
||||||
// 1. The pointer originates from a `&mut` — no aliasing with shared refs.
|
.par_iter()
|
||||||
// 2. Events in the same color group touch disjoint `Index` slots in the
|
.map(|ev| {
|
||||||
// underlying Vec, so concurrent writes from different threads land on
|
|
||||||
// different memory locations — no data race.
|
|
||||||
// 3. `self.events[range]` and `self.skills` are separate struct fields,
|
|
||||||
// so the borrow splits cleanly.
|
|
||||||
let skills_addr: usize = (&mut self.skills as *mut SkillStore) as usize;
|
|
||||||
self.events[range].par_iter_mut().for_each(move |ev| {
|
|
||||||
// SAFETY: see above.
|
|
||||||
let skills: &mut SkillStore = unsafe { &mut *(skills_addr as *mut SkillStore) };
|
|
||||||
ARENA.with(|cell| {
|
ARENA.with(|cell| {
|
||||||
let mut arena = cell.borrow_mut();
|
let mut arena = cell.borrow_mut();
|
||||||
arena.reset();
|
arena.reset();
|
||||||
ev.iteration_direct(skills, agents, p_draw, convergence, &mut arena);
|
|
||||||
});
|
ev.compute(skills, agents, p_draw, convergence, &mut arena)
|
||||||
});
|
})
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
for (ev, update) in self.events[range].iter_mut().zip(updates) {
|
||||||
|
ev.apply(&mut self.skills, update);
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
for ev in &mut self.events[range] {
|
for ev in &mut self.events[range] {
|
||||||
ev.iteration_direct(
|
ev.iteration_direct(
|
||||||
@@ -454,18 +532,29 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)]
|
/// Iterate this slice alone until its posteriors stop moving, returning
|
||||||
|
/// the number of iterations taken.
|
||||||
|
///
|
||||||
|
/// Used by `filtered_step` to drive a scratch copy of the slice, and by
|
||||||
|
/// tests. Production convergence across slices is driven by
|
||||||
|
/// `History::converge`, which calls `iteration` directly.
|
||||||
|
///
|
||||||
|
/// Honours `self.convergence`; it previously hard-coded an epsilon and a
|
||||||
|
/// 20-iteration cap that matched neither `ConvergenceOptions` nor the
|
||||||
|
/// schedule default.
|
||||||
pub(crate) fn iterate_to_convergence<D: Drift<T>>(
|
pub(crate) fn iterate_to_convergence<D: Drift<T>>(
|
||||||
&mut self,
|
&mut self,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
) -> usize {
|
) -> usize {
|
||||||
let epsilon = 1e-6;
|
use crate::{tuple_gt, tuple_max};
|
||||||
let iterations = 20;
|
|
||||||
|
let epsilon = self.convergence.epsilon;
|
||||||
|
let max_iter = self.convergence.max_iter;
|
||||||
|
|
||||||
let mut step = (f64::INFINITY, f64::INFINITY);
|
let mut step = (f64::INFINITY, f64::INFINITY);
|
||||||
let mut i = 0;
|
let mut i = 0;
|
||||||
|
|
||||||
while tuple_gt(step, epsilon) && i < iterations {
|
while tuple_gt(step, epsilon) && i < max_iter {
|
||||||
let old = self.posteriors();
|
let old = self.posteriors();
|
||||||
|
|
||||||
self.iteration(0, agents);
|
self.iteration(0, agents);
|
||||||
@@ -477,6 +566,10 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
});
|
});
|
||||||
|
|
||||||
i += 1;
|
i += 1;
|
||||||
|
|
||||||
|
if !crate::step_is_finite(step) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
i
|
i
|
||||||
@@ -497,14 +590,13 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
n.forget(
|
n.forget(
|
||||||
agents[*agent]
|
agents[*agent]
|
||||||
.rating
|
.rating
|
||||||
.drift
|
.drift_variance_for_elapsed(skill.elapsed),
|
||||||
.variance_for_elapsed(skill.elapsed),
|
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn new_backward_info<D: Drift<T>>(&mut self, agents: &CompetitorStore<T, D>) {
|
pub(crate) fn new_backward_info<D: Drift<T>>(&mut self, agents: &CompetitorStore<T, D>) {
|
||||||
for (agent, skill) in self.skills.iter_mut() {
|
for (agent, skill) in self.skills.iter_mut() {
|
||||||
skill.backward = agents[agent].message;
|
skill.backward = agents[agent].message.unwrap_or(N_INF);
|
||||||
}
|
}
|
||||||
self.iteration(0, agents);
|
self.iteration(0, agents);
|
||||||
}
|
}
|
||||||
@@ -516,21 +608,99 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
self.iteration(0, agents);
|
self.iteration(0, agents);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Run this slice's events on forward (filtering) information alone.
|
||||||
|
///
|
||||||
|
/// `incoming` holds each competitor's forward message out of their
|
||||||
|
/// previous appearance; a competitor absent from it starts at their
|
||||||
|
/// configured prior. The sweep runs on a scratch copy, so the real slice
|
||||||
|
/// is untouched — which is what makes the filtered estimates independent
|
||||||
|
/// of whether `History::converge` has run.
|
||||||
|
pub(crate) fn filtered_step<D: Drift<T>>(
|
||||||
|
&self,
|
||||||
|
incoming: &HashMap<Index, Gaussian>,
|
||||||
|
agents: &CompetitorStore<T, D>,
|
||||||
|
) -> FilteredStep {
|
||||||
|
let mut scratch = TimeSlice {
|
||||||
|
events: self.events.clone(),
|
||||||
|
skills: SkillStore::new(),
|
||||||
|
time: self.time,
|
||||||
|
p_draw: self.p_draw,
|
||||||
|
convergence: self.convergence,
|
||||||
|
arena: ScratchArena::new(),
|
||||||
|
color_groups: ColorGroups::new(),
|
||||||
|
color_groups_dirty: true,
|
||||||
|
};
|
||||||
|
|
||||||
|
for event in &mut scratch.events {
|
||||||
|
for team in &mut event.teams {
|
||||||
|
for item in &mut team.items {
|
||||||
|
item.likelihood = N_INF;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
event.log_evidence = 0.0;
|
||||||
|
}
|
||||||
|
|
||||||
|
for (agent, skill) in self.skills.iter() {
|
||||||
|
let rating = &agents[agent].rating;
|
||||||
|
|
||||||
|
let forward = match incoming.get(&agent) {
|
||||||
|
Some(message) => message.forget(rating.drift_variance_for_elapsed(skill.elapsed)),
|
||||||
|
None => rating.prior,
|
||||||
|
};
|
||||||
|
|
||||||
|
let slot = scratch.skills.insert(
|
||||||
|
agent,
|
||||||
|
Skill {
|
||||||
|
forward,
|
||||||
|
backward: N_INF,
|
||||||
|
likelihood: N_INF,
|
||||||
|
elapsed: skill.elapsed,
|
||||||
|
},
|
||||||
|
);
|
||||||
|
|
||||||
|
// The cloned events carry slots resolved against the REAL store, so
|
||||||
|
// the scratch must assign the same ones. It does because `iter()`
|
||||||
|
// yields slot order and `insert` allocates slots in call order —
|
||||||
|
// but that is a coupling between two types, so pin it here rather
|
||||||
|
// than leave it to be rediscovered after it breaks.
|
||||||
|
debug_assert_eq!(
|
||||||
|
Some(slot),
|
||||||
|
self.skills.slot_of(agent),
|
||||||
|
"scratch slot must match the real slice's slot for {agent:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch.iterate_to_convergence(agents);
|
||||||
|
|
||||||
|
FilteredStep {
|
||||||
|
log_evidence: scratch.events.iter().map(|event| event.log_evidence).sum(),
|
||||||
|
posteriors: scratch
|
||||||
|
.skills
|
||||||
|
.iter()
|
||||||
|
.map(|(agent, skill)| (agent, skill.posterior()))
|
||||||
|
.collect(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn log_evidence<D: Drift<T>>(
|
pub(crate) fn log_evidence<D: Drift<T>>(
|
||||||
&self,
|
&self,
|
||||||
online: bool,
|
|
||||||
targets: &[Index],
|
targets: &[Index],
|
||||||
forward: bool,
|
forward: bool,
|
||||||
agents: &CompetitorStore<T, D>,
|
agents: &CompetitorStore<T, D>,
|
||||||
) -> f64 {
|
) -> f64 {
|
||||||
|
// Hashed once rather than scanned per player per event, so a
|
||||||
|
// `log_evidence_for` with many keys is not quadratic.
|
||||||
|
let target_set: std::collections::HashSet<Index> = targets.iter().copied().collect();
|
||||||
// log_evidence is infrequent; a local arena avoids needing &mut self.
|
// log_evidence is infrequent; a local arena avoids needing &mut self.
|
||||||
let mut arena = ScratchArena::new();
|
let mut arena = ScratchArena::new();
|
||||||
|
|
||||||
let run_event = |event: &Event, arena: &mut ScratchArena| -> f64 {
|
let run_event = |event: &Event, arena: &mut ScratchArena| -> f64 {
|
||||||
let teams = event.within_priors(online, forward, &self.skills, agents);
|
let teams = event.within_priors(forward, &self.skills, agents);
|
||||||
let result = event.outputs();
|
let result = event.outputs();
|
||||||
match event.kind {
|
match event.kind {
|
||||||
EventKind::Ranked => Game::ranked_with_arena(
|
EventKind::Ranked => {
|
||||||
|
Game::ranked_with_arena(
|
||||||
teams,
|
teams,
|
||||||
&result,
|
&result,
|
||||||
&event.weights,
|
&event.weights,
|
||||||
@@ -538,9 +708,10 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
self.convergence,
|
self.convergence,
|
||||||
arena,
|
arena,
|
||||||
)
|
)
|
||||||
.evidence
|
.log_evidence
|
||||||
.ln(),
|
}
|
||||||
EventKind::Scored { score_sigma } => Game::scored_with_arena(
|
EventKind::Scored { score_sigma } => {
|
||||||
|
Game::scored_with_arena(
|
||||||
teams,
|
teams,
|
||||||
&result,
|
&result,
|
||||||
&event.weights,
|
&event.weights,
|
||||||
@@ -548,21 +719,21 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
self.convergence,
|
self.convergence,
|
||||||
arena,
|
arena,
|
||||||
)
|
)
|
||||||
.evidence
|
.log_evidence
|
||||||
.ln(),
|
}
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
if targets.is_empty() {
|
if targets.is_empty() {
|
||||||
if online || forward {
|
if forward {
|
||||||
self.events
|
self.events
|
||||||
.iter()
|
.iter()
|
||||||
.map(|event| run_event(event, &mut arena))
|
.map(|event| run_event(event, &mut arena))
|
||||||
.sum()
|
.sum()
|
||||||
} else {
|
} else {
|
||||||
self.events.iter().map(|event| event.evidence.ln()).sum()
|
self.events.iter().map(|event| event.log_evidence).sum()
|
||||||
}
|
}
|
||||||
} else if online || forward {
|
} else if forward {
|
||||||
self.events
|
self.events
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|event| {
|
.filter(|event| {
|
||||||
@@ -570,7 +741,7 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
.teams
|
.teams
|
||||||
.iter()
|
.iter()
|
||||||
.flat_map(|team| &team.items)
|
.flat_map(|team| &team.items)
|
||||||
.any(|item| targets.contains(&item.agent))
|
.any(|item| target_set.contains(&item.agent))
|
||||||
})
|
})
|
||||||
.map(|event| run_event(event, &mut arena))
|
.map(|event| run_event(event, &mut arena))
|
||||||
.sum()
|
.sum()
|
||||||
@@ -582,9 +753,9 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
.teams
|
.teams
|
||||||
.iter()
|
.iter()
|
||||||
.flat_map(|team| &team.items)
|
.flat_map(|team| &team.items)
|
||||||
.any(|item| targets.contains(&item.agent))
|
.any(|item| target_set.contains(&item.agent))
|
||||||
})
|
})
|
||||||
.map(|event| event.evidence.ln())
|
.map(|event| event.log_evidence)
|
||||||
.sum()
|
.sum()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -616,8 +787,26 @@ impl<T: Time> TimeSlice<T> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Elapsed time from a competitor's previous appearance to `current`.
|
||||||
|
///
|
||||||
|
/// A negative elapsed means slices are being visited out of time order, which
|
||||||
|
/// would make drift *reduce* uncertainty. Release builds clamp to zero so a
|
||||||
|
/// bad timestamp degrades to "no drift" rather than corrupting the posterior;
|
||||||
|
/// debug builds trip instead, because reaching here is a bug in slice ordering
|
||||||
|
/// rather than something callers can cause with ordinary data.
|
||||||
pub(crate) fn compute_elapsed<T: Time>(last: Option<&T>, current: &T) -> i64 {
|
pub(crate) fn compute_elapsed<T: Time>(last: Option<&T>, current: &T) -> i64 {
|
||||||
last.map(|l| l.elapsed_to(current).max(0)).unwrap_or(0)
|
let Some(last) = last else {
|
||||||
|
return 0;
|
||||||
|
};
|
||||||
|
|
||||||
|
let elapsed = last.elapsed_to(current);
|
||||||
|
|
||||||
|
debug_assert!(
|
||||||
|
elapsed >= 0,
|
||||||
|
"negative elapsed ({elapsed}) — slices visited out of time order"
|
||||||
|
);
|
||||||
|
|
||||||
|
elapsed.max(0)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -665,8 +854,8 @@ mod tests {
|
|||||||
vec![vec![c], vec![d]],
|
vec![vec![c], vec![d]],
|
||||||
vec![vec![e], vec![f]],
|
vec![vec![e], vec![f]],
|
||||||
],
|
],
|
||||||
vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]],
|
Some(vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]]),
|
||||||
vec![],
|
None,
|
||||||
vec![EventKind::Ranked; 3],
|
vec![EventKind::Ranked; 3],
|
||||||
&agents,
|
&agents,
|
||||||
);
|
);
|
||||||
@@ -742,8 +931,8 @@ mod tests {
|
|||||||
vec![vec![a], vec![c]],
|
vec![vec![a], vec![c]],
|
||||||
vec![vec![b], vec![c]],
|
vec![vec![b], vec![c]],
|
||||||
],
|
],
|
||||||
vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]],
|
Some(vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]]),
|
||||||
vec![],
|
None,
|
||||||
vec![EventKind::Ranked; 3],
|
vec![EventKind::Ranked; 3],
|
||||||
&agents,
|
&agents,
|
||||||
);
|
);
|
||||||
@@ -822,8 +1011,8 @@ mod tests {
|
|||||||
vec![vec![a], vec![c]],
|
vec![vec![a], vec![c]],
|
||||||
vec![vec![b], vec![c]],
|
vec![vec![b], vec![c]],
|
||||||
],
|
],
|
||||||
vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]],
|
Some(vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]]),
|
||||||
vec![],
|
None,
|
||||||
vec![EventKind::Ranked; 3],
|
vec![EventKind::Ranked; 3],
|
||||||
&agents,
|
&agents,
|
||||||
);
|
);
|
||||||
@@ -854,8 +1043,8 @@ mod tests {
|
|||||||
vec![vec![a], vec![c]],
|
vec![vec![a], vec![c]],
|
||||||
vec![vec![b], vec![c]],
|
vec![vec![b], vec![c]],
|
||||||
],
|
],
|
||||||
vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]],
|
Some(vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 0.0]]),
|
||||||
vec![],
|
None,
|
||||||
vec![EventKind::Ranked; 3],
|
vec![EventKind::Ranked; 3],
|
||||||
&agents,
|
&agents,
|
||||||
);
|
);
|
||||||
@@ -866,19 +1055,24 @@ mod tests {
|
|||||||
|
|
||||||
let post = time_slice.posteriors();
|
let post = time_slice.posteriors();
|
||||||
|
|
||||||
|
// These are convergence residuals, not exact values: by symmetry the
|
||||||
|
// true mean is 25.0 and the iteration approaches it from above. The
|
||||||
|
// previous expectation of 25.000003 was the residual after the
|
||||||
|
// hard-coded 20-iteration cap; honouring `ConvergenceOptions` runs to
|
||||||
|
// 30 and lands nearer the truth.
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
post[&a],
|
post[&a],
|
||||||
Gaussian::from_ms(25.000003, 3.880150),
|
Gaussian::from_ms(25.000001, 3.880150),
|
||||||
epsilon = 1e-6
|
epsilon = 1e-6
|
||||||
);
|
);
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
post[&b],
|
post[&b],
|
||||||
Gaussian::from_ms(25.000003, 3.880150),
|
Gaussian::from_ms(25.000001, 3.880150),
|
||||||
epsilon = 1e-6
|
epsilon = 1e-6
|
||||||
);
|
);
|
||||||
assert_ulps_eq!(
|
assert_ulps_eq!(
|
||||||
post[&c],
|
post[&c],
|
||||||
Gaussian::from_ms(25.000003, 3.880150),
|
Gaussian::from_ms(25.000001, 3.880150),
|
||||||
epsilon = 1e-6
|
epsilon = 1e-6
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
@@ -920,8 +1114,8 @@ mod tests {
|
|||||||
vec![vec![c], vec![d]],
|
vec![vec![c], vec![d]],
|
||||||
vec![vec![a], vec![c]],
|
vec![vec![a], vec![c]],
|
||||||
],
|
],
|
||||||
vec![vec![1.0, 0.0], vec![1.0, 0.0], vec![1.0, 0.0]],
|
Some(vec![vec![1.0, 0.0], vec![1.0, 0.0], vec![1.0, 0.0]]),
|
||||||
vec![],
|
None,
|
||||||
vec![EventKind::Ranked; 3],
|
vec![EventKind::Ranked; 3],
|
||||||
&agents,
|
&agents,
|
||||||
);
|
);
|
||||||
|
|||||||
+52
-5
@@ -203,7 +203,7 @@ fn predict_quality_two_teams() {
|
|||||||
h.record_winner(&"a", &"b", 1).unwrap();
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
h.converge().unwrap();
|
h.converge().unwrap();
|
||||||
|
|
||||||
let q = h.predict_quality(&[&[&"a"], &[&"b"]]);
|
let q = h.predict_quality(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
assert!(q > 0.0 && q <= 1.0);
|
assert!(q > 0.0 && q <= 1.0);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -219,10 +219,14 @@ fn predict_outcome_two_teams_sums_to_one() {
|
|||||||
h.record_winner(&"a", &"b", 1).unwrap();
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
h.converge().unwrap();
|
h.converge().unwrap();
|
||||||
|
|
||||||
let p = h.predict_outcome(&[&[&"a"], &[&"b"]]);
|
let p = h.predict_outcome(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
assert_eq!(p.len(), 2);
|
let wins = p.win_probabilities();
|
||||||
assert!((p[0] + p[1] - 1.0).abs() < 1e-9);
|
assert_eq!(wins.len(), 2);
|
||||||
assert!(p[0] > p[1]);
|
// With p_draw == 0 there is no draw outcome, so the two win
|
||||||
|
// probabilities are the whole space.
|
||||||
|
assert!((p.total() - 1.0).abs() < 1e-9, "total = {}", p.total());
|
||||||
|
assert!((wins[0] + wins[1] - 1.0).abs() < 1e-9);
|
||||||
|
assert!(wins[0] > wins[1]);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -247,3 +251,46 @@ fn fluent_event_builder_scores() {
|
|||||||
let b = h.current_skill(&"bob").unwrap();
|
let b = h.current_skill(&"bob").unwrap();
|
||||||
assert!(a.mu() > b.mu());
|
assert!(a.mu() > b.mu());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Every field of `ConvergenceReport` must carry real information.
|
||||||
|
///
|
||||||
|
/// `slices_skipped` was public, hardcoded to `0`, and reported a plausible
|
||||||
|
/// value for a feature that never existed — the same shape as the inert
|
||||||
|
/// `online` flag in #19. It was removed in #33. This pins the remaining fields
|
||||||
|
/// so the next always-constant member has to survive an assertion rather than
|
||||||
|
/// just a reviewer's attention.
|
||||||
|
#[test]
|
||||||
|
fn every_convergence_report_field_is_populated() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
|
||||||
|
for time in 1..=6i64 {
|
||||||
|
h.record_winner(&"a", &"b", time).unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
report.iterations > 0,
|
||||||
|
"iterations is zero on a real converge"
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(report.converged, "fixture must converge");
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
report.final_step.0.is_finite() && report.final_step.1.is_finite(),
|
||||||
|
"final_step is not finite: {:?}",
|
||||||
|
report.final_step
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
report.log_evidence.is_finite() && report.log_evidence < 0.0,
|
||||||
|
"log_evidence is not a finite negative log probability: {}",
|
||||||
|
report.log_evidence
|
||||||
|
);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
report.per_iteration_time.len(),
|
||||||
|
report.iterations,
|
||||||
|
"per_iteration_time must carry one duration per iteration"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
//! Helpers shared across the integration suites.
|
||||||
|
//!
|
||||||
|
//! Each integration file is its own binary, so `mod common;` compiles a copy
|
||||||
|
//! per suite. Anything unused in a given suite would warn, hence the
|
||||||
|
//! `#![allow(dead_code)]`.
|
||||||
|
|
||||||
|
#![allow(dead_code)]
|
||||||
|
|
||||||
|
use trueskill_tt::Gaussian;
|
||||||
|
|
||||||
|
/// A posterior must be finite with a strictly positive sigma.
|
||||||
|
///
|
||||||
|
/// A non-finite posterior is the failure mode this crate is most prone to —
|
||||||
|
/// EP breaking down produces NaN rather than an error — and a zero or negative
|
||||||
|
/// sigma means the precision went non-positive, which `Gaussian::sigma` reports
|
||||||
|
/// as improper rather than trapping.
|
||||||
|
pub fn assert_finite(g: Gaussian, what: &str) {
|
||||||
|
assert!(
|
||||||
|
g.mu().is_finite(),
|
||||||
|
"{what}: mu is not finite (mu={}, sigma={})",
|
||||||
|
g.mu(),
|
||||||
|
g.sigma()
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
g.sigma().is_finite() && g.sigma() > 0.0,
|
||||||
|
"{what}: sigma must be finite and positive (mu={}, sigma={})",
|
||||||
|
g.mu(),
|
||||||
|
g.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Every point on every learning curve must be finite.
|
||||||
|
pub fn assert_curve_finite(curve: &[(i64, Gaussian)], who: &str) {
|
||||||
|
for (time, g) in curve {
|
||||||
|
assert_finite(*g, &format!("{who} at t={time}"));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,222 @@
|
|||||||
|
//! `Member::with_prior` / `with_drift_scale` — competitor configuration.
|
||||||
|
//!
|
||||||
|
//! Both were previously consumed only on the branch that *creates* a
|
||||||
|
//! competitor, so configuration supplied for a key the history already knew was
|
||||||
|
//! dropped with no error. `with_prior` had no coverage in this directory at
|
||||||
|
//! all, which is how that survived.
|
||||||
|
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{
|
||||||
|
ConvergenceOptions, Event, Gaussian, History, InferenceError, Member, Outcome, Team,
|
||||||
|
};
|
||||||
|
|
||||||
|
const CONVERGENCE: ConvergenceOptions = ConvergenceOptions {
|
||||||
|
max_iter: 2_000,
|
||||||
|
epsilon: 1e-12,
|
||||||
|
alpha: 1.0,
|
||||||
|
};
|
||||||
|
|
||||||
|
fn history() -> History {
|
||||||
|
History::builder()
|
||||||
|
.mu(25.0)
|
||||||
|
.sigma(25.0 / 3.0)
|
||||||
|
.beta(25.0 / 6.0)
|
||||||
|
.p_draw(0.0)
|
||||||
|
.convergence(CONVERGENCE)
|
||||||
|
.build()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// One event, optionally configuring `a`.
|
||||||
|
fn bout(
|
||||||
|
a: &'static str,
|
||||||
|
b: &'static str,
|
||||||
|
time: i64,
|
||||||
|
prior: Option<Gaussian>,
|
||||||
|
scale: Option<f64>,
|
||||||
|
) -> Event<i64, &'static str> {
|
||||||
|
let mut member = Member::new(a);
|
||||||
|
if let Some(p) = prior {
|
||||||
|
member = member.with_prior(p);
|
||||||
|
}
|
||||||
|
if let Some(s) = scale {
|
||||||
|
member = member.with_drift_scale(s);
|
||||||
|
}
|
||||||
|
|
||||||
|
Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([member]),
|
||||||
|
Team::with_members([Member::new(b)]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn skill_of(h: &History, key: &str) -> Gaussian {
|
||||||
|
h.current_skill(&key).expect("key in history")
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Baseline: the mechanism works at all on a competitor's first appearance.
|
||||||
|
#[test]
|
||||||
|
fn a_prior_applies_to_a_new_competitor() {
|
||||||
|
let seeded = Gaussian::from_ms(40.0, 1.0);
|
||||||
|
|
||||||
|
let mut with = history();
|
||||||
|
with.add_events(vec![bout("a", "b", 0, Some(seeded), None)])
|
||||||
|
.unwrap();
|
||||||
|
with.converge().unwrap();
|
||||||
|
|
||||||
|
let mut without = history();
|
||||||
|
without
|
||||||
|
.add_events(vec![bout("a", "b", 0, None, None)])
|
||||||
|
.unwrap();
|
||||||
|
without.converge().unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(skill_of(&with, "a").mu() - skill_of(&without, "a").mu()).abs() > 1.0,
|
||||||
|
"a seeded prior should move the fit"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The defect in #10: a prior supplied for a competitor the history already
|
||||||
|
/// knows was silently discarded, and the caller got output computed from the
|
||||||
|
/// default prior with no indication anything had been dropped.
|
||||||
|
#[test]
|
||||||
|
fn a_prior_applies_to_a_competitor_the_history_already_knows() {
|
||||||
|
let seeded = Gaussian::from_ms(40.0, 1.0);
|
||||||
|
|
||||||
|
let mut late = history();
|
||||||
|
late.add_events(vec![bout("a", "b", 0, None, None)])
|
||||||
|
.unwrap();
|
||||||
|
// "a" now exists. Configuring it here used to do nothing whatsoever.
|
||||||
|
late.add_events(vec![bout("a", "b", 1, Some(seeded), None)])
|
||||||
|
.unwrap();
|
||||||
|
late.converge().unwrap();
|
||||||
|
|
||||||
|
let mut never = history();
|
||||||
|
never
|
||||||
|
.add_events(vec![
|
||||||
|
bout("a", "b", 0, None, None),
|
||||||
|
bout("a", "b", 1, None, None),
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
never.converge().unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(skill_of(&late, "a").mu() - skill_of(&never, "a").mu()).abs() > 1.0,
|
||||||
|
"a late prior must not be silently dropped: {} vs {}",
|
||||||
|
skill_of(&late, "a").mu(),
|
||||||
|
skill_of(&never, "a").mu()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Configuration is competitor-scoped, not event-scoped, and `converge` refits
|
||||||
|
/// from competitor state — so seeding late reaches the same fit as seeding from
|
||||||
|
/// the start. This is the documented scope, asserted rather than assumed.
|
||||||
|
#[test]
|
||||||
|
fn a_prior_is_whole_history_scoped_not_per_event() {
|
||||||
|
let seeded = Gaussian::from_ms(40.0, 1.0);
|
||||||
|
|
||||||
|
let mut late = history();
|
||||||
|
late.add_events(vec![bout("a", "b", 0, None, None)])
|
||||||
|
.unwrap();
|
||||||
|
late.add_events(vec![bout("a", "b", 1, Some(seeded), None)])
|
||||||
|
.unwrap();
|
||||||
|
late.converge().unwrap();
|
||||||
|
|
||||||
|
let mut early = history();
|
||||||
|
early
|
||||||
|
.add_events(vec![
|
||||||
|
bout("a", "b", 0, Some(seeded), None),
|
||||||
|
bout("a", "b", 1, Some(seeded), None),
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
early.converge().unwrap();
|
||||||
|
|
||||||
|
let (l, e) = (skill_of(&late, "a"), skill_of(&early, "a"));
|
||||||
|
assert!(
|
||||||
|
(l.mu() - e.mu()).abs() < 1e-9 && (l.sigma() - e.sigma()).abs() < 1e-9,
|
||||||
|
"late seeding should refit the whole history: {l:?} vs {e:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn repeating_the_same_prior_is_inert() {
|
||||||
|
let seeded = Gaussian::from_ms(40.0, 1.0);
|
||||||
|
|
||||||
|
let mut once = history();
|
||||||
|
once.add_events(vec![
|
||||||
|
bout("a", "b", 0, Some(seeded), None),
|
||||||
|
bout("a", "b", 1, None, None),
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
once.converge().unwrap();
|
||||||
|
|
||||||
|
let mut every_time = history();
|
||||||
|
every_time
|
||||||
|
.add_events(vec![
|
||||||
|
bout("a", "b", 0, Some(seeded), None),
|
||||||
|
bout("a", "b", 1, Some(seeded), None),
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
every_time.converge().unwrap();
|
||||||
|
|
||||||
|
let (o, e) = (skill_of(&once, "a"), skill_of(&every_time, "a"));
|
||||||
|
assert!(
|
||||||
|
(o.mu() - e.mu()).abs() < 1e-12 && (o.sigma() - e.sigma()).abs() < 1e-12,
|
||||||
|
"declaring the same prior repeatedly changed the fit: {o:?} vs {e:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Events within a batch have no order, so two different values for one
|
||||||
|
/// competitor have no well-defined winner. Rejecting is what keeps the answer
|
||||||
|
/// independent of iteration order.
|
||||||
|
#[test]
|
||||||
|
fn a_batch_declaring_two_different_priors_is_rejected() {
|
||||||
|
let mut h = history();
|
||||||
|
let err = h
|
||||||
|
.add_events(vec![
|
||||||
|
bout("a", "b", 0, Some(Gaussian::from_ms(40.0, 1.0)), None),
|
||||||
|
bout("a", "b", 1, Some(Gaussian::from_ms(10.0, 1.0)), None),
|
||||||
|
])
|
||||||
|
.expect_err("two different priors for one competitor in one batch");
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::ConflictingCompetitorConfig { field: "prior", .. }
|
||||||
|
),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A member setting only `drift_scale` must not also assert the default prior,
|
||||||
|
/// or it would silently undo a prior seeded earlier. This is why the collected
|
||||||
|
/// configuration tracks each field separately rather than a merged `Rating`.
|
||||||
|
#[test]
|
||||||
|
fn setting_one_field_late_leaves_the_other_alone() {
|
||||||
|
let seeded = Gaussian::from_ms(40.0, 1.0);
|
||||||
|
|
||||||
|
let mut h = history();
|
||||||
|
h.add_events(vec![bout("a", "b", 0, Some(seeded), None)])
|
||||||
|
.unwrap();
|
||||||
|
// Only the scale this time — the prior above must survive.
|
||||||
|
h.add_events(vec![bout("a", "b", 1, None, Some(0.5))])
|
||||||
|
.unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let mut both_upfront = history();
|
||||||
|
both_upfront
|
||||||
|
.add_events(vec![
|
||||||
|
bout("a", "b", 0, Some(seeded), Some(0.5)),
|
||||||
|
bout("a", "b", 1, None, None),
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
both_upfront.converge().unwrap();
|
||||||
|
|
||||||
|
let (a, b) = (skill_of(&h, "a"), skill_of(&both_upfront, "a"));
|
||||||
|
assert!(
|
||||||
|
(a.mu() - b.mu()).abs() < 1e-9 && (a.sigma() - b.sigma()).abs() < 1e-9,
|
||||||
|
"setting drift_scale late clobbered the earlier prior: {a:?} vs {b:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -0,0 +1,423 @@
|
|||||||
|
//! Degenerate, boundary, and error-path coverage.
|
||||||
|
//!
|
||||||
|
//! These run in both debug and release: the defects they pin were all
|
||||||
|
//! guarded only by `debug_assert!`, so a debug-only suite never saw them.
|
||||||
|
|
||||||
|
mod common;
|
||||||
|
|
||||||
|
use common::assert_finite;
|
||||||
|
use trueskill_tt::{
|
||||||
|
ConstantDrift, ConvergenceOptions, Game, GameOptions, Gaussian, History, InferenceError,
|
||||||
|
NullObserver, Outcome, Rating,
|
||||||
|
};
|
||||||
|
|
||||||
|
type R = Rating<i64, ConstantDrift>;
|
||||||
|
|
||||||
|
fn rating() -> R {
|
||||||
|
R::new(
|
||||||
|
Gaussian::from_ms(25.0, 25.0 / 3.0),
|
||||||
|
25.0 / 6.0,
|
||||||
|
ConstantDrift(25.0 / 300.0),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn record_draw_without_draw_probability_is_rejected() {
|
||||||
|
let mut h = History::default();
|
||||||
|
let err = h.record_draw(&"a", &"b", 1).unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::TieWithoutDrawProbability { .. }
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn builder_draw_without_draw_probability_is_rejected() {
|
||||||
|
let mut h = History::default();
|
||||||
|
let err = h
|
||||||
|
.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.team(["b"])
|
||||||
|
.draw()
|
||||||
|
.commit()
|
||||||
|
.unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::TieWithoutDrawProbability { .. }
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn draw_with_positive_draw_probability_is_finite() {
|
||||||
|
let mut h = History::builder().p_draw(0.25).build();
|
||||||
|
h.record_draw(&"a", &"b", 1).unwrap();
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
|
||||||
|
assert_finite(h.current_skill("a").unwrap(), "drawn competitor skill");
|
||||||
|
assert_finite(h.current_skill("b").unwrap(), "drawn competitor skill");
|
||||||
|
assert!(report.log_evidence.is_finite());
|
||||||
|
assert!(report.converged);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn game_ranked_rejects_tie_without_draw_probability() {
|
||||||
|
let a = [rating()];
|
||||||
|
let b = [rating()];
|
||||||
|
let teams: Vec<&[R]> = vec![&a, &b];
|
||||||
|
let err = Game::ranked(&teams, Outcome::draw(2), &GameOptions::default()).unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::TieWithoutDrawProbability { .. }
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `Outcome::winner(w, n)` ties every loser, so any n >= 3 free-for-all hits
|
||||||
|
/// the tie path even though the caller never asked for a draw.
|
||||||
|
#[test]
|
||||||
|
fn winner_of_three_or_more_requires_draw_probability() {
|
||||||
|
let a = [rating()];
|
||||||
|
let b = [rating()];
|
||||||
|
let c = [rating()];
|
||||||
|
let teams: Vec<&[R]> = vec![&a, &b, &c];
|
||||||
|
|
||||||
|
let err = Game::ranked(&teams, Outcome::winner(0, 3), &GameOptions::default()).unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::TieWithoutDrawProbability { .. }
|
||||||
|
));
|
||||||
|
|
||||||
|
let opts = GameOptions {
|
||||||
|
p_draw: 0.1,
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
let game = Game::ranked(&teams, Outcome::winner(0, 3), &opts).unwrap();
|
||||||
|
for team in game.posteriors() {
|
||||||
|
for skill in team {
|
||||||
|
assert_finite(skill, "3-team winner posterior");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn full_ranking_without_ties_needs_no_draw_probability() {
|
||||||
|
let a = [rating()];
|
||||||
|
let b = [rating()];
|
||||||
|
let c = [rating()];
|
||||||
|
let teams: Vec<&[R]> = vec![&a, &b, &c];
|
||||||
|
let game = Game::ranked(&teams, Outcome::ranking([0, 1, 2]), &GameOptions::default()).unwrap();
|
||||||
|
|
||||||
|
for team in game.posteriors() {
|
||||||
|
for skill in team {
|
||||||
|
assert_finite(skill, "strict ranking posterior");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_history_converges_trivially() {
|
||||||
|
let mut h = History::default();
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
assert_eq!(report.iterations, 0);
|
||||||
|
assert!(report.converged);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Issue #27's exact reproduction: a non-default key type reaching `converge`
|
||||||
|
/// with no events at all. The underflow it reported trapped in debug and
|
||||||
|
/// indexed out of bounds in release, so this must run in both profiles.
|
||||||
|
#[test]
|
||||||
|
fn converge_on_an_empty_history_with_owned_keys() {
|
||||||
|
let mut history: History<i64, ConstantDrift, NullObserver, String> =
|
||||||
|
History::builder_with_key().score_sigma(5.0).build();
|
||||||
|
|
||||||
|
let report = history.converge().unwrap();
|
||||||
|
|
||||||
|
assert_eq!(report.iterations, 0);
|
||||||
|
assert!(report.converged);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A weights/team length mismatch used to be a `debug_assert!`, so release
|
||||||
|
/// builds ingested the event with the weights silently unapplied. This file's
|
||||||
|
/// CI job runs in release too, which is the point of pinning it here.
|
||||||
|
#[test]
|
||||||
|
fn event_builder_rejects_a_weights_length_mismatch() {
|
||||||
|
let mut h = History::default();
|
||||||
|
|
||||||
|
let err = h
|
||||||
|
.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.weights([1.0, 2.0])
|
||||||
|
.team(["b"])
|
||||||
|
.winner(0)
|
||||||
|
.commit()
|
||||||
|
.unwrap_err();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::MismatchedShape {
|
||||||
|
kind: "weights",
|
||||||
|
expected: 1,
|
||||||
|
got: 2,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
"expected a weights MismatchedShape, got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The mismatch must not be applied even partially — a half-weighted team
|
||||||
|
/// reaching the history would be worse than the error.
|
||||||
|
#[test]
|
||||||
|
fn event_builder_weights_mismatch_leaves_the_history_untouched() {
|
||||||
|
let mut h = History::default();
|
||||||
|
|
||||||
|
// Two teams, so ingestion would otherwise succeed — a one-team event is
|
||||||
|
// rejected for an unrelated reason and would pass this vacuously.
|
||||||
|
let _ = h
|
||||||
|
.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.weights([1.0, 2.0])
|
||||||
|
.team(["b"])
|
||||||
|
.winner(0)
|
||||||
|
.commit();
|
||||||
|
|
||||||
|
assert!(h.learning_curve("a").is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_event_stream_then_converge() {
|
||||||
|
let mut h = History::default();
|
||||||
|
h.add_events(std::iter::empty()).unwrap();
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
assert_eq!(report.iterations, 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_history_queries_do_not_panic() {
|
||||||
|
let h = History::default();
|
||||||
|
assert!(h.learning_curves().is_empty());
|
||||||
|
assert!(h.learning_curve("nobody").is_empty());
|
||||||
|
assert!(h.current_skill("nobody").is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn single_event_history_converges() {
|
||||||
|
let mut h = History::default();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
assert!(report.converged);
|
||||||
|
assert_finite(h.current_skill("a").unwrap(), "single-event skill");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn scored_event_rejects_non_positive_sigma() {
|
||||||
|
let mut h = History::builder().score_sigma(2.0).build();
|
||||||
|
let err = h
|
||||||
|
.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.team(["b"])
|
||||||
|
.scores_with_sigma([3.0, 1.0], f64::NAN)
|
||||||
|
.commit()
|
||||||
|
.unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::InvalidParameter {
|
||||||
|
name: "score_sigma",
|
||||||
|
..
|
||||||
|
}
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn convergence_reports_are_finite_across_many_teams() {
|
||||||
|
let opts = GameOptions {
|
||||||
|
p_draw: 0.1,
|
||||||
|
convergence: ConvergenceOptions::default(),
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
let holders: Vec<[R; 1]> = (0..12).map(|_| [rating()]).collect();
|
||||||
|
let teams: Vec<&[R]> = holders.iter().map(|t| t.as_slice()).collect();
|
||||||
|
let game = Game::ranked(&teams, Outcome::ranking(0..12), &opts).unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
game.log_evidence().is_finite(),
|
||||||
|
"12-team log-evidence must be finite, got {}",
|
||||||
|
game.log_evidence()
|
||||||
|
);
|
||||||
|
for team in game.posteriors() {
|
||||||
|
for skill in team {
|
||||||
|
assert_finite(skill, "12-team posterior");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A long diff chain underflows a linear evidence product: each link
|
||||||
|
/// contributes a probability in (0, 1], so ~1000 links flush the product to
|
||||||
|
/// exactly 0.0 and `ln(0.0)` is `-inf`. Accumulating in log space keeps it
|
||||||
|
/// finite.
|
||||||
|
#[test]
|
||||||
|
fn log_evidence_survives_a_long_diff_chain() {
|
||||||
|
let holders: Vec<[R; 1]> = (0..1200).map(|_| [rating()]).collect();
|
||||||
|
let teams: Vec<&[R]> = holders.iter().map(|t| t.as_slice()).collect();
|
||||||
|
let game = Game::ranked(
|
||||||
|
&teams,
|
||||||
|
Outcome::ranking(0..holders.len() as u32),
|
||||||
|
&GameOptions::default(),
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let log_evidence = game.log_evidence();
|
||||||
|
assert!(
|
||||||
|
log_evidence.is_finite(),
|
||||||
|
"1200-team log-evidence must be finite, got {log_evidence}"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
log_evidence < 0.0,
|
||||||
|
"log-evidence of a probability must be negative, got {log_evidence}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A near-certain outcome rounds the losing tail to exactly zero in the
|
||||||
|
/// `erfc` approximation; the evidence floor keeps `ln` finite.
|
||||||
|
#[test]
|
||||||
|
fn log_evidence_finite_for_near_certain_outcome() {
|
||||||
|
let overwhelming = R::new(Gaussian::from_ms(5_000.0, 0.5), 1.0, ConstantDrift(0.0));
|
||||||
|
let hopeless = R::new(Gaussian::from_ms(-5_000.0, 0.5), 1.0, ConstantDrift(0.0));
|
||||||
|
let a = [overwhelming];
|
||||||
|
let b = [hopeless];
|
||||||
|
let teams: Vec<&[R]> = vec![&a, &b];
|
||||||
|
|
||||||
|
let game = Game::ranked(&teams, Outcome::winner(0, 2), &GameOptions::default()).unwrap();
|
||||||
|
assert!(
|
||||||
|
game.log_evidence().is_finite(),
|
||||||
|
"got {}",
|
||||||
|
game.log_evidence()
|
||||||
|
);
|
||||||
|
|
||||||
|
// And the reverse — a colossal upset — must also stay finite.
|
||||||
|
let upset = Game::ranked(&teams, Outcome::winner(1, 2), &GameOptions::default()).unwrap();
|
||||||
|
assert!(
|
||||||
|
upset.log_evidence().is_finite(),
|
||||||
|
"upset log-evidence must be finite, got {}",
|
||||||
|
upset.log_evidence()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn empty_history_has_no_filtered_estimates() {
|
||||||
|
let history: History = History::builder().build();
|
||||||
|
|
||||||
|
assert_eq!(history.filtered_log_evidence(), 0.0);
|
||||||
|
|
||||||
|
assert!(history.filtered_learning_curves().is_empty());
|
||||||
|
|
||||||
|
assert!(history.filtered_learning_curve("nobody").is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Boundary inputs (#26) ----------------------------------------------
|
||||||
|
|
||||||
|
fn tight() -> ConvergenceOptions {
|
||||||
|
ConvergenceOptions {
|
||||||
|
max_iter: 2_000,
|
||||||
|
epsilon: 1e-12,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn assert_curve_finite(h: &History, keys: &[&str], what: &str) {
|
||||||
|
for key in keys {
|
||||||
|
for (time, g) in h.learning_curve(*key) {
|
||||||
|
assert!(
|
||||||
|
g.mu().is_finite() && g.sigma().is_finite(),
|
||||||
|
"{what}: non-finite posterior for {key} at t={time} (mu={} sigma={})",
|
||||||
|
g.mu(),
|
||||||
|
g.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A zero weight reaches `(m - performance.exclude(..)) * (1.0 / w)`, i.e. a
|
||||||
|
/// division by zero. The commit is accepted today, so this pins that the
|
||||||
|
/// resulting posterior is still finite rather than quietly NaN.
|
||||||
|
#[test]
|
||||||
|
fn zero_weight_does_not_produce_a_non_finite_posterior() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
|
||||||
|
h.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.weights([0.0])
|
||||||
|
.team(["b"])
|
||||||
|
.winner(0)
|
||||||
|
.commit()
|
||||||
|
.expect("a zero weight is accepted today; update this test if that changes");
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
assert_curve_finite(&h, &["a", "b"], "zero weight");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn negative_weight_does_not_produce_a_non_finite_posterior() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
|
||||||
|
h.event(1)
|
||||||
|
.team(["a"])
|
||||||
|
.weights([-1.0])
|
||||||
|
.team(["b"])
|
||||||
|
.winner(0)
|
||||||
|
.commit()
|
||||||
|
.expect("a negative weight is accepted today; update this test if that changes");
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
assert_curve_finite(&h, &["a", "b"], "negative weight");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Events supplied newest-first must land in the same slices as oldest-first:
|
||||||
|
/// ingestion sorts by time rather than trusting arrival order.
|
||||||
|
#[test]
|
||||||
|
fn out_of_order_timestamps_converge_to_the_same_answer() {
|
||||||
|
fn build(descending: bool) -> History {
|
||||||
|
let mut h = History::builder().convergence(tight()).build();
|
||||||
|
|
||||||
|
let mut times: Vec<i64> = (1..=6).collect();
|
||||||
|
if descending {
|
||||||
|
times.reverse();
|
||||||
|
}
|
||||||
|
|
||||||
|
for time in times {
|
||||||
|
h.record_winner(&"a", &"b", time).unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
h
|
||||||
|
}
|
||||||
|
|
||||||
|
let ascending = build(false);
|
||||||
|
let descending = build(true);
|
||||||
|
|
||||||
|
let one = ascending.current_skill("a").unwrap();
|
||||||
|
let other = descending.current_skill("a").unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(one.mu() - other.mu()).abs() < 1e-8 && (one.sigma() - other.sigma()).abs() < 1e-8,
|
||||||
|
"arrival order changed the answer: ascending mu={} sigma={}, descending mu={} sigma={}",
|
||||||
|
one.mu(),
|
||||||
|
one.sigma(),
|
||||||
|
other.mu(),
|
||||||
|
other.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn extreme_beta_and_sigma_stay_finite() {
|
||||||
|
for (beta, sigma) in [(1e-6, 1e-6), (1e6, 1e6), (1e-6, 1e6), (1e6, 1e-6)] {
|
||||||
|
let mut h = History::builder().beta(beta).sigma(sigma).build();
|
||||||
|
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.record_winner(&"a", &"b", 2).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
assert_curve_finite(&h, &["a", "b"], &format!("beta={beta} sigma={sigma}"));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,499 @@
|
|||||||
|
//! Per-competitor drift scaling via `Member::with_drift_scale`.
|
||||||
|
//!
|
||||||
|
//! The scale multiplies the *variance* the history's `Drift` contributes for
|
||||||
|
//! that competitor, so `scale` is in the same units as `gamma`:
|
||||||
|
//! `ConstantDrift(g)` at `scale = s` behaves as `ConstantDrift(g * s)` would.
|
||||||
|
//! `scale = 0.0` pins a competitor still — an anchor, a rating floor, a course
|
||||||
|
//! difficulty — while everyone around them keeps drifting.
|
||||||
|
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{
|
||||||
|
ConstantDrift, ConvergenceOptions, Event, Gaussian, History, InferenceError, Member,
|
||||||
|
NullObserver, Outcome, Team,
|
||||||
|
};
|
||||||
|
|
||||||
|
type Fit = History<i64, ConstantDrift, NullObserver, &'static str>;
|
||||||
|
|
||||||
|
const CONVERGENCE: ConvergenceOptions = ConvergenceOptions {
|
||||||
|
max_iter: 64,
|
||||||
|
epsilon: 1e-9,
|
||||||
|
alpha: 1.0,
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Two events separated by a long gap, so drift has room to matter.
|
||||||
|
fn distant_pair(anchor_scale: Option<f64>) -> Vec<Event<i64, &'static str>> {
|
||||||
|
let anchor = |s: Option<f64>| match s {
|
||||||
|
Some(scale) => Member::new("anchor").with_drift_scale(scale),
|
||||||
|
None => Member::new("anchor"),
|
||||||
|
};
|
||||||
|
|
||||||
|
vec![
|
||||||
|
Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([anchor(anchor_scale)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 1000,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([anchor(anchor_scale)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(1, 2),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
fn fit(events: Vec<Event<i64, &'static str>>, gamma: f64) -> Fit {
|
||||||
|
let mut h = History::builder()
|
||||||
|
.mu(25.0)
|
||||||
|
.sigma(25.0 / 3.0)
|
||||||
|
.beta(25.0 / 6.0)
|
||||||
|
.p_draw(0.0)
|
||||||
|
.drift(ConstantDrift(gamma))
|
||||||
|
.convergence(CONVERGENCE)
|
||||||
|
.build();
|
||||||
|
|
||||||
|
h.add_events(events).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
h
|
||||||
|
}
|
||||||
|
|
||||||
|
fn curve(h: &Fit, key: &str) -> Vec<(i64, Gaussian)> {
|
||||||
|
let mut c = h.learning_curves().remove(key).expect("key in curves");
|
||||||
|
c.sort_by_key(|(t, _)| *t);
|
||||||
|
c
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A competitor at `scale = 0.0` is one latent skill observed twice, so the
|
||||||
|
/// posterior is the same distribution at both times — and strictly tighter
|
||||||
|
/// than the same competitor left to drift.
|
||||||
|
#[test]
|
||||||
|
fn zero_scale_pins_a_competitor_still() {
|
||||||
|
let pinned = fit(distant_pair(Some(0.0)), 25.0 / 300.0);
|
||||||
|
let drifting = fit(distant_pair(None), 25.0 / 300.0);
|
||||||
|
|
||||||
|
let pinned_curve = curve(&pinned, "anchor");
|
||||||
|
assert_eq!(pinned_curve.len(), 2);
|
||||||
|
|
||||||
|
let (t0, first) = pinned_curve[0];
|
||||||
|
let (t1, second) = pinned_curve[1];
|
||||||
|
assert_eq!((t0, t1), (0, 1000));
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(first.sigma() - second.sigma()).abs() < 1e-9,
|
||||||
|
"a pinned competitor's uncertainty must not move between t=0 and t=1000: \
|
||||||
|
{} vs {}",
|
||||||
|
first.sigma(),
|
||||||
|
second.sigma()
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(first.mu() - second.mu()).abs() < 1e-9,
|
||||||
|
"a pinned competitor's mean must not move: {} vs {}",
|
||||||
|
first.mu(),
|
||||||
|
second.mu()
|
||||||
|
);
|
||||||
|
|
||||||
|
let drifting_curve = curve(&drifting, "anchor");
|
||||||
|
assert!(
|
||||||
|
drifting_curve[0].1.sigma() > first.sigma() + 1e-6,
|
||||||
|
"drift must leave the anchor less certain than pinning does: {} vs {}",
|
||||||
|
drifting_curve[0].1.sigma(),
|
||||||
|
first.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The scale is composable with `gamma`: scaling every competitor by `s` is
|
||||||
|
/// exactly the same fit as scaling the history's drift by `s`.
|
||||||
|
#[test]
|
||||||
|
fn scale_is_equivalent_to_scaling_gamma() {
|
||||||
|
let scaled: Vec<Event<i64, &'static str>> = vec![
|
||||||
|
Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a").with_drift_scale(0.5)]),
|
||||||
|
Team::with_members([Member::new("b").with_drift_scale(0.5)]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 400,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("b").with_drift_scale(0.5)]),
|
||||||
|
Team::with_members([Member::new("a").with_drift_scale(0.5)]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
];
|
||||||
|
|
||||||
|
let plain: Vec<Event<i64, &'static str>> = vec![
|
||||||
|
Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 400,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
];
|
||||||
|
|
||||||
|
let by_scale = fit(scaled, 0.3);
|
||||||
|
let by_gamma = fit(plain, 0.15);
|
||||||
|
|
||||||
|
for key in ["a", "b"] {
|
||||||
|
let lhs = curve(&by_scale, key);
|
||||||
|
let rhs = curve(&by_gamma, key);
|
||||||
|
assert_eq!(lhs.len(), rhs.len());
|
||||||
|
|
||||||
|
for ((t_l, g_l), (t_r, g_r)) in lhs.iter().zip(rhs.iter()) {
|
||||||
|
assert_eq!(t_l, t_r);
|
||||||
|
assert!(
|
||||||
|
(g_l.mu() - g_r.mu()).abs() < 1e-9 && (g_l.sigma() - g_r.sigma()).abs() < 1e-9,
|
||||||
|
"ConstantDrift(0.3) at scale 0.5 must equal ConstantDrift(0.15) for {key} at \
|
||||||
|
t={t_l}: ({}, {}) vs ({}, {})",
|
||||||
|
g_l.mu(),
|
||||||
|
g_l.sigma(),
|
||||||
|
g_r.mu(),
|
||||||
|
g_r.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `None` means 1.0: an explicit unit scale changes nothing.
|
||||||
|
#[test]
|
||||||
|
fn unset_scale_matches_an_explicit_unit_scale() {
|
||||||
|
let implicit = fit(distant_pair(None), 25.0 / 300.0);
|
||||||
|
let explicit = fit(distant_pair(Some(1.0)), 25.0 / 300.0);
|
||||||
|
|
||||||
|
for key in ["anchor", "player"] {
|
||||||
|
let lhs = curve(&implicit, key);
|
||||||
|
let rhs = curve(&explicit, key);
|
||||||
|
assert_eq!(lhs.len(), rhs.len());
|
||||||
|
|
||||||
|
for ((t_l, g_l), (t_r, g_r)) in lhs.iter().zip(rhs.iter()) {
|
||||||
|
assert_eq!(t_l, t_r);
|
||||||
|
assert_eq!(
|
||||||
|
(g_l.mu(), g_l.sigma()),
|
||||||
|
(g_r.mu(), g_r.sigma()),
|
||||||
|
"an explicit scale of 1.0 must be bit-identical to leaving it unset, \
|
||||||
|
for {key} at t={t_l}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The use case from the issue: a static difficulty alongside drifting players,
|
||||||
|
/// in one graph. The anchor must hold still without absorbing drift through its
|
||||||
|
/// neighbours, and everything must stay finite.
|
||||||
|
#[test]
|
||||||
|
fn mixed_static_and_drifting_graph_converges() {
|
||||||
|
let mut events: Vec<Event<i64, &'static str>> = Vec::new();
|
||||||
|
let players = ["p0", "p1", "p2"];
|
||||||
|
|
||||||
|
for (i, p) in players.iter().cycle().take(9).enumerate() {
|
||||||
|
events.push(Event {
|
||||||
|
time: (i as i64) * 100,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(*p)]),
|
||||||
|
Team::with_members([Member::new("layout").with_drift_scale(0.0)]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner((i % 2) as u32, 2),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut h = History::builder()
|
||||||
|
.mu(25.0)
|
||||||
|
.sigma(25.0 / 3.0)
|
||||||
|
.beta(25.0 / 6.0)
|
||||||
|
.p_draw(0.0)
|
||||||
|
.drift(ConstantDrift(25.0 / 300.0))
|
||||||
|
.convergence(CONVERGENCE)
|
||||||
|
.build();
|
||||||
|
|
||||||
|
h.add_events(events).unwrap();
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
assert!(report.converged, "mixed graph must converge: {report:?}");
|
||||||
|
|
||||||
|
let curves = h.learning_curves();
|
||||||
|
for (key, points) in &curves {
|
||||||
|
for (t, g) in points {
|
||||||
|
assert!(
|
||||||
|
g.mu().is_finite() && g.sigma().is_finite() && g.sigma() > 0.0,
|
||||||
|
"{key} at t={t} is not a usable posterior: mu={}, sigma={}",
|
||||||
|
g.mu(),
|
||||||
|
g.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let layout = curve(&h, "layout");
|
||||||
|
assert_eq!(layout.len(), 9);
|
||||||
|
let (_, first) = layout[0];
|
||||||
|
for (t, g) in &layout {
|
||||||
|
assert!(
|
||||||
|
(g.sigma() - first.sigma()).abs() < 1e-9,
|
||||||
|
"a static layout must not accumulate uncertainty; t={t} has sigma {} vs {}",
|
||||||
|
g.sigma(),
|
||||||
|
first.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
let p0 = curve(&h, "p0");
|
||||||
|
assert!(
|
||||||
|
p0.last().unwrap().1.sigma() > 0.0,
|
||||||
|
"a drifting player should still have a proper posterior"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn reject(scale: f64) -> InferenceError {
|
||||||
|
let mut h = History::builder()
|
||||||
|
.drift(ConstantDrift(25.0 / 300.0))
|
||||||
|
.build();
|
||||||
|
|
||||||
|
let events: Vec<Event<i64, &'static str>> = vec![Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a").with_drift_scale(scale)]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}];
|
||||||
|
|
||||||
|
h.add_events(events)
|
||||||
|
.expect_err("an out-of-range drift_scale must be rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn negative_scale_is_rejected() {
|
||||||
|
assert_eq!(
|
||||||
|
reject(-1.0),
|
||||||
|
InferenceError::InvalidParameter {
|
||||||
|
name: "drift_scale",
|
||||||
|
value: -1.0
|
||||||
|
}
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn non_finite_scale_is_rejected() {
|
||||||
|
for scale in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
|
||||||
|
assert!(
|
||||||
|
matches!(
|
||||||
|
reject(scale),
|
||||||
|
InferenceError::InvalidParameter {
|
||||||
|
name: "drift_scale",
|
||||||
|
..
|
||||||
|
}
|
||||||
|
),
|
||||||
|
"a drift_scale of {scale} must be rejected as an invalid parameter"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The scale must reach the filtering pass too, not just `converge()`.
|
||||||
|
/// `filtered_learning_curves` runs its own drift application, so a pinned
|
||||||
|
/// competitor has to stay pinned there as well.
|
||||||
|
#[test]
|
||||||
|
fn zero_scale_pins_a_competitor_in_the_filtered_pass() {
|
||||||
|
let pinned = fit(distant_pair(Some(0.0)), 25.0 / 300.0);
|
||||||
|
let drifting = fit(distant_pair(None), 25.0 / 300.0);
|
||||||
|
|
||||||
|
let filtered = |h: &Fit| -> Vec<(i64, Gaussian)> {
|
||||||
|
let mut c = h
|
||||||
|
.filtered_learning_curves()
|
||||||
|
.remove("anchor")
|
||||||
|
.expect("anchor in filtered curves");
|
||||||
|
c.sort_by_key(|(t, _)| *t);
|
||||||
|
c
|
||||||
|
};
|
||||||
|
|
||||||
|
let pinned_curve = filtered(&pinned);
|
||||||
|
let drifting_curve = filtered(&drifting);
|
||||||
|
assert_eq!(pinned_curve.len(), 2);
|
||||||
|
assert_eq!(drifting_curve.len(), 2);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
pinned_curve[1].1.sigma() < pinned_curve[0].1.sigma(),
|
||||||
|
"a pinned competitor's filtered uncertainty must shrink with a second \
|
||||||
|
observation, not be re-inflated by drift: {} then {}",
|
||||||
|
pinned_curve[0].1.sigma(),
|
||||||
|
pinned_curve[1].1.sigma()
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
pinned_curve[1].1.sigma() < drifting_curve[1].1.sigma() - 1e-6,
|
||||||
|
"pinning must leave the filtered estimate tighter than drifting does: \
|
||||||
|
{} vs {}",
|
||||||
|
pinned_curve[1].1.sigma(),
|
||||||
|
drifting_curve[1].1.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `drift_scale` is competitor configuration, and configuration supplied for a
|
||||||
|
/// competitor the history already knows is now *applied* rather than dropped.
|
||||||
|
///
|
||||||
|
/// This test previously asserted the opposite. It was written as a deliberate
|
||||||
|
/// change-detector — "moving the capture would be a visible break, not a silent
|
||||||
|
/// one" — and that is exactly what happened: the capture moved, and the
|
||||||
|
/// assertion inverted rather than being deleted.
|
||||||
|
///
|
||||||
|
/// Because configuration lives on the competitor and `converge` refits from
|
||||||
|
/// competitor state, a late pin applies to the *whole* history, not just to
|
||||||
|
/// events after it. So a scale set on the second batch must reach the same fit
|
||||||
|
/// as one set from the very first event.
|
||||||
|
#[test]
|
||||||
|
fn drift_scale_applies_when_set_after_first_appearance() {
|
||||||
|
let mut late = History::builder()
|
||||||
|
.mu(25.0)
|
||||||
|
.sigma(25.0 / 3.0)
|
||||||
|
.beta(25.0 / 6.0)
|
||||||
|
.p_draw(0.0)
|
||||||
|
.drift(ConstantDrift(25.0 / 300.0))
|
||||||
|
.convergence(CONVERGENCE)
|
||||||
|
.build();
|
||||||
|
|
||||||
|
// First batch creates "anchor" with the default scale.
|
||||||
|
late.add_events(vec![Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("anchor")]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}])
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
// Second batch asks for a pin. No longer too late.
|
||||||
|
late.add_events(vec![Event {
|
||||||
|
time: 1000,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("anchor").with_drift_scale(0.0)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(1, 2),
|
||||||
|
}])
|
||||||
|
.unwrap();
|
||||||
|
late.converge().unwrap();
|
||||||
|
|
||||||
|
let applied = curve(&late, "anchor");
|
||||||
|
let pinned_from_the_start = curve(&fit(distant_pair(Some(0.0)), 25.0 / 300.0), "anchor");
|
||||||
|
let never_pinned = curve(&fit(distant_pair(None), 25.0 / 300.0), "anchor");
|
||||||
|
|
||||||
|
for ((t_l, g_l), (t_r, g_r)) in applied.iter().zip(pinned_from_the_start.iter()) {
|
||||||
|
assert_eq!(t_l, t_r);
|
||||||
|
assert!(
|
||||||
|
(g_l.sigma() - g_r.sigma()).abs() < 1e-9,
|
||||||
|
"a late pin should refit the whole history: t={t_l}, {} vs {}",
|
||||||
|
g_l.sigma(),
|
||||||
|
g_r.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// And it must actually have done something.
|
||||||
|
assert!(
|
||||||
|
applied
|
||||||
|
.iter()
|
||||||
|
.zip(never_pinned.iter())
|
||||||
|
.any(|((_, a), (_, b))| (a.sigma() - b.sigma()).abs() > 1e-9),
|
||||||
|
"the pin had no effect at all — the silent drop is back"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Re-declaring the same configuration must be inert. This is the shape a
|
||||||
|
/// caller gets when the configuration is a property of the domain — "layouts
|
||||||
|
/// are static" — so every ingestion path repeats it on every event.
|
||||||
|
///
|
||||||
|
/// Both histories see exactly the same events; only how many times the scale
|
||||||
|
/// is declared differs.
|
||||||
|
#[test]
|
||||||
|
fn repeating_the_same_configuration_changes_nothing() {
|
||||||
|
let events = |declare_every_time: bool| {
|
||||||
|
let anchor = |first: bool| {
|
||||||
|
if first || declare_every_time {
|
||||||
|
Member::new("anchor").with_drift_scale(0.0)
|
||||||
|
} else {
|
||||||
|
Member::new("anchor")
|
||||||
|
}
|
||||||
|
};
|
||||||
|
vec![
|
||||||
|
Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([anchor(true)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 1000,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([anchor(false)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(1, 2),
|
||||||
|
},
|
||||||
|
]
|
||||||
|
};
|
||||||
|
|
||||||
|
let once = curve(&fit(events(false), 25.0 / 300.0), "anchor");
|
||||||
|
let every_time = curve(&fit(events(true), 25.0 / 300.0), "anchor");
|
||||||
|
|
||||||
|
for ((t_l, a), (t_r, b)) in once.iter().zip(every_time.iter()) {
|
||||||
|
assert_eq!(t_l, t_r);
|
||||||
|
assert!(
|
||||||
|
(a.sigma() - b.sigma()).abs() < 1e-12,
|
||||||
|
"t={t_l}: declaring the same scale repeatedly changed the fit, {} vs {}",
|
||||||
|
a.sigma(),
|
||||||
|
b.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn a_batch_that_contradicts_itself_is_rejected() {
|
||||||
|
let mut h = History::builder().convergence(CONVERGENCE).build();
|
||||||
|
|
||||||
|
let err = h
|
||||||
|
.add_events(vec![
|
||||||
|
Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("anchor").with_drift_scale(0.0)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 1,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("anchor").with_drift_scale(1.0)]),
|
||||||
|
Team::with_members([Member::new("player")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
])
|
||||||
|
.expect_err("two different scales for one competitor in one batch");
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::ConflictingCompetitorConfig {
|
||||||
|
field: "drift_scale",
|
||||||
|
..
|
||||||
|
}
|
||||||
|
),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
+7
-12
@@ -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,
|
||||||
@@ -48,15 +49,9 @@ fn game_1v1_draw_golden() {
|
|||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
let p = g.posteriors();
|
let p = g.posteriors();
|
||||||
// Historical golden from pre-T2 test_1vs1_draw:
|
// Historical golden from pre-T2 test_1vs1_draw. The mean is 25.0 exactly
|
||||||
assert_ulps_eq!(
|
// by symmetry — two identical competitors drawing cannot move apart — and
|
||||||
p[0][0],
|
// the reference's 24.999999 is that value transcribed to six decimals.
|
||||||
Gaussian::from_ms(24.999999, 6.469480),
|
assert_ulps_eq!(p[0][0], Gaussian::from_ms(25.0, 6.469480), epsilon = 1e-6);
|
||||||
epsilon = 1e-6
|
assert_ulps_eq!(p[1][0], Gaussian::from_ms(25.0, 6.469480), epsilon = 1e-6);
|
||||||
);
|
|
||||||
assert_ulps_eq!(
|
|
||||||
p[1][0],
|
|
||||||
Gaussian::from_ms(24.999999, 6.469480),
|
|
||||||
epsilon = 1e-6
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,254 @@
|
|||||||
|
//! Forward-only (filtering) estimates: what the model knew at the time,
|
||||||
|
//! as opposed to the smoothed posteriors `learning_curve` reports.
|
||||||
|
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{ConvergenceOptions, Event, History, Member, Outcome, Team};
|
||||||
|
|
||||||
|
/// `games` one-on-one matches at successive times, won by "a" every time,
|
||||||
|
/// built with the given convergence options.
|
||||||
|
fn repeated_winner_with(games: i64, convergence: ConvergenceOptions) -> History {
|
||||||
|
let mut history = History::builder().convergence(convergence).build();
|
||||||
|
|
||||||
|
for time in 1..=games {
|
||||||
|
history
|
||||||
|
.add_events([Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}])
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
history
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `games` one-on-one matches at successive times, won by "a" every time.
|
||||||
|
///
|
||||||
|
/// This is the fixture from issue #19, where `online(true)` reported
|
||||||
|
/// `games * ln(0.5)`.
|
||||||
|
fn repeated_winner(games: i64) -> History {
|
||||||
|
repeated_winner_with(games, ConvergenceOptions::default())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The default 30-iteration cap leaves a residual around 1e-6, which would
|
||||||
|
/// swamp these comparisons. Drive both sides well past the fixed point.
|
||||||
|
fn tight() -> ConvergenceOptions {
|
||||||
|
ConvergenceOptions {
|
||||||
|
max_iter: 2_000,
|
||||||
|
epsilon: 1e-12,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn filtered_evidence_sits_between_coin_flip_and_batch() {
|
||||||
|
let mut history = repeated_winner(5);
|
||||||
|
|
||||||
|
history.converge().unwrap();
|
||||||
|
|
||||||
|
let coin_flip = 5.0 * 0.5f64.ln();
|
||||||
|
let batch = history.log_evidence();
|
||||||
|
let filtered = history.filtered_log_evidence();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
filtered > coin_flip,
|
||||||
|
"filtered evidence {filtered} is at or below {coin_flip}, the all-coin-flip \
|
||||||
|
value the inert online flag reported; game one is a coin flip but games two \
|
||||||
|
through five are not"
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
filtered < batch,
|
||||||
|
"filtered evidence {filtered} is not below the smoothed {batch}; filtering \
|
||||||
|
scores each game on strictly less information than smoothing does"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn filtered_first_point_is_less_certain_than_smoothed() {
|
||||||
|
let mut history = repeated_winner(12);
|
||||||
|
|
||||||
|
history.converge().unwrap();
|
||||||
|
|
||||||
|
let smoothed = history.learning_curve("a");
|
||||||
|
let filtered = history.filtered_learning_curve("a");
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
smoothed.len(),
|
||||||
|
filtered.len(),
|
||||||
|
"both curves must cover the same time points"
|
||||||
|
);
|
||||||
|
|
||||||
|
let (smoothed_time, first_smoothed) = smoothed[0];
|
||||||
|
let (filtered_time, first_filtered) = filtered[0];
|
||||||
|
|
||||||
|
assert_eq!(smoothed_time, filtered_time);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
first_filtered.sigma() > first_smoothed.sigma(),
|
||||||
|
"filtered sigma {} at the first point is not above smoothed {}; the smoother \
|
||||||
|
collapses uncertainty before the first round is drawn, which is the whole \
|
||||||
|
reason this method exists",
|
||||||
|
first_filtered.sigma(),
|
||||||
|
first_smoothed.sigma()
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
first_filtered.sigma() < trueskill_tt::SIGMA,
|
||||||
|
"filtered sigma {} at the first point is not below the prior {}; one game was \
|
||||||
|
played, so some uncertainty must have been resolved",
|
||||||
|
first_filtered.sigma(),
|
||||||
|
trueskill_tt::SIGMA
|
||||||
|
);
|
||||||
|
|
||||||
|
for pair in filtered.windows(2) {
|
||||||
|
assert!(
|
||||||
|
pair[1].1.mu() > pair[0].1.mu(),
|
||||||
|
"filtered mu must climb at every step for a competitor who wins every \
|
||||||
|
game: t={} mu={} then t={} mu={}",
|
||||||
|
pair[0].0,
|
||||||
|
pair[0].1.mu(),
|
||||||
|
pair[1].0,
|
||||||
|
pair[1].1.mu()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn filtered_curves_plural_agrees_with_singular() {
|
||||||
|
let mut history = repeated_winner(4);
|
||||||
|
|
||||||
|
history.converge().unwrap();
|
||||||
|
|
||||||
|
let curves = history.filtered_learning_curves();
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
curves["b"],
|
||||||
|
history.filtered_learning_curve("b"),
|
||||||
|
"the plural form must agree with the singular for the same key"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn filtered_evidence_is_invariant_to_convergence() {
|
||||||
|
let mut history = repeated_winner_with(6, tight());
|
||||||
|
|
||||||
|
let before = history.filtered_log_evidence();
|
||||||
|
|
||||||
|
let report = history.converge().unwrap();
|
||||||
|
assert!(
|
||||||
|
report.converged,
|
||||||
|
"fixture must converge: {:?}",
|
||||||
|
report.final_step
|
||||||
|
);
|
||||||
|
|
||||||
|
let after = history.filtered_log_evidence();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(before - after).abs() < 1e-8,
|
||||||
|
"filtered evidence moved across converge(): {before} -> {after}. The pass must \
|
||||||
|
carry its own forward messages; anything reading skill.forward shows exactly \
|
||||||
|
this drift, because converge() contaminates it with backward information."
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn single_slice_filtered_matches_smoothed() {
|
||||||
|
let mut history = History::builder().convergence(tight()).build();
|
||||||
|
|
||||||
|
history
|
||||||
|
.add_events([
|
||||||
|
Event {
|
||||||
|
time: 1,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
Event {
|
||||||
|
time: 1,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("c")]),
|
||||||
|
Team::with_members([Member::new("d")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
},
|
||||||
|
])
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
history.converge().unwrap();
|
||||||
|
|
||||||
|
let smoothed = history.learning_curve("a");
|
||||||
|
let filtered = history.filtered_learning_curve("a");
|
||||||
|
|
||||||
|
assert_eq!(smoothed.len(), 1);
|
||||||
|
assert_eq!(filtered.len(), 1);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(smoothed[0].1.mu() - filtered[0].1.mu()).abs() < 1e-8
|
||||||
|
&& (smoothed[0].1.sigma() - filtered[0].1.sigma()).abs() < 1e-8,
|
||||||
|
"one slice has no future to propagate back, so filtered and smoothed must \
|
||||||
|
agree: smoothed mu={} sigma={}, filtered mu={} sigma={}",
|
||||||
|
smoothed[0].1.mu(),
|
||||||
|
smoothed[0].1.sigma(),
|
||||||
|
filtered[0].1.mu(),
|
||||||
|
filtered[0].1.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn filtered_curves_do_not_depend_on_ingestion_order() {
|
||||||
|
let events = |time: i64, winner: &'static str, loser: &'static str| Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(winner)]),
|
||||||
|
Team::with_members([Member::new(loser)]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
};
|
||||||
|
|
||||||
|
let all = vec![
|
||||||
|
events(1, "a", "b"),
|
||||||
|
events(1, "c", "d"),
|
||||||
|
events(1, "a", "c"),
|
||||||
|
events(1, "b", "d"),
|
||||||
|
events(2, "a", "d"),
|
||||||
|
events(2, "b", "c"),
|
||||||
|
events(2, "a", "b"),
|
||||||
|
];
|
||||||
|
|
||||||
|
let mut batched = History::builder().convergence(tight()).build();
|
||||||
|
batched.add_events(all.clone()).unwrap();
|
||||||
|
batched.converge().unwrap();
|
||||||
|
|
||||||
|
let mut incremental = History::builder().convergence(tight()).build();
|
||||||
|
for event in all {
|
||||||
|
incremental.add_events([event]).unwrap();
|
||||||
|
}
|
||||||
|
incremental.converge().unwrap();
|
||||||
|
|
||||||
|
let from_batched = batched.filtered_learning_curve("a");
|
||||||
|
let from_incremental = incremental.filtered_learning_curve("a");
|
||||||
|
|
||||||
|
assert_eq!(from_batched.len(), from_incremental.len());
|
||||||
|
|
||||||
|
for ((time_b, gaussian_b), (time_i, gaussian_i)) in
|
||||||
|
from_batched.iter().zip(from_incremental.iter())
|
||||||
|
{
|
||||||
|
assert_eq!(time_b, time_i);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(gaussian_b.mu() - gaussian_i.mu()).abs() < 1e-8
|
||||||
|
&& (gaussian_b.sigma() - gaussian_i.sigma()).abs() < 1e-8,
|
||||||
|
"at t={time_b}: batched mu={} sigma={}, incremental mu={} sigma={}",
|
||||||
|
gaussian_b.mu(),
|
||||||
|
gaussian_b.sigma(),
|
||||||
|
gaussian_i.mu(),
|
||||||
|
gaussian_i.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
+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,225 @@
|
|||||||
|
//! Ingesting the same events must give the same answer however they were
|
||||||
|
//! batched.
|
||||||
|
//!
|
||||||
|
//! The numerical goldens all ingest in a single call with one slice per
|
||||||
|
//! timestamp, so they never exercise the "append to an existing slice" path.
|
||||||
|
//! These do.
|
||||||
|
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{ConvergenceOptions, Event, Gaussian, History, Member, Outcome, Team};
|
||||||
|
|
||||||
|
/// Converge tightly: the default cap of 30 iterations leaves a residual around
|
||||||
|
/// 1e-6, which would swamp the comparison. Both paths must reach the same
|
||||||
|
/// fixed point, so drive both well past it.
|
||||||
|
fn tight() -> ConvergenceOptions {
|
||||||
|
ConvergenceOptions {
|
||||||
|
max_iter: 2_000,
|
||||||
|
epsilon: 1e-12,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn event(a: &str, b: &str, time: i64) -> Event<i64, String> {
|
||||||
|
Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(a.to_string())]),
|
||||||
|
Team::with_members([Member::new(b.to_string())]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Like [`event`], but `a` carries competitor configuration.
|
||||||
|
///
|
||||||
|
/// `prior` and `drift_scale` configure the competitor rather than the event, so
|
||||||
|
/// they are the part of ingestion most exposed to order: they are consumed once,
|
||||||
|
/// where the competitor's state is written.
|
||||||
|
fn configured_event(a: &str, b: &str, time: i64, scale: f64) -> Event<i64, String> {
|
||||||
|
Event {
|
||||||
|
time,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(a.to_string()).with_drift_scale(scale)]),
|
||||||
|
Team::with_members([Member::new(b.to_string())]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn converged_skills(events: Vec<Event<i64, String>>, batched: bool) -> Vec<(String, Gaussian)> {
|
||||||
|
let mut h: History<i64, _, _, String> =
|
||||||
|
History::builder_with_key().convergence(tight()).build();
|
||||||
|
|
||||||
|
if batched {
|
||||||
|
h.add_events(events).unwrap();
|
||||||
|
} else {
|
||||||
|
for ev in events {
|
||||||
|
h.add_events(std::iter::once(ev)).unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let report = h.converge().unwrap();
|
||||||
|
assert!(
|
||||||
|
report.converged,
|
||||||
|
"fixture must converge before results can be compared; final step {:?}",
|
||||||
|
report.final_step
|
||||||
|
);
|
||||||
|
|
||||||
|
let mut skills: Vec<(String, Gaussian)> = h
|
||||||
|
.learning_curves()
|
||||||
|
.into_iter()
|
||||||
|
.map(|(key, curve)| (key, curve.last().unwrap().1))
|
||||||
|
.collect();
|
||||||
|
skills.sort_by(|a, b| a.0.cmp(&b.0));
|
||||||
|
skills
|
||||||
|
}
|
||||||
|
|
||||||
|
fn assert_same(batched: &[(String, Gaussian)], incremental: &[(String, Gaussian)], what: &str) {
|
||||||
|
assert_eq!(
|
||||||
|
batched.len(),
|
||||||
|
incremental.len(),
|
||||||
|
"{what}: competitor count differs"
|
||||||
|
);
|
||||||
|
|
||||||
|
for ((kb, gb), (ki, gi)) in batched.iter().zip(incremental.iter()) {
|
||||||
|
assert_eq!(kb, ki, "{what}: key order differs");
|
||||||
|
assert!(
|
||||||
|
(gb.mu() - gi.mu()).abs() < 1e-8 && (gb.sigma() - gi.sigma()).abs() < 1e-8,
|
||||||
|
"{what}: {kb} differs — batched mu={} sigma={}, incremental mu={} sigma={}",
|
||||||
|
gb.mu(),
|
||||||
|
gb.sigma(),
|
||||||
|
gi.mu(),
|
||||||
|
gi.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// All events share one timestamp, so incremental ingestion repeatedly appends
|
||||||
|
/// to an existing slice.
|
||||||
|
#[test]
|
||||||
|
fn same_slice_incremental_matches_batched() {
|
||||||
|
let events = vec![
|
||||||
|
event("a", "b", 1),
|
||||||
|
event("c", "d", 1),
|
||||||
|
event("e", "f", 1),
|
||||||
|
event("a", "c", 1),
|
||||||
|
event("b", "e", 1),
|
||||||
|
];
|
||||||
|
|
||||||
|
let batched = converged_skills(events.clone(), true);
|
||||||
|
let incremental = converged_skills(events, false);
|
||||||
|
assert_same(&batched, &incremental, "single shared slice");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Distinct timestamps, so each append lands in a fresh slice appended after
|
||||||
|
/// the existing ones.
|
||||||
|
#[test]
|
||||||
|
fn distinct_slices_incremental_matches_batched() {
|
||||||
|
let events = vec![
|
||||||
|
event("a", "b", 1),
|
||||||
|
event("b", "c", 2),
|
||||||
|
event("c", "a", 3),
|
||||||
|
event("a", "c", 4),
|
||||||
|
];
|
||||||
|
|
||||||
|
let batched = converged_skills(events.clone(), true);
|
||||||
|
let incremental = converged_skills(events, false);
|
||||||
|
assert_same(&batched, &incremental, "distinct slices");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Several events per timestamp across several timestamps — appends to
|
||||||
|
/// existing slices interleaved with new ones.
|
||||||
|
#[test]
|
||||||
|
fn mixed_slices_incremental_matches_batched() {
|
||||||
|
let events = vec![
|
||||||
|
event("a", "b", 1),
|
||||||
|
event("c", "d", 1),
|
||||||
|
event("a", "c", 2),
|
||||||
|
event("b", "d", 2),
|
||||||
|
event("a", "d", 3),
|
||||||
|
event("b", "c", 3),
|
||||||
|
];
|
||||||
|
|
||||||
|
let batched = converged_skills(events.clone(), true);
|
||||||
|
let incremental = converged_skills(events, false);
|
||||||
|
assert_same(&batched, &incremental, "mixed slices");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Appending an event to a slice that is *not* the most recent one exercises
|
||||||
|
/// the forward refresh of every later slice.
|
||||||
|
#[test]
|
||||||
|
fn back_dated_event_matches_batched() {
|
||||||
|
let events = vec![
|
||||||
|
event("a", "b", 1),
|
||||||
|
event("b", "c", 5),
|
||||||
|
event("c", "a", 9),
|
||||||
|
// arrives last, but belongs to the middle slice
|
||||||
|
event("a", "c", 5),
|
||||||
|
];
|
||||||
|
|
||||||
|
let batched = converged_skills(events.clone(), true);
|
||||||
|
let incremental = converged_skills(events, false);
|
||||||
|
assert_same(&batched, &incremental, "back-dated event");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The invariant this file protects was only ever checked for *unconfigured*
|
||||||
|
/// competitors — every helper above built members with `Member::new`.
|
||||||
|
///
|
||||||
|
/// Configuration is the part most exposed to ordering, because it is consumed
|
||||||
|
/// once at the point the competitor's state is written rather than replayed per
|
||||||
|
/// event. These cover it.
|
||||||
|
#[test]
|
||||||
|
fn configured_competitors_are_order_independent() {
|
||||||
|
let events = vec![
|
||||||
|
configured_event("a", "b", 0, 0.0),
|
||||||
|
configured_event("a", "c", 1, 0.0),
|
||||||
|
configured_event("a", "b", 2, 0.0),
|
||||||
|
event("b", "c", 3),
|
||||||
|
];
|
||||||
|
|
||||||
|
assert_same(
|
||||||
|
&converged_skills(events.clone(), true),
|
||||||
|
&converged_skills(events, false),
|
||||||
|
"configuration repeated on every appearance",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Configuration supplied only on a *later* event is the case that used to be
|
||||||
|
/// silently dropped. It must now reach the same fit either way it is ingested.
|
||||||
|
#[test]
|
||||||
|
fn late_configuration_is_order_independent() {
|
||||||
|
let events = vec![
|
||||||
|
event("a", "b", 0),
|
||||||
|
configured_event("a", "c", 1, 0.0),
|
||||||
|
event("a", "b", 2),
|
||||||
|
];
|
||||||
|
|
||||||
|
assert_same(
|
||||||
|
&converged_skills(events.clone(), true),
|
||||||
|
&converged_skills(events, false),
|
||||||
|
"configuration supplied after first appearance",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// And it must actually be doing something — an implementation that dropped
|
||||||
|
/// configuration entirely would pass both tests above.
|
||||||
|
#[test]
|
||||||
|
fn configuration_changes_the_fit_however_it_is_ingested() {
|
||||||
|
let configured = vec![
|
||||||
|
event("a", "b", 0),
|
||||||
|
configured_event("a", "c", 1, 0.0),
|
||||||
|
event("a", "b", 2),
|
||||||
|
];
|
||||||
|
let plain = vec![event("a", "b", 0), event("a", "c", 1), event("a", "b", 2)];
|
||||||
|
|
||||||
|
for batched in [true, false] {
|
||||||
|
let with = converged_skills(configured.clone(), batched);
|
||||||
|
let without = converged_skills(plain.clone(), batched);
|
||||||
|
assert!(
|
||||||
|
with.iter()
|
||||||
|
.zip(&without)
|
||||||
|
.any(|((_, x), (_, y))| (x.sigma() - y.sigma()).abs() > 1e-9),
|
||||||
|
"batched={batched}: configuration had no effect, so the order tests are vacuous"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,161 @@
|
|||||||
|
//! `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};
|
||||||
|
|
||||||
|
/// Plain fields. `Arc<O>` implements `Observer`, so the caller shares the
|
||||||
|
/// observer itself rather than wrapping each field in its own `Arc`.
|
||||||
|
#[derive(Default)]
|
||||||
|
struct Recorder {
|
||||||
|
iterations: Mutex<Vec<usize>>,
|
||||||
|
slices: Mutex<Vec<(i64, usize, usize)>>,
|
||||||
|
converged: 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 = Arc::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(Arc::clone(&recorder)).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 = Arc::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(Arc::clone(&recorder)).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 = Arc::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(Arc::clone(&recorder)).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));
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The gap #40 closed: without `impl Observer for Arc<O>`, an observer that
|
||||||
|
/// accumulates anything had to wrap every field in its own `Arc` and derive
|
||||||
|
/// `Clone`, because `History` consumes the observer and never hands it back.
|
||||||
|
#[test]
|
||||||
|
fn a_shared_observer_reaches_the_callers_handle() {
|
||||||
|
let recorder = Arc::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(Arc::clone(&recorder)).build();
|
||||||
|
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
assert!(!recorder.iterations.lock().unwrap().is_empty());
|
||||||
|
assert!(!recorder.slices.lock().unwrap().is_empty());
|
||||||
|
assert!(!recorder.converged.lock().unwrap().is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `?Sized` on the blanket impls means the observer can be chosen at runtime.
|
||||||
|
#[test]
|
||||||
|
fn a_trait_object_observer_works() {
|
||||||
|
let boxed: Box<dyn Observer<i64>> = Box::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(boxed).build();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let shared: Arc<dyn Observer<i64>> = Arc::new(Recorder::default());
|
||||||
|
let mut h = History::builder().observer(Arc::clone(&shared)).build();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A non-shared observer can be reclaimed after convergence instead.
|
||||||
|
#[test]
|
||||||
|
fn into_observer_returns_the_accumulated_state() {
|
||||||
|
let mut h = History::builder().observer(Recorder::default()).build();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
// Readable in place...
|
||||||
|
assert!(!h.observer().iterations.lock().unwrap().is_empty());
|
||||||
|
|
||||||
|
// ...and reclaimable by value.
|
||||||
|
let recorder = h.into_observer();
|
||||||
|
assert!(!recorder.slices.lock().unwrap().is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Borrowing works too, for an observer that outlives the history.
|
||||||
|
#[test]
|
||||||
|
fn a_borrowed_observer_works() {
|
||||||
|
let recorder = Recorder::default();
|
||||||
|
{
|
||||||
|
let mut h = History::builder().observer(&recorder).build();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
}
|
||||||
|
assert!(!recorder.iterations.lock().unwrap().is_empty());
|
||||||
|
}
|
||||||
@@ -0,0 +1,291 @@
|
|||||||
|
//! Prediction API: N-team outcomes, draw mass, and the error paths that used
|
||||||
|
//! to be panics or silent wrong answers.
|
||||||
|
|
||||||
|
use trueskill_tt::{History, InferenceError, MAX_PREDICTED_TEAMS};
|
||||||
|
|
||||||
|
fn history_with(names: &[&'static str], p_draw: f64) -> History {
|
||||||
|
let mut h = History::builder().p_draw(p_draw).build();
|
||||||
|
// Give every competitor a recorded skill by playing a small round robin.
|
||||||
|
for pair in names.windows(2) {
|
||||||
|
h.record_winner(&pair[0], &pair[1], 1).unwrap();
|
||||||
|
}
|
||||||
|
h.converge().unwrap();
|
||||||
|
h
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn unknown_keys_are_reported_not_silently_dropped() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
|
||||||
|
let err = h
|
||||||
|
.predict_outcome(&[&[&"a"], &[&"ghost"]])
|
||||||
|
.expect_err("an unknown key must not yield a confident prediction");
|
||||||
|
assert_eq!(err, InferenceError::UnknownKey { team: 1, member: 0 });
|
||||||
|
|
||||||
|
// Every prediction entry point, not just one.
|
||||||
|
assert!(
|
||||||
|
h.predict_win_probabilities(&[&[&"a"], &[&"ghost"]])
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
assert!(h.predict_quality(&[&[&"a"], &[&"ghost"]]).is_err());
|
||||||
|
assert!(h.predict_ranking(&[&[&"a"], &[&"ghost"]], &[0, 1]).is_err());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn an_entirely_unknown_team_is_an_error() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
let err = h.predict_outcome(&[&[&"a"], &[&"x", &"y"]]).unwrap_err();
|
||||||
|
assert_eq!(err, InferenceError::UnknownKey { team: 1, member: 0 });
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn degenerate_team_shapes_are_errors_rather_than_panics() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
h.predict_outcome(&[&[&"a"]]).unwrap_err(),
|
||||||
|
InferenceError::NotEnoughTeams { got: 1 }
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
h.predict_outcome(&[]).unwrap_err(),
|
||||||
|
InferenceError::NotEnoughTeams { got: 0 }
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
h.predict_outcome(&[&[&"a"], &[]]).unwrap_err(),
|
||||||
|
InferenceError::EmptyTeam { team: 1 }
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn more_than_two_teams_no_longer_panics() {
|
||||||
|
let h = history_with(&["a", "b", "c"], 0.0);
|
||||||
|
let p = h
|
||||||
|
.predict_outcome(&[&[&"a"], &[&"b"], &[&"c"]])
|
||||||
|
.expect("three teams must be supported");
|
||||||
|
assert!((p.total() - 1.0).abs() < 1e-6, "total = {}", p.total());
|
||||||
|
// Three teams, no draws possible: exactly the six strict orderings.
|
||||||
|
assert_eq!(p.outcomes().len(), 6);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn the_outcome_space_is_capped_rather_than_hanging() {
|
||||||
|
let names: Vec<&'static str> = vec!["a", "b", "c", "d", "e", "f", "g", "h"];
|
||||||
|
let h = history_with(&names, 0.0);
|
||||||
|
|
||||||
|
let teams: Vec<&[&&'static str]> = Vec::new();
|
||||||
|
let _ = teams;
|
||||||
|
|
||||||
|
let too_many: Vec<Vec<&&str>> = names.iter().map(|n| vec![n]).collect();
|
||||||
|
let refs: Vec<&[&&str]> = too_many.iter().map(Vec::as_slice).collect();
|
||||||
|
|
||||||
|
let err = h.predict_outcome(&refs).unwrap_err();
|
||||||
|
assert_eq!(
|
||||||
|
err,
|
||||||
|
InferenceError::TooManyTeams {
|
||||||
|
got: 8,
|
||||||
|
max: MAX_PREDICTED_TEAMS
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
// The cheap paths stay available at any size.
|
||||||
|
let wins = h.predict_win_probabilities(&refs).unwrap();
|
||||||
|
assert_eq!(wins.len(), 8);
|
||||||
|
assert!(
|
||||||
|
(wins.iter().sum::<f64>() - 1.0).abs() < 1e-6,
|
||||||
|
"win probabilities must still sum to one: {wins:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The defect that made every draw-enabled prediction wrong: `[p, 1 - p]`
|
||||||
|
/// allocated no mass to a draw even with `p_draw > 0`.
|
||||||
|
#[test]
|
||||||
|
fn a_draw_carries_probability_mass_when_p_draw_is_positive() {
|
||||||
|
let h = history_with(&["a", "b"], 0.25);
|
||||||
|
let p = h.predict_outcome(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
|
|
||||||
|
let draw = p.probability_of(&[0, 0]);
|
||||||
|
assert!(draw > 0.0, "a draw-enabled model must give draws mass");
|
||||||
|
assert!((p.total() - 1.0).abs() < 1e-6, "total = {}", p.total());
|
||||||
|
|
||||||
|
let wins = p.win_probabilities();
|
||||||
|
assert!(
|
||||||
|
(wins.iter().sum::<f64>() + draw - 1.0).abs() < 1e-6,
|
||||||
|
"wins {wins:?} plus draw {draw} must be the whole space"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
(p.shared_first_place() - draw).abs() < 1e-12,
|
||||||
|
"a two-team draw is a shared first place"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn a_zero_draw_probability_admits_no_ties() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
let p = h.predict_outcome(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
|
assert_eq!(p.probability_of(&[0, 0]), 0.0);
|
||||||
|
assert!(p.shared_first_place() < 1e-12);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The two routes to a win probability run through entirely different
|
||||||
|
/// algorithms — adaptive quadrature versus the enumerated chain recursion —
|
||||||
|
/// so agreement between them is a real cross-check, not a tautology.
|
||||||
|
#[test]
|
||||||
|
fn the_cheap_and_exhaustive_paths_agree() {
|
||||||
|
for p_draw in [0.0, 0.1] {
|
||||||
|
let h = history_with(&["a", "b", "c"], p_draw);
|
||||||
|
let teams: &[&[&&str]] = &[&[&"a"], &[&"b"], &[&"c"]];
|
||||||
|
|
||||||
|
let cheap = h.predict_win_probabilities(teams).unwrap();
|
||||||
|
let exhaustive = h.predict_outcome(teams).unwrap().win_probabilities();
|
||||||
|
|
||||||
|
for (i, (a, b)) in cheap.iter().zip(&exhaustive).enumerate() {
|
||||||
|
assert!(
|
||||||
|
(a - b).abs() < 1e-6,
|
||||||
|
"p_draw={p_draw} team {i}: quadrature {a} vs enumeration {b}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn predict_ranking_agrees_with_the_distribution() {
|
||||||
|
let h = history_with(&["a", "b", "c"], 0.1);
|
||||||
|
let teams: &[&[&&str]] = &[&[&"a"], &[&"b"], &[&"c"]];
|
||||||
|
let dist = h.predict_outcome(teams).unwrap();
|
||||||
|
|
||||||
|
for (ranks, expected) in dist.outcomes() {
|
||||||
|
let direct = h.predict_ranking(teams, ranks).unwrap();
|
||||||
|
assert!(
|
||||||
|
(direct - expected).abs() < 1e-9,
|
||||||
|
"ranks {ranks:?}: {direct} vs {expected}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn predict_ranking_checks_its_shape() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
let err = h
|
||||||
|
.predict_ranking(&[&[&"a"], &[&"b"]], &[0, 1, 2])
|
||||||
|
.unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
InferenceError::MismatchedShape {
|
||||||
|
expected: 2,
|
||||||
|
got: 3,
|
||||||
|
..
|
||||||
|
}
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn the_stronger_competitor_is_favoured() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
for t in 1..=10 {
|
||||||
|
h.record_winner(&"strong", &"weak", t).unwrap();
|
||||||
|
}
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let p = h.predict_outcome(&[&[&"strong"], &[&"weak"]]).unwrap();
|
||||||
|
let (best, _) = p.most_likely().expect("a most likely outcome");
|
||||||
|
assert_eq!(best, &[0, 1], "the winner should be favoured");
|
||||||
|
|
||||||
|
let wins = p.win_probabilities();
|
||||||
|
assert!(wins[0] > wins[1], "{wins:?}");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unequal team sizes change the draw margin, because inference derives it
|
||||||
|
/// from the teams' betas. Prediction has to follow, or it describes a
|
||||||
|
/// different model than the one that will be fitted.
|
||||||
|
#[test]
|
||||||
|
fn team_size_affects_the_prediction() {
|
||||||
|
let mut h = History::builder().p_draw(0.2).build();
|
||||||
|
h.event(1)
|
||||||
|
.team(["a", "b"])
|
||||||
|
.team(["c"])
|
||||||
|
.winner(0)
|
||||||
|
.commit()
|
||||||
|
.unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let p = h.predict_outcome(&[&[&"a", &"b"], &[&"c"]]).unwrap();
|
||||||
|
assert!((p.total() - 1.0).abs() < 1e-6, "total = {}", p.total());
|
||||||
|
assert!(p.probability_of(&[0, 0]) > 0.0);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// Expected information gain
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
/// The whole point of #39: "which comparison should I run next?" is a
|
||||||
|
/// different question from "who will win?" or "is this fair?".
|
||||||
|
#[test]
|
||||||
|
fn information_gain_prefers_the_uncertain_pairing() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
|
||||||
|
// "known" and "rival" have played a lot; "newcomer" has played once.
|
||||||
|
for t in 1..=15 {
|
||||||
|
h.record_winner(&"known", &"rival", t).unwrap();
|
||||||
|
h.record_winner(&"rival", &"known", t + 100).unwrap();
|
||||||
|
}
|
||||||
|
h.record_winner(&"known", &"newcomer", 500).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let settled = h
|
||||||
|
.expected_information_gain(&[&[&"known"], &[&"rival"]])
|
||||||
|
.unwrap();
|
||||||
|
let unknown = h
|
||||||
|
.expected_information_gain(&[&[&"known"], &[&"newcomer"]])
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
unknown > settled,
|
||||||
|
"pairing against the newcomer should teach more: {unknown} vs {settled}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The analytic ceiling, through the `History` entry point rather than the
|
||||||
|
/// standalone one.
|
||||||
|
#[test]
|
||||||
|
fn information_gain_respects_the_entropy_ceiling() {
|
||||||
|
let h = history_with(&["a", "b", "c"], 0.0);
|
||||||
|
|
||||||
|
let two = h.expected_information_gain(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
|
assert!(
|
||||||
|
(0.0..=std::f64::consts::LN_2).contains(&two),
|
||||||
|
"two-team EIG {two} outside [0, ln 2]"
|
||||||
|
);
|
||||||
|
|
||||||
|
let three = h
|
||||||
|
.expected_information_gain(&[&[&"a"], &[&"b"], &[&"c"]])
|
||||||
|
.unwrap();
|
||||||
|
assert!(
|
||||||
|
(0.0..=6.0f64.ln()).contains(&three),
|
||||||
|
"three-team EIG {three} outside [0, ln 6]"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn information_gain_reports_unknown_keys() {
|
||||||
|
let h = history_with(&["a", "b"], 0.0);
|
||||||
|
assert_eq!(
|
||||||
|
h.expected_information_gain(&[&[&"a"], &[&"ghost"]])
|
||||||
|
.unwrap_err(),
|
||||||
|
InferenceError::UnknownKey { team: 1, member: 0 }
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A draw-enabled history has three outcomes to weigh rather than two, so the
|
||||||
|
/// draw branch must actually be reachable through this path.
|
||||||
|
#[test]
|
||||||
|
fn information_gain_accounts_for_draws() {
|
||||||
|
let with_draws = history_with(&["a", "b"], 0.25);
|
||||||
|
let g = with_draws
|
||||||
|
.expected_information_gain(&[&[&"a"], &[&"b"]])
|
||||||
|
.unwrap();
|
||||||
|
assert!(g > 0.0 && g <= 3.0f64.ln(), "{g}");
|
||||||
|
|
||||||
|
// The draw outcome carries mass, so it is genuinely being weighed.
|
||||||
|
let dist = with_draws.predict_outcome(&[&[&"a"], &[&"b"]]).unwrap();
|
||||||
|
assert!(dist.probability_of(&[0, 0]) > 0.0);
|
||||||
|
}
|
||||||
@@ -0,0 +1,167 @@
|
|||||||
|
//! Property-based tests over generated histories.
|
||||||
|
//!
|
||||||
|
//! The golden suite pins exact values against the Python/Julia reference on a
|
||||||
|
//! handful of fixtures. These pin *invariants* over inputs nobody wrote by
|
||||||
|
//! hand, which is where the defects this crate has actually shipped were
|
||||||
|
//! hiding: a linear evidence product that underflowed only past ~1000 teams,
|
||||||
|
//! and a batching path no golden exercised because every golden ingests in one
|
||||||
|
//! call.
|
||||||
|
|
||||||
|
mod common;
|
||||||
|
|
||||||
|
use common::assert_finite;
|
||||||
|
use proptest::prelude::*;
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{ConvergenceOptions, Event, History, Member, Outcome, Team};
|
||||||
|
|
||||||
|
/// Distinct competitors, so no event pits someone against themselves.
|
||||||
|
fn pairs() -> impl Strategy<Value = Vec<(usize, usize)>> {
|
||||||
|
prop::collection::vec((0usize..8, 0usize..8), 1..24)
|
||||||
|
.prop_map(|v| v.into_iter().filter(|(a, b)| a != b).collect::<Vec<_>>())
|
||||||
|
.prop_filter("needs at least one valid pair", |v| !v.is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
const KEYS: [&str; 8] = ["a", "b", "c", "d", "e", "f", "g", "h"];
|
||||||
|
|
||||||
|
fn history_from(games: &[(usize, usize)]) -> History {
|
||||||
|
let mut h = History::builder()
|
||||||
|
.convergence(ConvergenceOptions {
|
||||||
|
max_iter: 200,
|
||||||
|
epsilon: 1e-10,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
})
|
||||||
|
.build();
|
||||||
|
|
||||||
|
let events: Vec<Event<i64, &'static str>> = games
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(i, &(a, b))| Event {
|
||||||
|
time: i as i64 + 1,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(KEYS[a])]),
|
||||||
|
Team::with_members([Member::new(KEYS[b])]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
h.add_events(events).unwrap();
|
||||||
|
|
||||||
|
h
|
||||||
|
}
|
||||||
|
|
||||||
|
proptest! {
|
||||||
|
#![proptest_config(ProptestConfig::with_cases(48))]
|
||||||
|
|
||||||
|
/// Whatever the schedule of games, convergence must not produce NaN or an
|
||||||
|
/// improper posterior. `converge` returns `NonFiniteResult` rather than
|
||||||
|
/// silently reporting a NaN step as converged, so a break shows up here as
|
||||||
|
/// either an Err or a non-finite curve point.
|
||||||
|
#[test]
|
||||||
|
fn converged_posteriors_are_always_finite(games in pairs()) {
|
||||||
|
let mut h = history_from(&games);
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
for key in KEYS {
|
||||||
|
for (time, g) in h.learning_curve(key) {
|
||||||
|
assert_finite(g, &format!("{key} at t={time}"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Log-evidence is a log probability: finite, and never above zero.
|
||||||
|
///
|
||||||
|
/// The linear-product implementation this replaced underflowed to zero on
|
||||||
|
/// long chains, making `ln(0)` = -inf — finite-ness is the property that
|
||||||
|
/// would have caught it.
|
||||||
|
#[test]
|
||||||
|
fn log_evidence_is_a_finite_log_probability(games in pairs()) {
|
||||||
|
let mut h = history_from(&games);
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let batch = h.log_evidence();
|
||||||
|
let filtered = h.filtered_log_evidence();
|
||||||
|
|
||||||
|
prop_assert!(batch.is_finite(), "batch log-evidence {batch} is not finite");
|
||||||
|
prop_assert!(batch <= 0.0, "batch log-evidence {batch} exceeds zero");
|
||||||
|
prop_assert!(filtered.is_finite(), "filtered log-evidence {filtered} is not finite");
|
||||||
|
prop_assert!(filtered <= 0.0, "filtered log-evidence {filtered} exceeds zero");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Filtered estimates must not depend on whether `converge` has run — the
|
||||||
|
/// property the whole forward-only design rests on.
|
||||||
|
#[test]
|
||||||
|
fn filtered_evidence_is_invariant_to_convergence(games in pairs()) {
|
||||||
|
let mut h = history_from(&games);
|
||||||
|
|
||||||
|
let before = h.filtered_log_evidence();
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let after = h.filtered_log_evidence();
|
||||||
|
|
||||||
|
prop_assert!(
|
||||||
|
(before - after).abs() < 1e-8,
|
||||||
|
"filtered evidence moved across converge(): {before} -> {after}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Ingesting the same games one at a time must reach the same fixed point
|
||||||
|
/// as ingesting them in one call.
|
||||||
|
#[test]
|
||||||
|
fn ingestion_order_does_not_change_the_answer(games in pairs()) {
|
||||||
|
let batched = {
|
||||||
|
let mut h = history_from(&games);
|
||||||
|
h.converge().unwrap();
|
||||||
|
h
|
||||||
|
};
|
||||||
|
|
||||||
|
let incremental = {
|
||||||
|
let mut h = History::builder()
|
||||||
|
.convergence(ConvergenceOptions {
|
||||||
|
max_iter: 200,
|
||||||
|
epsilon: 1e-10,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
})
|
||||||
|
.build();
|
||||||
|
|
||||||
|
for (i, &(a, b)) in games.iter().enumerate() {
|
||||||
|
h.add_events([Event {
|
||||||
|
time: i as i64 + 1,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new(KEYS[a])]),
|
||||||
|
Team::with_members([Member::new(KEYS[b])]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::winner(0, 2),
|
||||||
|
}])
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
h.converge().unwrap();
|
||||||
|
h
|
||||||
|
};
|
||||||
|
|
||||||
|
for key in KEYS {
|
||||||
|
let one = batched.current_skill(key);
|
||||||
|
let other = incremental.current_skill(key);
|
||||||
|
|
||||||
|
match (one, other) {
|
||||||
|
(Some(one), Some(other)) => {
|
||||||
|
prop_assert!(
|
||||||
|
(one.mu() - other.mu()).abs() < 1e-6
|
||||||
|
&& (one.sigma() - other.sigma()).abs() < 1e-6,
|
||||||
|
"{key}: batched mu={} sigma={}, incremental mu={} sigma={}",
|
||||||
|
one.mu(),
|
||||||
|
one.sigma(),
|
||||||
|
other.mu(),
|
||||||
|
other.sigma()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
(None, None) => {}
|
||||||
|
_ => prop_assert!(false, "{key} present in only one history"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
//! `quality()` beyond two rating groups.
|
||||||
|
//!
|
||||||
|
//! The historical golden (two equal singletons) is asserted in
|
||||||
|
//! `src/lib.rs::tests::test_quality`. These cover the N-group generalisation,
|
||||||
|
//! which previously panicked with an out-of-bounds index at 3+ groups.
|
||||||
|
|
||||||
|
use trueskill_tt::{Gaussian, quality};
|
||||||
|
|
||||||
|
const BETA: f64 = 25.0 / 3.0 / 2.0;
|
||||||
|
|
||||||
|
fn rating(mu: f64, sigma: f64) -> Gaussian {
|
||||||
|
Gaussian::from_ms(mu, sigma)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn three_equal_groups_is_finite_and_in_range() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
let q = quality(&[&[r], &[r], &[r]], BETA);
|
||||||
|
|
||||||
|
assert!(q.is_finite(), "quality must be finite, got {q}");
|
||||||
|
assert!((0.0..=1.0).contains(&q), "quality out of range: {q}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn quality_supports_many_groups() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
for n in 2..=8 {
|
||||||
|
let holders: Vec<[Gaussian; 1]> = (0..n).map(|_| [r]).collect();
|
||||||
|
let groups: Vec<&[Gaussian]> = holders.iter().map(|g| g.as_slice()).collect();
|
||||||
|
let q = quality(&groups, BETA);
|
||||||
|
assert!(q.is_finite(), "n={n}: quality must be finite, got {q}");
|
||||||
|
assert!((0.0..=1.0).contains(&q), "n={n}: out of range: {q}");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Equal-strength groups are the best-matched case: introducing a skill gap
|
||||||
|
/// must lower quality.
|
||||||
|
#[test]
|
||||||
|
fn imbalance_lowers_quality() {
|
||||||
|
let strong = rating(40.0, 3.0);
|
||||||
|
let average = rating(25.0, 3.0);
|
||||||
|
|
||||||
|
let balanced = quality(&[&[average], &[average], &[average]], BETA);
|
||||||
|
let lopsided = quality(&[&[strong], &[average], &[average]], BETA);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
lopsided < balanced,
|
||||||
|
"expected imbalanced quality {lopsided} < balanced {balanced}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Quality is a property of the multiset of groups, not their order.
|
||||||
|
#[test]
|
||||||
|
fn quality_is_permutation_invariant() {
|
||||||
|
let a = rating(30.0, 2.0);
|
||||||
|
let b = rating(25.0, 3.0);
|
||||||
|
let c = rating(20.0, 4.0);
|
||||||
|
|
||||||
|
let forward = quality(&[&[a], &[b], &[c]], BETA);
|
||||||
|
let reversed = quality(&[&[c], &[b], &[a]], BETA);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(forward - reversed).abs() < 1e-9,
|
||||||
|
"permutation changed quality: {forward} vs {reversed}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn multi_player_groups_work() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
let q = quality(&[&[r, r], &[r, r], &[r, r]], BETA);
|
||||||
|
assert!(q.is_finite());
|
||||||
|
assert!((0.0..=1.0).contains(&q));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn uneven_group_sizes_work() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
let q = quality(&[&[r, r], &[r], &[r, r, r]], BETA);
|
||||||
|
assert!(q.is_finite(), "got {q}");
|
||||||
|
assert!((0.0..=1.0).contains(&q), "got {q}");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "at least 2 rating groups")]
|
||||||
|
fn single_group_panics_with_clear_message() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
let _ = quality(&[&[r]], BETA);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "at least 2 rating groups")]
|
||||||
|
fn zero_groups_panics_with_clear_message() {
|
||||||
|
let _ = quality(&[], BETA);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "non-empty")]
|
||||||
|
fn empty_group_panics_with_clear_message() {
|
||||||
|
let r = rating(25.0, 3.0);
|
||||||
|
let _ = quality(&[&[r], &[]], BETA);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn history_predict_quality_supports_three_teams() {
|
||||||
|
use trueskill_tt::History;
|
||||||
|
|
||||||
|
let mut h = History::default();
|
||||||
|
h.record_winner(&"a", &"b", 1).unwrap();
|
||||||
|
h.record_winner(&"b", &"c", 2).unwrap();
|
||||||
|
h.converge().unwrap();
|
||||||
|
|
||||||
|
let q = h.predict_quality(&[&[&"a"], &[&"b"], &[&"c"]]).unwrap();
|
||||||
|
assert!(
|
||||||
|
q.is_finite(),
|
||||||
|
"3-team predict_quality must be finite, got {q}"
|
||||||
|
);
|
||||||
|
assert!((0.0..=1.0).contains(&q), "out of range: {q}");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `quality()` for N identical teams has a closed form, which pins the N-group
|
||||||
|
/// determinant path across the whole range rather than at a single golden.
|
||||||
|
///
|
||||||
|
/// For two identical single-player teams the standard result is
|
||||||
|
/// `sqrt(2b^2 / (2b^2 + s1^2 + s2^2))`. With the conventional parameters
|
||||||
|
/// (`sigma = 25/3`, `beta = 25/6`) that ratio is exactly `1/5`, and the N-group
|
||||||
|
/// generalisation is `(1/5)^((n-1)/2)` — one factor per adjacent pair.
|
||||||
|
///
|
||||||
|
/// The n=3 and n=5 values this produces (0.200 and 0.040) are also what the
|
||||||
|
/// `trueskill` Python package returns for the same configuration, so this
|
||||||
|
/// doubles as the cross-implementation check the README asked for.
|
||||||
|
#[test]
|
||||||
|
fn quality_of_identical_teams_follows_its_closed_form() {
|
||||||
|
let g = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||||||
|
let beta = 25.0 / 6.0;
|
||||||
|
|
||||||
|
for n in 2..=10usize {
|
||||||
|
let groups: Vec<Vec<Gaussian>> = (0..n).map(|_| vec![g]).collect();
|
||||||
|
let refs: Vec<&[Gaussian]> = groups.iter().map(Vec::as_slice).collect();
|
||||||
|
|
||||||
|
let got = quality(&refs, beta);
|
||||||
|
let expected = 0.2f64.powf((n - 1) as f64 / 2.0);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
(got - expected).abs() / expected < 1e-9,
|
||||||
|
"n={n}: quality {got}, closed form {expected}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Spot-check against the two values the `trueskill` Python package is known
|
||||||
|
/// to produce for this configuration, stated as literals so a future change to
|
||||||
|
/// the closed-form reasoning above cannot quietly take these with it.
|
||||||
|
#[test]
|
||||||
|
fn quality_matches_the_reference_implementation() {
|
||||||
|
let g = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||||||
|
let beta = 25.0 / 6.0;
|
||||||
|
|
||||||
|
let three: Vec<Vec<Gaussian>> = (0..3).map(|_| vec![g]).collect();
|
||||||
|
let refs: Vec<&[Gaussian]> = three.iter().map(Vec::as_slice).collect();
|
||||||
|
assert!((quality(&refs, beta) - 0.200).abs() < 1e-9);
|
||||||
|
|
||||||
|
let five: Vec<Vec<Gaussian>> = (0..5).map(|_| vec![g]).collect();
|
||||||
|
let refs: Vec<&[Gaussian]> = five.iter().map(Vec::as_slice).collect();
|
||||||
|
assert!((quality(&refs, beta) - 0.040).abs() < 1e-9);
|
||||||
|
}
|
||||||
@@ -0,0 +1,186 @@
|
|||||||
|
//! Input validation must hold in **release**, where `debug_assert!` is gone.
|
||||||
|
//!
|
||||||
|
//! The engine guards itself with `debug_assert!`, which documents invariants
|
||||||
|
//! but vanishes in the profile users actually ship. Anything reachable from the
|
||||||
|
//! public API has to be rejected with an `InferenceError` instead, at the
|
||||||
|
//! boundary, rather than becoming NaN or an out-of-bounds panic deep inside
|
||||||
|
//! `run_chain`.
|
||||||
|
//!
|
||||||
|
//! `GameOptions` and `ConvergenceOptions` both have public fields, so the
|
||||||
|
//! eager asserts on `HistoryBuilder` do not cover the `Game` constructors —
|
||||||
|
//! a caller can build the options struct directly.
|
||||||
|
|
||||||
|
use smallvec::smallvec;
|
||||||
|
use trueskill_tt::{
|
||||||
|
ConstantDrift, ConvergenceOptions, Event, Game, GameOptions, Gaussian, History, InferenceError,
|
||||||
|
Member, Outcome, Rating, Team,
|
||||||
|
};
|
||||||
|
|
||||||
|
type R = Rating<i64, ConstantDrift>;
|
||||||
|
|
||||||
|
fn rating() -> R {
|
||||||
|
R::new(
|
||||||
|
Gaussian::from_ms(25.0, 25.0 / 3.0),
|
||||||
|
25.0 / 6.0,
|
||||||
|
ConstantDrift(0.0),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn options_with_alpha(alpha: f64) -> GameOptions {
|
||||||
|
GameOptions {
|
||||||
|
convergence: ConvergenceOptions {
|
||||||
|
alpha,
|
||||||
|
..ConvergenceOptions::default()
|
||||||
|
},
|
||||||
|
..GameOptions::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `alpha == 0.0` leaves every EP update unapplied, so inference silently
|
||||||
|
/// returns the priors — the worst possible failure, since the output looks
|
||||||
|
/// entirely reasonable.
|
||||||
|
#[test]
|
||||||
|
fn ranked_rejects_a_zero_damping_factor() {
|
||||||
|
let (a, b) = (rating(), rating());
|
||||||
|
let err = Game::<i64, _>::ranked(
|
||||||
|
&[&[a], &[b]],
|
||||||
|
Outcome::winner(0, 2),
|
||||||
|
&options_with_alpha(0.0),
|
||||||
|
)
|
||||||
|
.expect_err("alpha = 0 must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::InvalidParameter { name: "alpha", .. }),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn ranked_rejects_an_out_of_range_damping_factor() {
|
||||||
|
let (a, b) = (rating(), rating());
|
||||||
|
for alpha in [-0.5, 1.5, f64::NAN] {
|
||||||
|
let err = Game::<i64, _>::ranked(
|
||||||
|
&[&[a], &[b]],
|
||||||
|
Outcome::winner(0, 2),
|
||||||
|
&options_with_alpha(alpha),
|
||||||
|
)
|
||||||
|
.expect_err("alpha out of (0, 1] must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::InvalidParameter { name: "alpha", .. }),
|
||||||
|
"alpha={alpha}: got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn scored_rejects_a_bad_damping_factor() {
|
||||||
|
let (a, b) = (rating(), rating());
|
||||||
|
let err = Game::<i64, _>::scored(
|
||||||
|
&[&[a], &[b]],
|
||||||
|
Outcome::scores([21.0, 9.0]),
|
||||||
|
&options_with_alpha(0.0),
|
||||||
|
)
|
||||||
|
.expect_err("alpha = 0 must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::InvalidParameter { name: "alpha", .. }),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Already covered by `Game::ranked`, asserted here so the release-mode
|
||||||
|
/// guarantee is stated in one place.
|
||||||
|
#[test]
|
||||||
|
fn ranked_rejects_an_out_of_range_draw_probability() {
|
||||||
|
let (a, b) = (rating(), rating());
|
||||||
|
for p_draw in [-0.5, 1.0, 1.5] {
|
||||||
|
let options = GameOptions {
|
||||||
|
p_draw,
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
assert!(
|
||||||
|
Game::<i64, _>::ranked(&[&[a], &[b]], Outcome::winner(0, 2), &options).is_err(),
|
||||||
|
"p_draw={p_draw} must be rejected"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn scored_rejects_a_non_positive_noise() {
|
||||||
|
let (a, b) = (rating(), rating());
|
||||||
|
for score_sigma in [0.0, -1.0, f64::NAN] {
|
||||||
|
let options = GameOptions {
|
||||||
|
score_sigma,
|
||||||
|
..GameOptions::default()
|
||||||
|
};
|
||||||
|
assert!(
|
||||||
|
Game::<i64, _>::scored(&[&[a], &[b]], Outcome::scores([21.0, 9.0]), &options).is_err(),
|
||||||
|
"score_sigma={score_sigma} must be rejected"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A tie with no draw probability makes the truncation margin zero and the
|
||||||
|
/// two-sided update evaluate 0/0. Ingestion must refuse it.
|
||||||
|
#[test]
|
||||||
|
fn ingestion_rejects_a_tie_without_a_draw_probability() {
|
||||||
|
let mut h = History::builder().p_draw(0.0).build();
|
||||||
|
let err = h
|
||||||
|
.add_events(vec![Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::draw(2),
|
||||||
|
}])
|
||||||
|
.expect_err("a tie with p_draw = 0 must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::TieWithoutDrawProbability { .. }),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// `Outcome::scores_with_sigma` documents that a non-positive sigma is
|
||||||
|
/// accepted at construction and rejected at ingestion.
|
||||||
|
#[test]
|
||||||
|
fn ingestion_rejects_a_non_positive_per_event_score_sigma() {
|
||||||
|
for sigma in [0.0, -1.0, f64::NAN] {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
let err = h
|
||||||
|
.add_events(vec![Event {
|
||||||
|
time: 0,
|
||||||
|
teams: smallvec![
|
||||||
|
Team::with_members([Member::new("a")]),
|
||||||
|
Team::with_members([Member::new("b")]),
|
||||||
|
],
|
||||||
|
outcome: Outcome::scores_with_sigma([21.0, 9.0], sigma),
|
||||||
|
}])
|
||||||
|
.expect_err("a non-positive per-event sigma must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::InvalidParameter { .. }),
|
||||||
|
"sigma={sigma}: got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Per-team weights must match that team's membership. The top-level length
|
||||||
|
/// checks in ingestion do not cover the inner dimension.
|
||||||
|
#[test]
|
||||||
|
fn ingestion_rejects_weights_that_do_not_match_their_team() {
|
||||||
|
let mut h = History::builder().build();
|
||||||
|
let mut team = Team::with_members([Member::new("a"), Member::new("b")]);
|
||||||
|
team.members[0].weight = 1.0;
|
||||||
|
|
||||||
|
let err = h
|
||||||
|
.event(0)
|
||||||
|
.team(["a", "b"])
|
||||||
|
.team(["c"])
|
||||||
|
// Three weights for a two-member team.
|
||||||
|
.weights([1.0, 1.0, 1.0])
|
||||||
|
.winner(0)
|
||||||
|
.commit()
|
||||||
|
.expect_err("a weight/member length mismatch must be rejected");
|
||||||
|
assert!(
|
||||||
|
matches!(err, InferenceError::MismatchedShape { .. }),
|
||||||
|
"got {err:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user