Skip to content

Commit 96889de

Browse files
author
yolo h8cker 93
committed
Test faster matrix eval with parallel games
1 parent e4b0283 commit 96889de

1 file changed

Lines changed: 133 additions & 49 deletions

File tree

codeclash/analysis/matrix.py

Lines changed: 133 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,9 @@
22
import json
33
import time
44
import uuid
5+
from concurrent.futures import ThreadPoolExecutor, as_completed
56
from pathlib import Path
7+
from threading import Lock
68

79
from codeclash.agents.dummy_agent import Dummy
810
from codeclash.agents.utils import GameContext
@@ -16,17 +18,21 @@
1618

1719

1820
class PvPMatrixEvaluator:
19-
def __init__(self, pvp_output_dir: Path, n_repetitions: int = 3):
21+
def __init__(self, pvp_output_dir: Path, n_repetitions: int = 3, max_workers: int = 4):
2022
self.pvp_output_dir = Path(pvp_output_dir)
2123
self.n_repetitions = n_repetitions
24+
self.max_workers = max_workers
2225
self.metadata = json.loads((self.pvp_output_dir / "metadata.json").read_text())
2326

2427
assert len(self.players) == 2, f"Expected exactly 2 players, got {len(self.players)}"
2528

26-
# Set up logging
29+
# Set up logging with thread safety
2730
self.logger = get_logger("MatrixEvaluator", log_path=self.pvp_output_dir / "matrix_eval.log", emoji="📊")
2831
add_file_handler(get_logger("."), self.pvp_output_dir / "matrix_eval.log")
2932

33+
# Thread safety for progress saving
34+
self._save_lock = Lock()
35+
3036
# Initialize metadata similar to tournament class
3137
self._metadata = {
3238
"name": "MatrixEvaluator",
@@ -35,27 +41,21 @@ def __init__(self, pvp_output_dir: Path, n_repetitions: int = 3):
3541
"p2_name": self.players[1],
3642
"rounds": self.rounds,
3743
"n_repetitions": n_repetitions,
44+
"max_workers": max_workers,
3845
"created_timestamp": int(time.time()),
3946
"matrices": {},
4047
}
4148

4249
# Load existing progress if available
4350
self._load_existing_progress()
4451

45-
# Create game instance for evaluation
46-
tournament_id = f"MatrixEval.{self.metadata['name']}.{time.strftime('%y%m%d%H%M%S')}"
47-
48-
self.config["game"]["sims_per_round"] = n_repetitions
49-
50-
self.game = get_game(
51-
self.config,
52-
tournament_id=tournament_id,
53-
local_output_dir=self.pvp_output_dir / "matrix_eval",
54-
keep_containers=False,
55-
)
52+
# Initialize game pool and agent pools
53+
self._initialize_game_pool()
54+
self._initialize_agent_pools()
5655

5756
self.logger.info(f"Initialized matrix evaluator for {self.players[0]} vs {self.players[1]}")
5857
self.logger.info(f"Will evaluate {self.rounds + 1} rounds with {n_repetitions} repetitions each")
58+
self.logger.info(f"Using {max_workers} parallel workers")
5959

6060
# Quick access properties
6161
# -----------------------
@@ -96,9 +96,48 @@ def _load_existing_progress(self):
9696
self._metadata["matrices"] = existing_data["matrices"]
9797

9898
def _save_progress(self):
99-
"""Save current progress to matrix.json."""
100-
self.output_file.write_text(json.dumps(self._metadata, indent=2))
101-
self.logger.debug("Progress saved to matrix.json")
99+
"""Save current progress to matrix.json in a thread-safe manner."""
100+
with self._save_lock:
101+
self.output_file.write_text(json.dumps(self._metadata, indent=2))
102+
self.logger.debug("Progress saved to matrix.json")
103+
104+
def _initialize_game_pool(self):
105+
"""Initialize a pool of game objects for parallel execution."""
106+
self.game_pool = []
107+
for i in range(self.max_workers):
108+
tournament_id = f"MatrixEval.{self.metadata['name']}.{time.strftime('%y%m%d%H%M%S')}.worker{i}"
109+
config = self.config.copy()
110+
config["game"]["sims_per_round"] = self.n_repetitions
111+
112+
game = get_game(
113+
config,
114+
tournament_id=tournament_id,
115+
local_output_dir=self.pvp_output_dir / "matrix_eval" / f"worker_{i}",
116+
keep_containers=False,
117+
)
118+
self.game_pool.append(game)
119+
120+
self.logger.info(f"Initialized {len(self.game_pool)} game workers")
121+
122+
def _initialize_agent_pools(self):
123+
"""Pre-initialize agents for all rounds for each player."""
124+
self.agent_pools = {}
125+
126+
for player_name in self.players:
127+
self.agent_pools[player_name] = {}
128+
for round_num in range(self.rounds + 1):
129+
# Pre-load the diff for this round
130+
patch = self._get_round_diff(player_name, round_num)
131+
if patch is not None:
132+
# Create agent for this round and player
133+
agent = self._create_dummy_agent(player_name, f"_r{round_num}")
134+
agent.reset_and_apply_patch(filter_git_diff(patch))
135+
self.agent_pools[player_name][round_num] = agent
136+
self.logger.debug(f"Pre-initialized agent for {player_name} round {round_num}")
137+
else:
138+
self.logger.warning(f"Missing changes file for {player_name} round {round_num}")
139+
140+
self.logger.info(f"Pre-initialized agents for {len(self.players)} players across {self.rounds + 1} rounds")
102141

