2121import torch .nn .functional as F
2222from torch .utils .data import DataLoader
2323
24+ import chess_engine as engine
2425from pawn .config import CLMConfig , TrainingConfig
2526from pawn .model import PAWNCLM
26- from pawn .data import CLMDataset , create_validation_set
27+ from pawn .data import CLMDataset , create_validation_set , shift_legal_mask
2728from pawn .logging import MetricsLogger
2829
2930from pawn .data_utils import unpack_grid
@@ -260,26 +261,25 @@ def compute_game_completion(
260261 Returns dict with:
261262 game_completion_rate: fraction of games with zero illegal moves
262263 avg_pct_completion: mean fraction of game completed before forfeit
263- avg_plies_to_forfeit: mean plies before first illegal move (inf-free)
264+ avg_plies_completed: mean plies completed before first illegal move.
265+ Games with no illegal moves contribute their full game_length.
264266 """
265267 B , T = preds .shape
266268
267269 with torch .no_grad ():
268270 n_complete = 0
269271 pct_completions = []
270- plies_to_forfeit = []
272+ plies_completed = []
271273
272274 for b in range (B ):
273275 gl = min (int (game_lengths [b ].item ()), T )
274276 forfeit_ply = - 1
275- n_checked = 0
276277 for p in range (gl ):
277278 if not loss_mask [b , p ]:
278279 continue
279280 # Skip plies with no legal moves (end-of-game padding)
280281 if not legal_mask [b , p ].any ():
281282 continue
282- n_checked += 1
283283 token = int (preds [b , p ].item ())
284284 if token < legal_mask .shape [2 ] and not legal_mask [b , p , token ]:
285285 forfeit_ply = p
@@ -291,15 +291,15 @@ def compute_game_completion(
291291 if forfeit_ply < 0 :
292292 n_complete += 1
293293 pct_completions .append (1.0 )
294- plies_to_forfeit .append (float (gl ))
294+ plies_completed .append (float (gl ))
295295 else :
296296 pct_completions .append (forfeit_ply / gl if gl > 0 else 0.0 )
297- plies_to_forfeit .append (float (forfeit_ply ))
297+ plies_completed .append (float (forfeit_ply ))
298298
299299 return {
300300 "game_completion_rate" : n_complete / B if B > 0 else 0.0 ,
301301 "avg_pct_completion" : sum (pct_completions ) / len (pct_completions ) if pct_completions else 0.0 ,
302- "avg_plies_to_forfeit " : sum (plies_to_forfeit ) / len (plies_to_forfeit ) if plies_to_forfeit else 0.0 ,
302+ "avg_plies_completed " : sum (plies_completed ) / len (plies_completed ) if plies_completed else 0.0 ,
303303 }
304304
305305
@@ -582,10 +582,8 @@ def evaluate(self) -> dict[str, float]:
582582 # without picking an illegal move? Uses a small subset to avoid
583583 # materializing a large dense (B, T, vocab) token mask.
584584 if "game_lengths" in self .val_data :
585- import chess_engine as engine_mod
586585 gc_n = min (64 , n )
587586 gc_input = self .val_data ["input_ids" ][:gc_n ].to (self .device )
588- gc_targets = self .val_data ["targets" ][:gc_n ].to (self .device )
589587 gc_loss_mask = self .val_data ["loss_mask" ][:gc_n ].to (self .device )
590588 gc_game_lengths = self .val_data ["game_lengths" ][:gc_n ].to (self .device )
591589 move_ids = self .val_data ["input_ids" ][:gc_n ].numpy ().astype (np .int16 )
@@ -598,18 +596,17 @@ def evaluate(self) -> dict[str, float]:
598596 gc_logits = model .lm_head (hidden )
599597 gc_preds = gc_logits .argmax (dim = - 1 )
600598
601- # Dense legal token mask, shifted to align with targets
602- legal_tokens = engine_mod .compute_legal_token_masks (move_ids , gl_np , vocab_size )
603- legal_tokens = np .roll (legal_tokens , - 1 , axis = 1 )
604- legal_tokens [:, - 1 , :] = False
605- legal_mask_t = torch .from_numpy (legal_tokens ).to (self .device )
599+ legal_tokens = engine .compute_legal_token_masks (move_ids , gl_np , vocab_size )
600+ legal_mask_t = torch .from_numpy (
601+ shift_legal_mask (legal_tokens )
602+ ).to (self .device )
606603
607604 gc = compute_game_completion (gc_preds , legal_mask_t , gc_loss_mask , gc_game_lengths )
608605 avg ["val/game_completion_rate" ] = gc ["game_completion_rate" ]
609606 avg ["val/avg_pct_completion" ] = gc ["avg_pct_completion" ]
610- avg ["val/avg_plies_to_forfeit " ] = gc ["avg_plies_to_forfeit " ]
607+ avg ["val/avg_plies_completed " ] = gc ["avg_plies_completed " ]
611608
612- del legal_mask_t , gc_logits , gc_preds
609+ del gc_input , gc_loss_mask , gc_game_lengths , legal_mask_t , gc_logits , gc_preds
613610 if self .device != "cpu" and torch .cuda .is_available ():
614611 torch .cuda .empty_cache ()
615612
@@ -705,7 +702,7 @@ def _graceful_exit(signum, frame):
705702 if "val/game_completion_rate" in val_metrics :
706703 val_msg += (
707704 f" | complete { val_metrics ['val/game_completion_rate' ]:.3f} "
708- f" | avg_ply { val_metrics ['val/avg_plies_to_forfeit ' ]:.0f} "
705+ f" | avg_ply { val_metrics ['val/avg_plies_completed ' ]:.0f} "
709706 )
710707
711708 # Compound early stopping
0 commit comments