DedeProGames commited on
Commit
c8def62
·
verified ·
1 Parent(s): 7b0b288

Add match logic, Elo and results store

Browse files
Files changed (1) hide show
  1. arena.py +162 -0
arena.py ADDED
@@ -0,0 +1,162 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Match logic, Elo and the public results store (a HF dataset repo)."""
2
+ from __future__ import annotations
3
+
4
+ import io
5
+ import json
6
+ import os
7
+ import random
8
+ import threading
9
+ import uuid
10
+ from datetime import datetime, timezone
11
+
12
+ from huggingface_hub import CommitOperationAdd, HfApi, hf_hub_download
13
+
14
+ from tetris import TetrisGame
15
+
16
+ MAX_PIECES = int(os.environ.get("MAX_PIECES", 500))
17
+ K_FACTOR = 32
18
+ START_ELO = 1000.0
19
+ APP_VERSION = "1.0"
20
+
21
+
22
+ def choose(game: TetrisGame, player, protocol: str, seed: int):
23
+ """Let `player` pick the next placement. Returns False on top-out."""
24
+ cands = game.candidates()
25
+ if not cands:
26
+ game.top_out()
27
+ return False
28
+ vals = player.values(protocol, cands)
29
+ best = max(vals)
30
+ ties = [i for i, v in enumerate(vals) if v >= best - 1e-9]
31
+ # deterministic tie-break, identical for every player at the same piece index
32
+ rng = random.Random(f"{seed}:{game.pieces}")
33
+ i = ties[rng.randrange(len(ties))]
34
+ game.apply(cands[i], vals[i])
35
+ return True
36
+
37
+
38
+ def result_key(g: TetrisGame):
39
+ return (g.score, g.lines, g.pieces)
40
+
41
+
42
+ def rank_games(games):
43
+ """Competition ranking (1,1,3...) by score, then lines, then pieces survived."""
44
+ keys = [result_key(g) for g in games]
45
+ return [1 + sum(1 for k2 in keys if k2 > k) for k in keys]
46
+
47
+
48
+ def elo_deltas(ratings: list[float], keys: list[tuple]) -> list[float]:
49
+ """Multiplayer Elo: every pair is a mini-game, scaled by 1/(N-1)."""
50
+ n = len(ratings)
51
+ out = []
52
+ for i in range(n):
53
+ s = 0.0
54
+ for j in range(n):
55
+ if i == j:
56
+ continue
57
+ actual = 1.0 if keys[i] > keys[j] else 0.5 if keys[i] == keys[j] else 0.0
58
+ expected = 1.0 / (1.0 + 10 ** ((ratings[j] - ratings[i]) / 400.0))
59
+ s += actual - expected
60
+ out.append(K_FACTOR * s / (n - 1))
61
+ return out
62
+
63
+
64
+ class ResultsStore:
65
+ """Leaderboard persisted in a public dataset repo: leaderboard.json + one JSON per match."""
66
+
67
+ def __init__(self, repo_id: str, token: str | None):
68
+ self.repo_id = repo_id
69
+ self.token = token
70
+ self.api = HfApi(token=token)
71
+ self.lock = threading.Lock()
72
+ self.load_error = None
73
+ self.save_error = None
74
+ self.data = self._load()
75
+
76
+ @property
77
+ def persistent(self):
78
+ return bool(self.token)
79
+
80
+ def _empty(self):
81
+ return {"version": 1, "updated": None, "protocols": {"guided": {}, "blind": {}}}
82
+
83
+ def _load(self):
84
+ try:
85
+ path = hf_hub_download(self.repo_id, "leaderboard.json", repo_type="dataset", token=self.token, force_download=True)
86
+ data = json.load(open(path))
87
+ data.setdefault("protocols", {}).setdefault("guided", {})
88
+ data["protocols"].setdefault("blind", {})
89
+ self.load_error = None
90
+ return data
91
+ except Exception as e: # first run (no file yet) or Hub unreachable
92
+ self.load_error = f"leaderboard not loaded ({type(e).__name__}); starting empty"
93
+ return self._empty()
94
+
95
+ def reload(self):
96
+ with self.lock:
97
+ self.data = self._load()
98
+
99
+ def rating(self, protocol, model_id):
100
+ e = self.data["protocols"][protocol].get(model_id)
101
+ return e["elo"] if e else START_ELO
102
+
103
+ def record(self, protocol: str, seed: int, players, games) -> dict:
104
+ now = datetime.now(timezone.utc).isoformat(timespec="seconds")
105
+ with self.lock:
106
+ table = self.data["protocols"][protocol]
107
+ ratings = [self.rating(protocol, p.model_id) for p in players]
108
+ keys = [result_key(g) for g in games]
109
+ deltas = elo_deltas(ratings, keys)
110
+ ranks = rank_games(games)
111
+ match = {
112
+ "id": uuid.uuid4().hex[:12],
113
+ "time": now,
114
+ "protocol": protocol,
115
+ "seed": seed,
116
+ "max_pieces": MAX_PIECES,
117
+ "app_version": APP_VERSION,
118
+ "players": [],
119
+ }
120
+ for p, g, r, d, rank in zip(players, games, ratings, deltas, ranks):
121
+ e = table.setdefault(
122
+ p.model_id,
123
+ {"model": p.model_id, "elo": START_ELO, "games": 0, "wins": 0, "total_pieces": 0,
124
+ "total_lines": 0, "best_score": 0, "best_lines": 0},
125
+ )
126
+ e["elo"] = round(r + d, 2)
127
+ e["games"] += 1
128
+ e["wins"] += int(rank == 1)
129
+ e["total_pieces"] += g.pieces
130
+ e["total_lines"] += g.lines
131
+ e["best_score"] = max(e["best_score"], g.score)
132
+ e["best_lines"] = max(e["best_lines"], g.lines)
133
+ e["params"] = p.n_params
134
+ e["custom_code"] = p.custom_code
135
+ e["last_sha"] = p.sha
136
+ e["last_played"] = now
137
+ match["players"].append({
138
+ "model": p.model_id, "sha": p.sha, "params": p.n_params, "rank": rank,
139
+ "score": g.score, "lines": g.lines, "pieces": g.pieces, "tetrises": g.tetrises,
140
+ "survived": g.alive, "elo_before": round(r, 2), "elo_after": e["elo"], "elo_delta": round(d, 2),
141
+ })
142
+ self.data["updated"] = now
143
+ if self.persistent:
144
+ try:
145
+ day = now[:10]
146
+ self.api.create_commit(
147
+ repo_id=self.repo_id,
148
+ repo_type="dataset",
149
+ operations=[
150
+ CommitOperationAdd("leaderboard.json", io.BytesIO(json.dumps(self.data, indent=1).encode())),
151
+ CommitOperationAdd(f"matches/{day}/{match['id']}.json", io.BytesIO(json.dumps(match, indent=1).encode())),
152
+ ],
153
+ commit_message=f"{protocol} match {match['id']}: " + " vs ".join(p.model_id for p in players),
154
+ )
155
+ self.save_error = None
156
+ except Exception as e:
157
+ self.save_error = f"could not save to {self.repo_id} ({type(e).__name__}: {str(e)[:120]})"
158
+ return match
159
+
160
+ def rows(self, protocol: str):
161
+ entries = sorted(self.data["protocols"][protocol].values(), key=lambda e: -e["elo"])
162
+ return entries