|
14 | 14 | from tqdm import tqdm |
15 | 15 |
|
16 | 16 | from codeclash.analysis.significance import calculate_p_value |
17 | | -from codeclash.analysis.viz.utils import ASSETS_DIR, FONT_BOLD, MODEL_TO_DISPLAY_NAME |
| 17 | +from codeclash.analysis.viz.utils import ASSETS_DIR, FONT_BOLD, MODEL_TO_DISPLAY_NAME, model_display_name |
18 | 18 | from codeclash.constants import LOCAL_LOG_DIR, RESULT_TIE |
19 | 19 | from codeclash.utils.log import add_file_handler, get_logger |
20 | 20 |
|
@@ -75,9 +75,6 @@ def __init__( |
75 | 75 | lambda: defaultdict(list) |
76 | 76 | ) |
77 | 77 |
|
78 | | - def _get_unique_model_name(self, model: str) -> str: |
79 | | - return model.rpartition("/")[2] |
80 | | - |
81 | 78 | def _get_sorted_pair(self, p1: str, p2: str) -> tuple[str, str]: |
82 | 79 | return tuple(sorted([p1, p2])) |
83 | 80 |
|
@@ -154,8 +151,6 @@ def _process_tournament(self, metadata_path: Path) -> None: |
154 | 151 | return |
155 | 152 |
|
156 | 153 | player_names = [p["name"] for p in players] |
157 | | - models = [p["config"]["model"]["model_name"].strip("@") for p in players] |
158 | | - |
159 | 154 | # Aggregate scores for each round |
160 | 155 | p1_round_scores = [] |
161 | 156 | p2_round_scores = [] |
@@ -199,7 +194,7 @@ def _process_tournament(self, metadata_path: Path) -> None: |
199 | 194 | p2_score = sum(p2_round_scores) |
200 | 195 |
|
201 | 196 | # Convert to unique names and sorted pair when updating matrix |
202 | | - unique_names = [self._get_unique_model_name(m) for m in models] |
| 197 | + unique_names = player_names |
203 | 198 | sorted_pair = self._get_sorted_pair(unique_names[0], unique_names[1]) |
204 | 199 |
|
205 | 200 | if unique_names[0] == sorted_pair[0]: |
@@ -550,7 +545,7 @@ def create_elo_plots(self, output_dir: Path) -> None: |
550 | 545 | player_order = [all_players[i] for i in all_indices] |
551 | 546 |
|
552 | 547 | # Translate to display names |
553 | | - display_names = [MODEL_TO_DISPLAY_NAME.get(p, p) for p in player_order] |
| 548 | + display_names = [model_display_name(p) for p in player_order] |
554 | 549 |
|
555 | 550 | # Create mapping from player to y-position |
556 | 551 | player_to_pos = {p: i for i, p in enumerate(player_order)} |
@@ -698,7 +693,7 @@ def create_validation_plots(self, output_dir: Path, regularization: float = 0.01 |
698 | 693 |
|
699 | 694 | ax.set_xlabel("BT Strength", fontproperties=FONT_BOLD, fontsize=12) |
700 | 695 | ax.set_ylabel("Negative Log-Likelihood", fontproperties=FONT_BOLD, fontsize=12) |
701 | | - display_name = MODEL_TO_DISPLAY_NAME.get(player, player) |
| 696 | + display_name = model_display_name(player) |
702 | 697 | ax.set_title(display_name, fontproperties=FONT_BOLD, fontsize=14) |
703 | 698 | legend = ax.legend(prop=FONT_BOLD, fontsize=10, loc="upper right") |
704 | 699 | legend.set_frame_on(False) |
@@ -777,7 +772,7 @@ def _create_rank_matrix_plot( |
777 | 772 | rank_matrix = (rank_matrix / self.n_bootstrap) * 100 |
778 | 773 |
|
779 | 774 | # Translate player names to display names |
780 | | - display_names = [MODEL_TO_DISPLAY_NAME.get(p, p) for p in players] |
| 775 | + display_names = [model_display_name(p) for p in players] |
781 | 776 |
|
782 | 777 | fig, ax = plt.subplots(figsize=(6, 6)) |
783 | 778 | im = ax.imshow(rank_matrix, cmap="YlOrRd", aspect="auto", vmin=0, vmax=100) |
@@ -826,7 +821,7 @@ def _create_elo_violin_plot( |
826 | 821 | elo_data = [elo_samples[p] for p in players] |
827 | 822 |
|
828 | 823 | # Translate player names to display names |
829 | | - display_names = [MODEL_TO_DISPLAY_NAME.get(p, p) for p in players] |
| 824 | + display_names = [model_display_name(p) for p in players] |
830 | 825 |
|
831 | 826 | fig, ax = plt.subplots(figsize=(6, 6)) |
832 | 827 |
|
@@ -1095,7 +1090,7 @@ def _plot_results(self, results_by_max_round: dict[int, dict[str, dict[str, floa |
1095 | 1090 | elos_list.append(results_by_max_round[max_round][game_name][player]) |
1096 | 1091 |
|
1097 | 1092 | if max_rounds_list: |
1098 | | - display_name = MODEL_TO_DISPLAY_NAME.get(player, player) |
| 1093 | + display_name = model_display_name(player) |
1099 | 1094 | ax.plot(max_rounds_list, elos_list, marker="o", label=display_name, linewidth=2, markersize=6) |
1100 | 1095 |
|
1101 | 1096 | ax.set_xlabel("Max Round", fontproperties=FONT_BOLD, fontsize=14) |
@@ -1212,7 +1207,7 @@ def _plot_results(self, results_by_round: dict[int, dict[str, dict[str, float]]] |
1212 | 1207 | elos_list.append(results_by_round[round_num][game_name][player]) |
1213 | 1208 |
|
1214 | 1209 | if rounds_list: |
1215 | | - display_name = MODEL_TO_DISPLAY_NAME.get(player, player) |
| 1210 | + display_name = model_display_name(player) |
1216 | 1211 | ax.plot(rounds_list, elos_list, marker="o", label=display_name, linewidth=2, markersize=6) |
1217 | 1212 |
|
1218 | 1213 | ax.set_xlabel("Round", fontproperties=FONT_BOLD, fontsize=14) |
@@ -1348,7 +1343,7 @@ def write_latex_table(results: dict[str, dict], output_dir: Path) -> None: |
1348 | 1343 | lines.append(r"\midrule") |
1349 | 1344 |
|
1350 | 1345 | for player, all_elo in sorted_players: |
1351 | | - display_name = MODEL_TO_DISPLAY_NAME.get(player, player) |
| 1346 | + display_name = model_display_name(player) |
1352 | 1347 | row_parts = [display_name.replace("_", r"\_")] |
1353 | 1348 |
|
1354 | 1349 | for game_name in games_in_table: |
@@ -1407,7 +1402,7 @@ def write_website_results(results: dict[str, dict], output_dir: Path) -> None: |
1407 | 1402 | # Create leaderboard entries |
1408 | 1403 | board = [] |
1409 | 1404 | for rank, (player, elo) in enumerate(sorted_players): |
1410 | | - entry = {"rank": rank + 1, "model": MODEL_TO_DISPLAY_NAME.get(player, player), "elo": int(round(elo))} |
| 1405 | + entry = {"rank": rank + 1, "model": model_display_name(player), "elo": int(round(elo))} |
1411 | 1406 | # Add confidence interval if available |
1412 | 1407 | if elo_std is not None: |
1413 | 1408 | player_idx = players.index(player) |
@@ -1506,7 +1501,7 @@ def write_latex_table_plain(results: dict[str, dict], output_dir: Path) -> None: |
1506 | 1501 | lines.append(r"\midrule") |
1507 | 1502 |
|
1508 | 1503 | for player, all_elo in sorted_players: |
1509 | | - display_name = MODEL_TO_DISPLAY_NAME.get(player, player) |
| 1504 | + display_name = model_display_name(player) |
1510 | 1505 | row_parts = [display_name.replace("_", r"\_")] |
1511 | 1506 |
|
1512 | 1507 | for game_name in games_in_table: |
|
0 commit comments