103142
def _get_round_diff(self, player_name: str, round_num: int) -> str | None:
104143
"""Read diff data from changes_r{round}.json file. Returns None if file doesn't exist."""
@@ -141,63 +180,98 @@ def _create_dummy_agent(self, player_name: str, agent_suffix: str = "") -> Dummy
141180

142181
return Dummy(original_config, environment, game_context)
143182

144-
def _evaluate_matrix_cell(
145-
self, agent1: Dummy, agent2: Dummy, player1_name: str, player2_name: str, i: int, j: int, matrix_id: str
146-
) -> dict | None:
147-
"""Evaluate a single matrix cell and return the stats object. Returns None if cell should be skipped."""
183+
def _evaluate_matrix_cell_parallel(
184+
self, game_worker, player1_name: str, player2_name: str, i: int, j: int, matrix_id: str
185+
) -> tuple[int, int, dict | None]:
186+
"""Evaluate a single matrix cell using pre-initialized agents. Returns (i, j, result)."""
148187
# Return existing result if already completed
149188
try:
150189
existing_result = self.matrices[matrix_id][str(i)][str(j)]
190+
if existing_result:
191+
self.logger.debug(f"Skipping {player1_name} round {i} vs {player2_name} round {j} - already completed")
192+
return (i, j, existing_result)
151193
except KeyError:
152-
existing_result = None
153-
if existing_result:
154-
self.logger.debug(f"Skipping {player1_name} round {i} vs {player2_name} round {j} - already completed")
155-
return existing_result
194+
pass
156195

157-
patch1 = self._get_round_diff(player1_name, i)
158-
patch2 = self._get_round_diff(player2_name, j)
159-
160-
# Skip if any required changes file is missing
161-
if patch1 is None:
196+
# Check if agents are available for these rounds
197+
if i not in self.agent_pools[player1_name]:
162198
self.logger.warning(
163-
f"Skipping {player1_name} round {i} vs {player2_name} round {j} - missing changes file for {player1_name} round {i}"
199+
f"Skipping {player1_name} round {i} vs {player2_name} round {j} - missing agent for {player1_name} round {i}"
164200
)
165-
return None
166-
if patch2 is None:
201+
return (i, j, None)
202+
if j not in self.agent_pools[player2_name]:
167203
self.logger.warning(
168-
f"Skipping {player1_name} round {i} vs {player2_name} round {j} - missing changes file for {player2_name} round {j}"
204+
f"Skipping {player1_name} round {i} vs {player2_name} round {j} - missing agent for {player2_name} round {j}"
169205
)
170-
return None
206+
return (i, j, None)
171207

172-
agent1.reset_and_apply_patch(filter_git_diff(patch1))
173-
agent2.reset_and_apply_patch(filter_git_diff(patch2))
208+
# Get pre-initialized agents
209+
agent1 = self.agent_pools[player1_name][i]
210+
agent2 = self.agent_pools[player2_name][j]
174211

175212
self.logger.info(f"Evaluating {player1_name} round {i} vs {player2_name} round {j}")
176213

177214
round_id = str(uuid.uuid4().hex)
178-
stats = self.game.run_round([agent1, agent2], round_id)
215+
stats = game_worker.run_round([agent1, agent2], round_id)
179216
self.logger.debug(f"Result: {stats.to_dict()}")
180217

181-
return stats.to_dict()
218+
return (i, j, stats.to_dict())
182219

