Download tinychess/batching.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 3.79 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/batching.py
- Command line
-
hf download hf://spaces/cazyundee/training/tinychess/batching.py
-
curl -L -o batching.py https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/batching.py
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), | |
| ) | |