Upload 3 files
Browse files- inference.py +149 -0
- model.bin +3 -0
- model.safetensors +3 -0
inference.py
ADDED
|
@@ -0,0 +1,149 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import math
|
| 2 |
+
from safetensors.numpy import load_file
|
| 3 |
+
import numpy as np
|
| 4 |
+
import chess
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class ChessNeuralNet:
|
| 8 |
+
def __init__(self, name: str):
|
| 9 |
+
self.FEATURE_SIZE = 768
|
| 10 |
+
self.HIDDEN_SIZE = 1600
|
| 11 |
+
self.INPUT_BUCKETS = 14
|
| 12 |
+
self.OUTPUT_BUCKETS = 8
|
| 13 |
+
self.INPUT_BUCKET_LAYOUT = [
|
| 14 |
+
0, 1, 2, 3,
|
| 15 |
+
4, 5, 6, 7,
|
| 16 |
+
8, 8, 9, 9,
|
| 17 |
+
10, 10, 11, 11,
|
| 18 |
+
10, 10, 11, 11,
|
| 19 |
+
12, 12, 13, 13,
|
| 20 |
+
12, 12, 13, 13,
|
| 21 |
+
12, 12, 13, 13,
|
| 22 |
+
]
|
| 23 |
+
self.QA = 255
|
| 24 |
+
self.QB = 64
|
| 25 |
+
self.SCALE = 400
|
| 26 |
+
|
| 27 |
+
self.feature_weights, self.feature_bias, self.output_weights, self.output_bias = self.load_weights(name)
|
| 28 |
+
|
| 29 |
+
@staticmethod
|
| 30 |
+
def mirror_square_as_needed(square: int, needs_vertical_mirror: bool, needs_horizontal_mirror: bool) -> int:
|
| 31 |
+
if needs_vertical_mirror:
|
| 32 |
+
square = chess.square_mirror(square)
|
| 33 |
+
if needs_horizontal_mirror:
|
| 34 |
+
square = square ^ 7
|
| 35 |
+
return square
|
| 36 |
+
|
| 37 |
+
@staticmethod
|
| 38 |
+
def get_color_and_piece_type(piece: str):
|
| 39 |
+
is_white_piece = piece.isupper()
|
| 40 |
+
piece_type = chess.PIECE_SYMBOLS.index(piece.lower()) - 1
|
| 41 |
+
return is_white_piece, piece_type
|
| 42 |
+
|
| 43 |
+
@staticmethod
|
| 44 |
+
def load_weights(file_name: str):
|
| 45 |
+
weights = load_file(file_name)
|
| 46 |
+
feature_weights = weights["feature_weights"]
|
| 47 |
+
feature_bias = weights["feature_bias"]
|
| 48 |
+
output_weights = weights["output_weights"]
|
| 49 |
+
output_bias = weights["output_bias"]
|
| 50 |
+
return (feature_weights, feature_bias, output_weights, output_bias)
|
| 51 |
+
|
| 52 |
+
def evaluate_position(self, fen: str) -> int:
|
| 53 |
+
# Create unprocessed board
|
| 54 |
+
chess_board = chess.Board(fen)
|
| 55 |
+
piece_square_pairs = [(piece.symbol(), square) for square, piece in chess_board.piece_map().items()]
|
| 56 |
+
|
| 57 |
+
# Get active king buckets
|
| 58 |
+
white_king_square = chess_board.king(chess.WHITE)
|
| 59 |
+
white_perspective_needs_horizontal_mirror = chess.square_file(white_king_square) >= 4
|
| 60 |
+
white_king_square_for_inference = self.mirror_square_as_needed(white_king_square, False, white_perspective_needs_horizontal_mirror)
|
| 61 |
+
white_active_king_bucket_index = chess.square_rank(white_king_square_for_inference) * 4 + chess.square_file(white_king_square_for_inference)
|
| 62 |
+
white_active_king_bucket = self.INPUT_BUCKET_LAYOUT[white_active_king_bucket_index]
|
| 63 |
+
|
| 64 |
+
black_king_square = chess_board.king(chess.BLACK)
|
| 65 |
+
black_perspective_needs_horizontal_mirror = chess.square_file(black_king_square) >= 4
|
| 66 |
+
black_king_square_for_inference = self.mirror_square_as_needed(black_king_square, True, black_perspective_needs_horizontal_mirror)
|
| 67 |
+
black_active_king_bucket_index = chess.square_rank(black_king_square_for_inference) * 4 + chess.square_file(black_king_square_for_inference)
|
| 68 |
+
black_active_king_bucket = self.INPUT_BUCKET_LAYOUT[black_active_king_bucket_index]
|
| 69 |
+
|
| 70 |
+
# Get active features
|
| 71 |
+
white_input_indexes = []
|
| 72 |
+
black_input_indexes = []
|
| 73 |
+
|
| 74 |
+
# White perspective features
|
| 75 |
+
for piece, square in piece_square_pairs:
|
| 76 |
+
is_white_piece, piece_type = self.get_color_and_piece_type(piece)
|
| 77 |
+
color_offset = 0 if is_white_piece else 6 * 64
|
| 78 |
+
transformed_sq = self.mirror_square_as_needed(square, False, white_perspective_needs_horizontal_mirror)
|
| 79 |
+
white_input_indexes.append(color_offset + piece_type * 64 + transformed_sq)
|
| 80 |
+
|
| 81 |
+
# Black perspective features
|
| 82 |
+
for piece, square in piece_square_pairs:
|
| 83 |
+
is_white_piece, piece_type = self.get_color_and_piece_type(piece)
|
| 84 |
+
color_offset = 0 if not is_white_piece else 6 * 64
|
| 85 |
+
transformed_sq = self.mirror_square_as_needed(square, True, black_perspective_needs_horizontal_mirror)
|
| 86 |
+
black_input_indexes.append(color_offset + piece_type * 64 + transformed_sq)
|
| 87 |
+
|
| 88 |
+
# Get active output bucket
|
| 89 |
+
piece_count = len(piece_square_pairs)
|
| 90 |
+
piece_count_divisor = math.ceil(32 / self.OUTPUT_BUCKETS)
|
| 91 |
+
active_output_bucket = (piece_count - 2) // piece_count_divisor
|
| 92 |
+
|
| 93 |
+
# Inference: create accumulator
|
| 94 |
+
white_accumulator = self.feature_bias.copy().astype(np.int32)
|
| 95 |
+
black_accumulator = self.feature_bias.copy().astype(np.int32)
|
| 96 |
+
|
| 97 |
+
for input_index in white_input_indexes:
|
| 98 |
+
white_accumulator += self.feature_weights[white_active_king_bucket, input_index]
|
| 99 |
+
for input_index in black_input_indexes:
|
| 100 |
+
black_accumulator += self.feature_weights[black_active_king_bucket, input_index]
|
| 101 |
+
|
| 102 |
+
# Inference: propagate
|
| 103 |
+
output = 0
|
| 104 |
+
our_accumulator, opp_accumulator = (white_accumulator, black_accumulator) if chess_board.turn == chess.WHITE else (black_accumulator, white_accumulator)
|
| 105 |
+
|
| 106 |
+
screlu_activation = lambda x: np.clip(x, 0, self.QA).astype(np.int32) ** 2
|
| 107 |
+
output += np.dot(screlu_activation(our_accumulator), self.output_weights[active_output_bucket, :self.HIDDEN_SIZE])
|
| 108 |
+
output += np.dot(screlu_activation(opp_accumulator), self.output_weights[active_output_bucket, self.HIDDEN_SIZE:2*self.HIDDEN_SIZE])
|
| 109 |
+
output = (output / self.QA + self.output_bias[active_output_bucket]) * self.SCALE / (self.QA * self.QB)
|
| 110 |
+
|
| 111 |
+
return output
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
if __name__ == "__main__":
|
| 115 |
+
|
| 116 |
+
network = ChessNeuralNet("model.safetensors")
|
| 117 |
+
|
| 118 |
+
fens = [
|
| 119 |
+
"rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR w KQkq - 0 1",
|
| 120 |
+
"2r3k1/2P2pp1/3Np2p/8/7P/5qP1/5P1K/2Q5 b - - 2 42",
|
| 121 |
+
"rnbqkb1r/pppppppp/8/3nP3/2P5/8/PP1P1PPP/RNBQKBNR b KQkq - 0 3",
|
| 122 |
+
"2rk1b1r/1Qp1p2p/p1P5/5p2/2PP2B1/4B3/5P1P/1R2K3 w - - 0 35",
|
| 123 |
+
"rnbqkb1r/pp2pppp/2p2n2/3p4/2PP4/5N2/PP2PPPP/RNBQKB1R w KQkq - 0 4",
|
| 124 |
+
"r3r1k1/ppp2ppp/3qbb2/2p1N2Q/5P2/1PPB4/P1P3PP/1R3RK1 b - - 2 15",
|
| 125 |
+
"7k/pp4rp/3p1Q2/1P1P4/2Pp4/3Pb2P/2q3P1/5R1K w - - 7 37",
|
| 126 |
+
"8/8/4p3/4P1p1/Pk6/2p5/1p3K2/1r6 b - - 1 52",
|
| 127 |
+
"r2qr1k1/pp2bppp/2n2n2/5b2/2P1N3/3P1N2/PP1BQPPP/R3KB1R w KQ - 2 11",
|
| 128 |
+
"r6k/p1nP3p/n1QP1qp1/1B6/1b6/7P/2P2PP1/1R1R2K1 b - - 0 32",
|
| 129 |
+
"1n3rk1/7p/1q4p1/p1p2pR1/2B2P2/4P3/P1QP1P2/4K3 b - - 1 24",
|
| 130 |
+
"2r4k/5Qp1/p6p/3Np1b1/4P3/P7/1P2n1PP/2r2R1K w - - 1 33",
|
| 131 |
+
"3r4/5p1k/5bpP/8/4K1P1/2p1PN2/8/7R b - - 9 105",
|
| 132 |
+
"r1bqk2r/p1p1bppp/1p2pB2/8/3P4/3B1N2/PPP2PPP/R2QK2R b KQkq - 0 9",
|
| 133 |
+
"1rbqk2r/4ppb1/2np2p1/p1p4p/N1P1P3/4BP1P/PP1Q2P1/2KR1B1R w k - 1 15",
|
| 134 |
+
"5k2/8/4p2p/4P1p1/3P2P1/3b3K/3B2P1/8 w - - 92 121",
|
| 135 |
+
"rnbqk2r/ppppppbp/5np1/8/2PP4/2N2N2/PP2PPPP/R1BQKB1R b KQkq - 3 4",
|
| 136 |
+
"2r1r1k1/p3B2p/3p1P2/2pP3q/P1P3n1/1R3p1P/2Q3P1/1R5K b - - 2 47",
|
| 137 |
+
"8/8/6R1/5K1p/8/5k2/6p1/8 w - - 0 79",
|
| 138 |
+
"3k4/8/1Q6/1b6/8/1P1p1P2/3K4/7q b - - 7 64",
|
| 139 |
+
"6Q1/8/8/7k/8/8/3p1pp1/3Kbrrb w - - 0 1",
|
| 140 |
+
"1rqbkrbn/1ppppp1p/1n6/p1N3p1/8/2P4P/PP1PPPP1/1RQBKRBN w FBfb - 0 9",
|
| 141 |
+
"rbbqn1kr/pp2p1pp/6n1/2pp1p2/2P4P/P7/BP1PPPP1/R1BQNNKR w HAha - 0 9",
|
| 142 |
+
"rqbbk1r1/1ppp2pp/p3n1n1/4pp2/P7/1PP1N3/1Q1PPPPP/R1BB1RKN b ga - 3 10",
|
| 143 |
+
"brnqnbkr/pppppppp/8/8/8/8/PPPPPPPP/BQNRNKRB w GDhb - 0 1"
|
| 144 |
+
]
|
| 145 |
+
|
| 146 |
+
for fen in fens:
|
| 147 |
+
score = network.evaluate_position(fen)
|
| 148 |
+
print(f"Position: {fen}")
|
| 149 |
+
print(f"Score: {round(score)}\n")
|
model.bin
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:87cefa667e1cc99f181a8f7987e97de48c697fb51d13a9b80883275c38ec0a5c
|
| 3 |
+
size 34460864
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d440626f447199d6b9721dd24c7b51c3dbb41127234d5286bd7832ecef42e7a6
|
| 3 |
+
size 34461144
|