Download common_chess.py from Joeyfully/ChessResNet-30M: direct link, hf CLI and curl.
- Browser
- Download file 4.19 kB
-
https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/common_chess.py
- Command line
-
hf download hf://Joeyfully/ChessResNet-30M/common_chess.py
-
curl -L -o common_chess.py https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/common_chess.py
4.19 kB
| """ | |
| Common chess encoding utilities shared across training, evaluation, and UCI engine. | |
| Must match the encoding used in `scripts/generate_data/label_stockfish_dataset.py`. | |
| """ | |
| import chess | |
| import numpy as np | |
| NUM_ACTIONS = 64 * 64 * 5 # 20480: canonical 64×64 squares × 5 promotion types | |
| NUM_BOARD_PLANES = 18 # 6 own pieces + 6 opponent + 4 castling + 1 ep + 1 side | |
| def encode_board(board: chess.Board) -> np.ndarray: | |
| """Encode a chess.Board into 18 canonical planes, shape (18, 8, 8), uint8. | |
| If black is to move, squares are mirrored vertically so that the side | |
| to move is always at the bottom (ranks 0-3 own, ranks 4-7 opponent). | |
| Planes: | |
| 0-5 : own pawn, knight, bishop, rook, queen, king | |
| 6-11 : opponent pawn, knight, bishop, rook, queen, king | |
| 12 : own kingside castling right (all-1 plane) | |
| 13 : own queenside castling right (all-1 plane) | |
| 14 : opponent kingside castling right (all-1 plane) | |
| 15 : opponent queenside castling right (all-1 plane) | |
| 16 : en-passant target square (1-hot) | |
| 17 : original side-to-move (1 = white, 0 = black) | |
| """ | |
| planes = np.zeros((18, 8, 8), dtype=np.uint8) | |
| turn = board.turn | |
| mirror = not turn # mirror squares if black to move | |
| # Piece planes (own = 0-5, opponent = 6-11) | |
| for sq in chess.SQUARES: | |
| piece = board.piece_at(sq) | |
| if piece is None: | |
| continue | |
| csq = chess.square_mirror(sq) if mirror else sq | |
| row, col = divmod(csq, 8) | |
| if piece.color == turn: | |
| idx = piece.piece_type - 1 # 0-5 own | |
| else: | |
| idx = piece.piece_type - 1 + 6 # 6-11 opponent | |
| planes[idx, row, col] = 1 | |
| # Castling rights (own / opponent perspective) | |
| if board.has_kingside_castling_rights(turn): | |
| planes[12, :, :] = 1 | |
| if board.has_queenside_castling_rights(turn): | |
| planes[13, :, :] = 1 | |
| if board.has_kingside_castling_rights(not turn): | |
| planes[14, :, :] = 1 | |
| if board.has_queenside_castling_rights(not turn): | |
| planes[15, :, :] = 1 | |
| # En-passant | |
| ep = board.ep_square | |
| if ep is not None: | |
| cep = chess.square_mirror(ep) if mirror else ep | |
| row, col = divmod(cep, 8) | |
| planes[16, row, col] = 1 | |
| # Original side-to-move indicator | |
| if turn == chess.WHITE: | |
| planes[17, :, :] = 1 | |
| return planes | |
| def move_to_action_id(move: chess.Move, turn: chess.Color) -> int: | |
| """Convert a chess.Move to canonical action id [0, 20480). | |
| If black to move, from_square and to_square are mirrored vertically | |
| so the encoding is invariant under board orientation. | |
| """ | |
| if turn == chess.BLACK: | |
| from_sq = chess.square_mirror(move.from_square) | |
| to_sq = chess.square_mirror(move.to_square) | |
| else: | |
| from_sq = move.from_square | |
| to_sq = move.to_square | |
| promo = move.promotion | |
| if promo is None: | |
| pid = 0 | |
| elif promo == chess.QUEEN: | |
| pid = 1 | |
| elif promo == chess.ROOK: | |
| pid = 2 | |
| elif promo == chess.BISHOP: | |
| pid = 3 | |
| elif promo == chess.KNIGHT: | |
| pid = 4 | |
| else: | |
| pid = 0 # should never happen | |
| return (from_sq * 64 + to_sq) * 5 + pid | |
| def action_id_to_move(action_id: int, turn: chess.Color) -> chess.Move: | |
| """Inverse of move_to_action_id.""" | |
| pid = action_id % 5 | |
| raw = action_id // 5 | |
| to_sq = raw % 64 | |
| from_sq = raw // 64 | |
| if turn == chess.BLACK: | |
| from_sq = chess.square_mirror(from_sq) | |
| to_sq = chess.square_mirror(to_sq) | |
| promo_map = {0: None, 1: chess.QUEEN, 2: chess.ROOK, | |
| 3: chess.BISHOP, 4: chess.KNIGHT} | |
| return chess.Move(from_sq, to_sq, promotion=promo_map[pid]) | |
| def legal_action_ids(board: chess.Board): | |
| """Return (action_ids, moves) for all legal moves on *board*. | |
| action_ids: list of int length len(moves) | |
| moves: list of chess.Move | |
| """ | |
| moves = list(board.legal_moves) | |
| action_ids = [move_to_action_id(m, board.turn) for m in moves] | |
| return action_ids, moves | |