pkrisz commited on
Commit
d5f8b01
·
verified ·
1 Parent(s): 510b9fc

Upload 3 files

Browse files
Files changed (3) hide show
  1. inference.py +149 -0
  2. model.bin +3 -0
  3. 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