22import json
33import time
44import uuid
5+ from concurrent .futures import ThreadPoolExecutor , as_completed
56from pathlib import Path
7+ from threading import Lock
68
79from codeclash .agents .dummy_agent import Dummy
810from codeclash .agents .utils import GameContext
1618
1719
1820class 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