ChessResNet-30M / common_chess.py
Joeyfully's picture
Upload 5 files
ba69de3 verified
Raw History Blame Contribute Delete
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