""" Batching helpers: python-chess boards -> model tensors. Contract (critical for correctness, tested in tests/test_phase1.py): * Every board is canonicalised to side-to-move == White. * Candidate moves are generated ON THE CANONICAL BOARD, so index i of the logits always corresponds to `cand_moves[b][i]` in canonical space, and to `real_moves[b][i]` on the original board. * The value head therefore always predicts from the SIDE-TO-MOVE perspective. """ from __future__ import annotations from typing import Sequence import chess import numpy as np import torch from .encoding import (encode_position, pack_move, unpack_fields, canonical_board, flip_move) class PositionBatch: """Tensors for a batch of positions plus the move bookkeeping.""" __slots__ = ("squares", "extras", "glob", "cand_fields", "cand_mask", "real_moves", "n_legal") def __init__(self, squares, extras, glob, cand_fields, cand_mask, real_moves, n_legal): self.squares = squares self.extras = extras self.glob = glob self.cand_fields = cand_fields self.cand_mask = cand_mask self.real_moves = real_moves # list[list[chess.Move]] in ORIGINAL board space self.n_legal = n_legal def to(self, device): self.squares = self.squares.to(device) self.extras = self.extras.to(device) self.glob = self.glob.to(device) self.cand_fields = self.cand_fields.to(device) self.cand_mask = self.cand_mask.to(device) return self def model_args(self): return dict(squares=self.squares, extras=self.extras, glob=self.glob, cand_fields=self.cand_fields, cand_mask=self.cand_mask) def encode_batch(boards: Sequence[chess.Board], last_moves: Sequence[chess.Move | None] | None = None, max_cand: int | None = None, reps: Sequence[int] | None = None) -> PositionBatch: """Encode boards + their legal candidate sets into padded tensors.""" n = len(boards) if last_moves is None: last_moves = [None] * n sq_l, ex_l, g_l, packed_l, real_l = [], [], [], [], [] for bi, (b, lm) in enumerate(zip(boards, last_moves)): if reps is not None: rep = int(reps[bi]) # explicit: FEN-restored boards lack history else: rep = 0 try: if b.is_repetition(2): rep = 2 except Exception: rep = 0 sq, ex, g, flipped = encode_position(b, lm, rep) cb = b.mirror() if flipped else b real, packed = [], [] for m in cb.legal_moves: packed.append(pack_move(cb, m)) real.append(flip_move(m) if flipped else m) sq_l.append(sq); ex_l.append(ex); g_l.append(g) packed_l.append(packed); real_l.append(real) L = max((len(p) for p in packed_l), default=1) if max_cand is not None: L = min(L, max_cand) L = max(L, 1) cand = np.zeros((n, L, 7), dtype=np.int64) mask = np.zeros((n, L), dtype=bool) n_legal = np.zeros(n, dtype=np.int64) for i, p in enumerate(packed_l): k = min(len(p), L) n_legal[i] = k if k: cand[i, :k] = unpack_fields(np.asarray(p[:k], dtype=np.int64)) mask[i, :k] = True real_l[i] = real_l[i][:k] return PositionBatch( squares=torch.from_numpy(np.stack(sq_l)).long(), extras=torch.from_numpy(np.stack(ex_l)).float(), glob=torch.from_numpy(np.stack(g_l)).float(), cand_fields=torch.from_numpy(cand), cand_mask=torch.from_numpy(mask), real_moves=real_l, n_legal=torch.from_numpy(n_legal), )