Skip to content

Commit d46f2a3

Browse files
committed
Feat(metrics): Add nonparam bootstrapping
1 parent 67e3f62 commit d46f2a3

1 file changed

Lines changed: 174 additions & 2 deletions

File tree

codeclash/analysis/metrics/elo2.py

Lines changed: 174 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,14 @@
22
import argparse
33
import json
44
from collections import defaultdict
5+
from dataclasses import dataclass
56
from pathlib import Path
67
from typing import Literal, TypeAlias, get_args
78

89
import matplotlib.pyplot as plt
910
import numpy as np
1011
from scipy.optimize import minimize
12+
from scipy.stats import kendalltau, spearmanr
1113
from tqdm import tqdm
1214

1315
from codeclash.analysis.metrics.elo import get_scores
@@ -277,12 +279,19 @@ def print_matrix(self) -> None:
277279
class BradleyTerryFitter:
278280
def __init__(
279281
self,
280-
matchups: dict[tuple[str, str], list[float]],
282+
win_matrix: dict[tuple[str, str], list[float]],
281283
*,
282284
regularization: float = 0.01,
283285
compute_uncertainties: bool = True,
284286
):
285-
self.matchups = matchups
287+
"""Fit Bradley-Terry model to a win matrix
288+
289+
Args:
290+
win_matrix: Dictionary mapping player pairs to win counts
291+
regularization: L2 regularization strength
292+
compute_uncertainties: Whether to compute uncertainties
293+
"""
294+
self.matchups = win_matrix
286295
self.regularization = regularization
287296
self.compute_uncertainties = compute_uncertainties
288297
self.result: dict | None = None
@@ -626,6 +635,150 @@ def create_validation_plots(self, output_dir: Path, regularization: float = 0.01
626635
logger.info(f"Saved validation plot: {output_path}")
627636

628637

638+
@dataclass
639+
class BootStrapRankStabilityConfig:
640+
n_bootstrap: int = 200
641+
game: str = "ALL"
642+
regularization: float = 0.01
643+
topks: list[int] | None = None
644+
rng_seed: int | None = None
645+
646+
def __post_init__(self) -> None:
647+
if self.topks is None:
648+
self.topks = [1, 3, 5]
649+
650+
651+
class BootStrapRankStability:
652+
def __init__(
653+
self,
654+
builder: ScoreMatrixBuilder,
655+
*,
656+
n_bootstrap: int = 200,
657+
game: str = "ALL",
658+
regularization: float = 0.01,
659+
topks: list[int] | None = None,
660+
):
661+
self.builder = builder
662+
self.n_bootstrap = n_bootstrap
663+
self.game = game
664+
self.regularization = regularization
665+
self.topks = topks
666+
667+
@staticmethod
668+
def _elos_from_result(result: dict) -> dict[str, float]:
669+
return {p: BradleyTerryFitter.bt_to_elo(s) for p, s in zip(result["players"], result["strengths"])}
670+
671+
@staticmethod
672+
def _ranking_from_elos(elos: dict[str, float]) -> list[str]:
673+
return [p for p, _ in sorted(elos.items(), key=lambda kv: kv[1], reverse=True)]
674+
675+
@staticmethod
676+
def _positions(ranking: list[str]) -> dict[str, int]:
677+
return {p: i for i, p in enumerate(ranking)}
678+
679+
@staticmethod
680+
def _max_footrule(n: int) -> float:
681+
return (n * n) / 2 if n % 2 == 0 else (n * n - 1) / 2
682+
683+
def _fit_on_matrix(self, matchups: dict[tuple[str, str], list[float]]) -> dict:
684+
fitter = BradleyTerryFitter(matchups, regularization=self.regularization, compute_uncertainties=False)
685+
return fitter.fit()
686+
687+
def run(self) -> None:
688+
game = self.game
689+
assert game in self.builder.win_matrix, f"Game '{game}' not found in win matrix"
690+
691+
baseline_res = self._fit_on_matrix(self.builder.win_matrix[game])
692+
baseline_elos = self._elos_from_result(baseline_res)
693+
baseline_ranking = self._ranking_from_elos(baseline_elos)
694+
players = baseline_ranking
695+
n = len(players)
696+
topks = list(range(1, n + 1)) if self.topks is None else [k for k in self.topks if k <= n]
697+
698+
rank_samples: dict[str, list[int]] = {p: [] for p in players}
699+
tau_vals: list[float] = []
700+
rho_vals: list[float] = []
701+
footrule_vals: list[float] = []
702+
topk_overlap: dict[int, list[float]] = {k: [] for k in topks}
703+
top1_match = 0
704+
pair_agree = 0
705+
total_pairs = n * (n - 1) // 2
706+
707+
base_pos = self._positions(baseline_ranking)
708+
709+
rng = np.random.default_rng(42)
710+
for _ in tqdm(range(self.n_bootstrap), desc="Bootstrap samples"):
711+
boot = self.builder.get_nonparametric_bootstrap(rng=rng)
712+
res = self._fit_on_matrix(boot[game])
713+
elos = self._elos_from_result(res)
714+
ranking = self._ranking_from_elos(elos)
715+
pos = self._positions(ranking)
716+
717+
for p in players:
718+
rank_samples[p].append(pos[p] + 1)
719+
720+
base_rank_arr = np.array([base_pos[p] + 1 for p in players])
721+
boot_rank_arr = np.array([pos[p] + 1 for p in players])
722+
tau = kendalltau(base_rank_arr, boot_rank_arr, variant="b").correlation
723+
rho = spearmanr(base_rank_arr, boot_rank_arr).correlation
724+
tau_vals.append(float(tau) if tau is not None else float("nan"))
725+
rho_vals.append(float(rho) if rho is not None else float("nan"))
726+
727+
foot = float(np.abs(base_rank_arr - boot_rank_arr).sum())
728+
footrule_vals.append(foot / self._max_footrule(n))
729+
730+
for k in topks:
731+
base_set = set(baseline_ranking[:k])
732+
boot_set = set(ranking[:k])
733+
inter = len(base_set & boot_set)
734+
topk_overlap[k].append(inter / k)
735+
736+
if ranking and baseline_ranking and ranking[0] == baseline_ranking[0]:
737+
top1_match += 1
738+
739+
agree = 0
740+
for i in range(n):
741+
for j in range(i + 1, n):
742+
pi, pj = players[i], players[j]
743+
agree += int((base_pos[pi] < base_pos[pj]) == (pos[pi] < pos[pj]))
744+
pair_agree += agree
745+
746+
mean_tau = float(np.nanmean(np.array(tau_vals))) if tau_vals else float("nan")
747+
mean_rho = float(np.nanmean(np.array(rho_vals))) if rho_vals else float("nan")
748+
mean_foot = float(np.nanmean(np.array(footrule_vals))) if footrule_vals else float("nan")
749+
top1_consistency = top1_match / self.n_bootstrap if self.n_bootstrap > 0 else float("nan")
750+
pairwise_agreement = (pair_agree / (self.n_bootstrap * total_pairs)) if total_pairs > 0 else float("nan")
751+
752+
lines = []
753+
lines.append("\nRank stability (bootstrap)")
754+
lines.append(f"Game: {game}")
755+
lines.append(f"Bootstraps: {self.n_bootstrap}")
756+
lines.append("")
757+
lines.append(f"{'Metric':<28} {'Value':>10}")
758+
lines.append("-" * 40)
759+
lines.append(f"{'Kendall tau (avg)':<28} {mean_tau:>10.3f}")
760+
lines.append(f"{'Spearman rho (avg)':<28} {mean_rho:>10.3f}")
761+
lines.append(f"{'Footrule (avg, norm)':<28} {mean_foot:>10.3f}")
762+
lines.append(f"{'Top-1 consistency':<28} {top1_consistency:>10.3f}")
763+
lines.append(f"{'Pairwise order agree':<28} {pairwise_agreement:>10.3f}")
764+
for k in topks:
765+
lines.append(f"{f'Top-{k} overlap (avg)':<28} {float(np.mean(topk_overlap[k])):>10.3f}")
766+
for ln in lines:
767+
logger.info(ln)
768+
769+
header = f"\n{'Model':<30} {'MeanRank':>9} {'StdRank':>8} " + " ".join([f"P@{k:>2}" for k in topks])
770+
logger.info(header)
771+
logger.info("-" * max(40, len(header)))
772+
for p in players:
773+
ranks = np.array(rank_samples[p], dtype=float)
774+
mean_r = float(np.mean(ranks))
775+
std_r = float(np.std(ranks, ddof=0))
776+
probs = []
777+
for k in topks:
778+
probs.append(np.mean(ranks <= k))
779+
logger.info(f"{p:<30} {mean_r:9.2f} {std_r:8.2f} " + " ".join([f"{float(pr):>5.2f}" for pr in probs]))
780+
781+
629782
def print_results(results: dict[str, dict]) -> None:
630783
"""Print fitted strengths and Elo ratings for all games.
631784
@@ -701,6 +854,16 @@ def print_results(results: dict[str, dict]) -> None:
701854
default=Path("elo2_plots"),
702855
help="Directory to save Elo plots (default: elo2_plots)",
703856
)
857+
parser.add_argument("--rank-stability", action="store_true", help="Run bootstrap rank stability analysis")
858+
parser.add_argument("--rank-stability-game", type=str, default="ALL", help="Game to analyze (default: ALL)")
859+
parser.add_argument("--rank-stability-n", type=int, default=200, help="Number of bootstrap samples")
860+
parser.add_argument(
861+
"--rank-stability-topk",
862+
type=int,
863+
nargs="*",
864+
default=None,
865+
help="Top-k cutoffs for overlap metrics (default: all players)",
866+
)
704867
args = parser.parse_args()
705868

706869
builder = ScoreMatrixBuilder(
@@ -736,3 +899,12 @@ def print_results(results: dict[str, dict]) -> None:
736899
plotter.create_validation_plots(args.validation_dir, regularization=args.regularization)
737900
if args.elo_plot:
738901
plotter.create_elo_plots(args.elo_plot_dir)
902+
903+
if args.rank_stability:
904+
BootStrapRankStability(
905+
builder,
906+
n_bootstrap=args.rank_stability_n,
907+
game=args.rank_stability_game,
908+
regularization=args.regularization,
909+
topks=args.rank_stability_topk if args.rank_stability_topk else None,
910+
).run()

0 commit comments

Comments
 (0)