system's picture
system HF Staff
Deploy from 1040f46
a5f2e46 verified
Raw
History Blame Contribute Delete
3.13 kB
import chess
_PIECE_TOKEN = {
(chess.PAWN, chess.WHITE): 1, (chess.KNIGHT, chess.WHITE): 2,
(chess.BISHOP, chess.WHITE): 3, (chess.ROOK, chess.WHITE): 4,
(chess.QUEEN, chess.WHITE): 5, (chess.KING, chess.WHITE): 6,
(chess.PAWN, chess.BLACK): 7, (chess.KNIGHT, chess.BLACK): 8,
(chess.BISHOP, chess.BLACK): 9, (chess.ROOK, chess.BLACK): 10,
(chess.QUEEN, chess.BLACK): 11, (chess.KING, chess.BLACK): 12,
}
class PostionTokenizer:
"""A simple tokenizer for chess board positions.
Each square on the board is represented by a token ID based on the piece occupying it.
Empty squares are represented by the token ID 0.
The board is represented as an 8x8 grid, flattened into a list of 64 tokens.
"""
def __init__(self):
self.vocab = {
"P": 1,
"N": 2,
"B": 3,
"R": 4,
"Q": 5,
"K": 6,
"p": 7,
"n": 8,
"b": 9,
"r": 10,
"q": 11,
"k": 12,
".": 0,
}
self.inv_vocab = {v: k for k, v in self.vocab.items()}
def encode(self, board: chess.Board) -> list[int]:
"""Encode a chess.Board object into a list of token IDs.
The board is represented as an 8x8 grid, flattened into a list of 64 tokens.
Args:
board: A chess.Board object
Returns:
List of token IDs representing the board position
"""
result = [0] * 64
for sq, piece in board.piece_map().items():
result[sq] = _PIECE_TOKEN[(piece.piece_type, piece.color)]
return result
def decode(self, token_ids: list[int]) -> chess.Board:
"""Decode a list of token IDs back into a chess board string.
Args:
token_ids: List of token IDs representing the board position
Returns:
A chess.Board object
"""
chars = [self.inv_vocab[token_id] for token_id in token_ids]
rows = ["".join(chars[i * 8 : (i + 1) * 8]) for i in range(8)][::-1] # Reverse to get the correct order
board_str = "\n".join(rows)
board = self._ascii2board(board_str)
return board
def _ascii2board(self, ascii_board: str) -> chess.Board:
"""Convert an ASCII board representation back to a chess.Board object.
Args:
ascii_board: String representation of the board (8 lines of 8 characters)
Returns:
A chess.Board object
"""
board = chess.Board.empty()
rows = ascii_board.split("\n")
for rank in range(8):
file = 0
for char in rows[7 - rank]: # Reverse the order of ranks
if char in self.vocab and char != ".":
square = chess.square(file, rank)
piece = chess.Piece.from_symbol(char)
board.set_piece_at(square, piece)
file += 1
return board
@property
def vocab_size(self) -> int:
"""Return the size of the vocabulary."""
return len(self.vocab)