import math from safetensors.numpy import load_file import numpy as np import chess class ChessNeuralNet: def __init__(self, name: str): self.FEATURE_SIZE = 768 self.HIDDEN_SIZE = 1600 self.INPUT_BUCKETS = 14 self.OUTPUT_BUCKETS = 8 self.INPUT_BUCKET_LAYOUT = [ 0, 1, 2, 3, 4, 5, 6, 7, 8, 8, 9, 9, 10, 10, 11, 11, 10, 10, 11, 11, 12, 12, 13, 13, 12, 12, 13, 13, 12, 12, 13, 13, ] self.QA = 255 self.QB = 64 self.SCALE = 400 self.feature_weights, self.feature_bias, self.output_weights, self.output_bias = self.load_weights(name) @staticmethod def mirror_square_as_needed(square: int, needs_vertical_mirror: bool, needs_horizontal_mirror: bool) -> int: if needs_vertical_mirror: square = chess.square_mirror(square) if needs_horizontal_mirror: square = square ^ 7 return square @staticmethod def get_color_and_piece_type(piece: str): is_white_piece = piece.isupper() piece_type = chess.PIECE_SYMBOLS.index(piece.lower()) - 1 return is_white_piece, piece_type @staticmethod def load_weights(file_name: str): weights = load_file(file_name) feature_weights = weights["feature_weights"] feature_bias = weights["feature_bias"] output_weights = weights["output_weights"] output_bias = weights["output_bias"] return (feature_weights, feature_bias, output_weights, output_bias) def evaluate_position(self, fen: str) -> int: # Create unprocessed board chess_board = chess.Board(fen) piece_square_pairs = [(piece.symbol(), square) for square, piece in chess_board.piece_map().items()] # Get active king buckets white_king_square = chess_board.king(chess.WHITE) white_perspective_needs_horizontal_mirror = chess.square_file(white_king_square) >= 4 white_king_square_for_inference = self.mirror_square_as_needed(white_king_square, False, white_perspective_needs_horizontal_mirror) white_active_king_bucket_index = chess.square_rank(white_king_square_for_inference) * 4 + chess.square_file(white_king_square_for_inference) white_active_king_bucket = self.INPUT_BUCKET_LAYOUT[white_active_king_bucket_index] black_king_square = chess_board.king(chess.BLACK) black_perspective_needs_horizontal_mirror = chess.square_file(black_king_square) >= 4 black_king_square_for_inference = self.mirror_square_as_needed(black_king_square, True, black_perspective_needs_horizontal_mirror) black_active_king_bucket_index = chess.square_rank(black_king_square_for_inference) * 4 + chess.square_file(black_king_square_for_inference) black_active_king_bucket = self.INPUT_BUCKET_LAYOUT[black_active_king_bucket_index] # Get active features white_input_indexes = [] black_input_indexes = [] # White perspective features for piece, square in piece_square_pairs: is_white_piece, piece_type = self.get_color_and_piece_type(piece) color_offset = 0 if is_white_piece else 6 * 64 transformed_sq = self.mirror_square_as_needed(square, False, white_perspective_needs_horizontal_mirror) white_input_indexes.append(color_offset + piece_type * 64 + transformed_sq) # Black perspective features for piece, square in piece_square_pairs: is_white_piece, piece_type = self.get_color_and_piece_type(piece) color_offset = 0 if not is_white_piece else 6 * 64 transformed_sq = self.mirror_square_as_needed(square, True, black_perspective_needs_horizontal_mirror) black_input_indexes.append(color_offset + piece_type * 64 + transformed_sq) # Get active output bucket piece_count = len(piece_square_pairs) piece_count_divisor = math.ceil(32 / self.OUTPUT_BUCKETS) active_output_bucket = (piece_count - 2) // piece_count_divisor # Inference: create accumulator white_accumulator = self.feature_bias.copy().astype(np.int32) black_accumulator = self.feature_bias.copy().astype(np.int32) for input_index in white_input_indexes: white_accumulator += self.feature_weights[white_active_king_bucket, input_index] for input_index in black_input_indexes: black_accumulator += self.feature_weights[black_active_king_bucket, input_index] # Inference: propagate output = 0 our_accumulator, opp_accumulator = (white_accumulator, black_accumulator) if chess_board.turn == chess.WHITE else (black_accumulator, white_accumulator) screlu_activation = lambda x: np.clip(x, 0, self.QA).astype(np.int32) ** 2 output += np.dot(screlu_activation(our_accumulator), self.output_weights[active_output_bucket, :self.HIDDEN_SIZE]) output += np.dot(screlu_activation(opp_accumulator), self.output_weights[active_output_bucket, self.HIDDEN_SIZE:2*self.HIDDEN_SIZE]) output = (output / self.QA + self.output_bias[active_output_bucket]) * self.SCALE / (self.QA * self.QB) return output if __name__ == "__main__": network = ChessNeuralNet("model.safetensors") fens = [ "rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR w KQkq - 0 1", "2r3k1/2P2pp1/3Np2p/8/7P/5qP1/5P1K/2Q5 b - - 2 42", "rnbqkb1r/pppppppp/8/3nP3/2P5/8/PP1P1PPP/RNBQKBNR b KQkq - 0 3", "2rk1b1r/1Qp1p2p/p1P5/5p2/2PP2B1/4B3/5P1P/1R2K3 w - - 0 35", "rnbqkb1r/pp2pppp/2p2n2/3p4/2PP4/5N2/PP2PPPP/RNBQKB1R w KQkq - 0 4", "r3r1k1/ppp2ppp/3qbb2/2p1N2Q/5P2/1PPB4/P1P3PP/1R3RK1 b - - 2 15", "7k/pp4rp/3p1Q2/1P1P4/2Pp4/3Pb2P/2q3P1/5R1K w - - 7 37", "8/8/4p3/4P1p1/Pk6/2p5/1p3K2/1r6 b - - 1 52", "r2qr1k1/pp2bppp/2n2n2/5b2/2P1N3/3P1N2/PP1BQPPP/R3KB1R w KQ - 2 11", "r6k/p1nP3p/n1QP1qp1/1B6/1b6/7P/2P2PP1/1R1R2K1 b - - 0 32", "1n3rk1/7p/1q4p1/p1p2pR1/2B2P2/4P3/P1QP1P2/4K3 b - - 1 24", "2r4k/5Qp1/p6p/3Np1b1/4P3/P7/1P2n1PP/2r2R1K w - - 1 33", "3r4/5p1k/5bpP/8/4K1P1/2p1PN2/8/7R b - - 9 105", "r1bqk2r/p1p1bppp/1p2pB2/8/3P4/3B1N2/PPP2PPP/R2QK2R b KQkq - 0 9", "1rbqk2r/4ppb1/2np2p1/p1p4p/N1P1P3/4BP1P/PP1Q2P1/2KR1B1R w k - 1 15", "5k2/8/4p2p/4P1p1/3P2P1/3b3K/3B2P1/8 w - - 92 121", "rnbqk2r/ppppppbp/5np1/8/2PP4/2N2N2/PP2PPPP/R1BQKB1R b KQkq - 3 4", "2r1r1k1/p3B2p/3p1P2/2pP3q/P1P3n1/1R3p1P/2Q3P1/1R5K b - - 2 47", "8/8/6R1/5K1p/8/5k2/6p1/8 w - - 0 79", "3k4/8/1Q6/1b6/8/1P1p1P2/3K4/7q b - - 7 64", "6Q1/8/8/7k/8/8/3p1pp1/3Kbrrb w - - 0 1", "1rqbkrbn/1ppppp1p/1n6/p1N3p1/8/2P4P/PP1PPPP1/1RQBKRBN w FBfb - 0 9", "rbbqn1kr/pp2p1pp/6n1/2pp1p2/2P4P/P7/BP1PPPP1/R1BQNNKR w HAha - 0 9", "rqbbk1r1/1ppp2pp/p3n1n1/4pp2/P7/1PP1N3/1Q1PPPPP/R1BB1RKN b ga - 3 10", "brnqnbkr/pppppppp/8/8/8/8/PPPPPPPP/BQNRNKRB w GDhb - 0 1" ] for fen in fens: score = network.evaluate_position(fen) print(f"Position: {fen}") print(f"Score: {round(score)}\n")