|
2 | 2 | import argparse |
3 | 3 | import json |
4 | 4 | from collections import defaultdict |
| 5 | +from dataclasses import dataclass |
5 | 6 | from pathlib import Path |
6 | 7 | from typing import Literal, TypeAlias, get_args |
7 | 8 |
|
8 | 9 | import matplotlib.pyplot as plt |
9 | 10 | import numpy as np |
10 | 11 | from scipy.optimize import minimize |
| 12 | +from scipy.stats import kendalltau, spearmanr |
11 | 13 | from tqdm import tqdm |
12 | 14 |
|
13 | 15 | from codeclash.analysis.metrics.elo import get_scores |
@@ -277,12 +279,19 @@ def print_matrix(self) -> None: |
277 | 279 | class BradleyTerryFitter: |
278 | 280 | def __init__( |
279 | 281 | self, |
280 | | - matchups: dict[tuple[str, str], list[float]], |
| 282 | + win_matrix: dict[tuple[str, str], list[float]], |
281 | 283 | *, |
282 | 284 | regularization: float = 0.01, |
283 | 285 | compute_uncertainties: bool = True, |
284 | 286 | ): |
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 |
286 | 295 | self.regularization = regularization |
287 | 296 | self.compute_uncertainties = compute_uncertainties |
288 | 297 | self.result: dict | None = None |
@@ -626,6 +635,150 @@ def create_validation_plots(self, output_dir: Path, regularization: float = 0.01 |
626 | 635 | logger.info(f"Saved validation plot: {output_path}") |
627 | 636 |
|
628 | 637 |
|
| 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 | + |
629 | 782 | def print_results(results: dict[str, dict]) -> None: |
630 | 783 | """Print fitted strengths and Elo ratings for all games. |
631 | 784 |
|
@@ -701,6 +854,16 @@ def print_results(results: dict[str, dict]) -> None: |
701 | 854 | default=Path("elo2_plots"), |
702 | 855 | help="Directory to save Elo plots (default: elo2_plots)", |
703 | 856 | ) |
| 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 | + ) |
704 | 867 | args = parser.parse_args() |
705 | 868 |
|
706 | 869 | builder = ScoreMatrixBuilder( |
@@ -736,3 +899,12 @@ def print_results(results: dict[str, dict]) -> None: |
736 | 899 | plotter.create_validation_plots(args.validation_dir, regularization=args.regularization) |
737 | 900 | if args.elo_plot: |
738 | 901 | 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