183220
def _evaluate_matrix(self, player1_name: str, player2_name: str):
184-
"""Generic method to evaluate a matrix between two players (or same player)."""
221+
"""Evaluate a matrix between two players using parallel execution."""
185222
symmetric = player1_name == player2_name
186223
matrix_id = f"{player1_name}_vs_{player2_name}"
187224
self.logger.info(f"Evaluating {matrix_id} matrix: {player1_name} vs {player2_name}")
188225

189-
agent1 = self._create_dummy_agent(player1_name, "_1" if player1_name == player2_name else "")
190-
agent2 = self._create_dummy_agent(player2_name, "_2" if player1_name == player2_name else "")
191-
226+
# Initialize matrix structure
192227
self.matrices.setdefault(matrix_id, {})
193228
for i in range(self.rounds + 1):
194229
self.matrices[matrix_id].setdefault(str(i), {})
230+
231+
# Collect all tasks to be executed
232+
tasks = []
233+
game_worker_index = 0
234+
235+
for i in range(self.rounds + 1):
195236
j_range = range(i + 1) if symmetric else range(self.rounds + 1)
196237
for j in j_range:
197-
result = self._evaluate_matrix_cell(agent1, agent2, player1_name, player2_name, i, j, matrix_id)
238+
# Skip if already completed
239+
try:
240+
if self.matrices[matrix_id][str(i)][str(j)]:
241+
continue
242+
except KeyError:
243+
pass
244+
245+
# Assign game worker in round-robin fashion
246+
game_worker = self.game_pool[game_worker_index % len(self.game_pool)]
247+
game_worker_index += 1
248+
249+
tasks.append((game_worker, player1_name, player2_name, i, j, matrix_id))
250+
251+
if not tasks:
252+
self.logger.info(f"All matrix cells for {matrix_id} already completed")
253+
return
254+
255+
self.logger.info(f"Executing {len(tasks)} matrix cells in parallel with {self.max_workers} workers")
256+
257+
# Execute tasks in parallel
258+
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
259+
# Submit all tasks
260+
future_to_task = {executor.submit(self._evaluate_matrix_cell_parallel, *task): task for task in tasks}
261+
262+
# Process completed tasks as they finish
263+
for future in as_completed(future_to_task):
264+
i, j, result = future.result()
198265
if result is not None:
199-
self.matrices[matrix_id][str(i)][str(j)] = result
200-
self._save_progress()
266+
with self._save_lock:
267+
self.matrices[matrix_id][str(i)][str(j)] = result
268+
self._save_progress()
269+
270+
# Log progress
271+
completed = len([f for f in future_to_task if f.done()])
272+
self.logger.info(f"Progress: {completed}/{len(tasks)} matrix cells completed for {matrix_id}")
273+
274+
self.logger.info(f"Completed matrix evaluation for {matrix_id}")
201275

202276
def evaluate_all_matrices(self) -> dict:
203277
"""Evaluate vs matrix between the two players."""
@@ -212,12 +286,21 @@ def end(self):
212286
"""Save metadata and clean up resources."""
213287
self.output_file.write_text(json.dumps(self._metadata, indent=2))
214288
self.logger.info(f"Matrix evaluation results saved to {self.output_file}")
215-
self.game.end(cleanup=True)
289+
290+
# Clean up all game workers
291+
for i, game in enumerate(self.game_pool):
292+
try:
293+
game.end(cleanup=True)
294+
self.logger.debug(f"Cleaned up game worker {i}")
295+
except Exception as e:
296+
self.logger.warning(f"Error cleaning up game worker {i}: {e}")
297+
298+
self.logger.info("All game workers cleaned up")
216299

217300

218-
def main(pvp_output_dir: Path, n_repetitions: int = 3):
301+
def main(pvp_output_dir: Path, n_repetitions: int = 3, max_workers: int = 4):
219302
"""Main function to evaluate PvP tournament matrices."""
220-
evaluator = PvPMatrixEvaluator(pvp_output_dir, n_repetitions)
303+
evaluator = PvPMatrixEvaluator(pvp_output_dir, n_repetitions, max_workers)
221304
return evaluator.evaluate_all_matrices()
222305

223306

@@ -227,6 +310,7 @@ def main(pvp_output_dir: Path, n_repetitions: int = 3):
227310
parser.add_argument(
228311
"--repetitions", "-r", type=int, default=3, help="Number of repetitions per matrix cell (default: 3)"
229312
)
313+
parser.add_argument("--max-workers", "-w", type=int, default=4, help="Number of parallel game workers (default: 4)")
230314

231315
args = parser.parse_args()
232-
main(args.pvp_output_dir, args.repetitions)
316+
main(args.pvp_output_dir, args.repetitions, args.max_workers)

0 commit comments

Comments
 (0)