4141from matplotlib import pyplot as plt
4242from tqdm .auto import tqdm
4343
44+ from codeclash .analysis .viz .utils import ASSETS_DIR , MODEL_TO_COLOR , MODEL_TO_DISPLAY_NAME
4445from codeclash .constants import LOCAL_LOG_DIR
4546
46- OUTPUT_FILE = "cdf_command_diversity.png"
47+ OUTPUT_FILE = ASSETS_DIR / "cdf_command_diversity.png"
48+ DATA_CACHE = ASSETS_DIR / "cdf_command_diversity.json"
4749
4850
4951def shannon_entropy (command_list ):
@@ -123,42 +125,49 @@ def main():
123125 """
124126 model_to_diversity = {}
125127
126- # Find all tournament directories by looking for metadata.json files
127- tournaments = [x .parent for x in LOCAL_LOG_DIR .rglob ("metadata.json" )]
128- for game_log_folder in tqdm (tournaments ):
129- # Load tournament metadata to get player-to-model mapping
130- with open (game_log_folder / "metadata.json" ) as f :
131- metadata = json .load (f )
132- try :
133- # Extract mapping from player name to model name
134- p2m = {x ["name" ]: x ["config" ]["model" ]["model_name" ].strip ("@" ) for x in metadata ["config" ]["players" ]}
135- # Initialize diversity list for each model we encounter
136- for model in p2m .values ():
137- if model not in model_to_diversity :
138- model_to_diversity [model ] = []
139- except KeyError :
140- # Skip tournaments with malformed metadata
141- continue
142-
143- # Process each player's trajectory files
144- for name in p2m .keys ():
145- traj_files = (game_log_folder / "players" / name ).rglob ("*.traj.json" )
146- for traj_file in traj_files :
147- try :
148- with open (traj_file ) as f :
149- traj = json .load (f )
150-
151- # Extract commands and calculate diversity for this session
152- commands = extract_commands_from_trajectory (traj )
153- if commands : # Only calculate entropy if there are commands
154- diversity = shannon_entropy (commands )
155- model_to_diversity [p2m [name ]].append (diversity )
156- except (json .JSONDecodeError , KeyError ):
157- # Skip malformed trajectory files
158- continue
159-
160- # Remove models with no valid data
161- model_to_diversity = {k : v for k , v in model_to_diversity .items () if v }
128+ if not DATA_CACHE .exists ():
129+ # Find all tournament directories by looking for metadata.json files
130+ tournaments = [x .parent for x in LOCAL_LOG_DIR .rglob ("metadata.json" )]
131+ for game_log_folder in tqdm (tournaments ):
132+ # Load tournament metadata to get player-to-model mapping
133+ with open (game_log_folder / "metadata.json" ) as f :
134+ metadata = json .load (f )
135+ try :
136+ # Extract mapping from player name to model name
137+ p2m = {x ["name" ]: x ["config" ]["model" ]["model_name" ].strip ("@" ) for x in metadata ["config" ]["players" ]}
138+ # Initialize diversity list for each model we encounter
139+ for model in p2m .values ():
140+ if model not in model_to_diversity :
141+ model_to_diversity [model ] = []
142+ except KeyError :
143+ # Skip tournaments with malformed metadata
144+ continue
145+
146+ # Process each player's trajectory files
147+ for name in p2m .keys ():
148+ traj_files = (game_log_folder / "players" / name ).rglob ("*.traj.json" )
149+ for traj_file in traj_files :
150+ try :
151+ with open (traj_file ) as f :
152+ traj = json .load (f )
153+
154+ # Extract commands and calculate diversity for this session
155+ commands = extract_commands_from_trajectory (traj )
156+ if commands : # Only calculate entropy if there are commands
157+ diversity = shannon_entropy (commands )
158+ model_to_diversity [p2m [name ]].append (diversity )
159+ except (json .JSONDecodeError , KeyError ):
160+ # Skip malformed trajectory files
161+ continue
162+
163+ # Remove models with no valid data
164+ model_to_diversity = {k : v for k , v in model_to_diversity .items () if v }
165+
166+ with open (DATA_CACHE , "w" ) as f :
167+ json .dump (model_to_diversity , f , indent = 2 )
168+
169+ with open (DATA_CACHE ) as f :
170+ model_to_diversity = json .load (f )
162171
163172 # Print summary statistics for each model
164173 print ("Command Diversity Summary:" )
@@ -170,18 +179,19 @@ def main():
170179 )
171180
172181 # Generate CDF plot comparing all models
173- plt .figure (figsize = (12 , 8 ))
174- colors = plt .cm .tab10 (range (len (model_to_diversity )))
182+ plt .figure (figsize = (8 , 8 ))
175183
176184 for i , (model , diversities ) in enumerate (model_to_diversity .items ()):
177185 # Sort diversity values and create cumulative probability values
178186 sorted_diversities = sorted (diversities )
179187 yvals = [i / len (sorted_diversities ) for i in range (len (sorted_diversities ))]
180188 # Plot as step function (standard for CDFs)
181- plt .step (sorted_diversities , yvals , label = model , where = "post" , color = colors [i ])
189+ plt .step (
190+ sorted_diversities , yvals , label = MODEL_TO_DISPLAY_NAME [model ], where = "post" , color = MODEL_TO_COLOR [model ]
191+ )
182192
183193 plt .xlabel ("Command Diversity (Shannon Entropy)" )
184- plt .ylabel ("Cumulative Probability" )
194+ # plt.ylabel("Cumulative Probability")
185195 plt .title ("CDF of Command Diversity by Model\n (Higher entropy = more diverse command usage)" )
186196 plt .legend (bbox_to_anchor = (1.05 , 1 ), loc = "upper left" ) # Legend outside plot area
187197 plt .grid (True , alpha = 0.3 )
0 commit comments