training / tinychess /batching.py
cazyundee's picture
fix: repetition feature lost in FEN round-trip (train/inference mismatch); greedy-health exp
79997ea verified
Raw History Blame Contribute Delete
3.79 kB
"""
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),
)