File size: 4,193 Bytes
ba69de3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""

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