ParallaxOpen commited on
Commit
fe8739d
·
verified ·
1 Parent(s): 25e96d1

Upload play_gui.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. play_gui.py +250 -0
play_gui.py ADDED
@@ -0,0 +1,250 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Parallax-Chess V3: Tkinter GUI."""
3
+ import chess
4
+ import tkinter as tk
5
+ from tkinter import font as tkfont
6
+ import sys, time
7
+ from pathlib import Path
8
+
9
+ sys.path.insert(0, str(Path(__file__).parent / "scripts"))
10
+ sys.path.insert(0, str(Path(__file__).parent / "src"))
11
+ sys.path.insert(0, str(Path(__file__).parent.parent))
12
+
13
+ import torch
14
+ from train_v3 import ParallaxChessV3
15
+ from centauri.models.base import SmallLMConfig
16
+ from chess_encoding import VOCAB_SIZE
17
+
18
+ PIECE_UNICODE = {
19
+ "K": "\u2654", "Q": "\u2655", "R": "\u2656", "B": "\u2657", "N": "\u2658", "P": "\u2659",
20
+ "k": "\u265A", "q": "\u265B", "r": "\u265C", "b": "\u265D", "n": "\u265E", "p": "\u265F",
21
+ }
22
+
23
+ LIGHT = "#F0D9B5"
24
+ DARK = "#B58863"
25
+ LAST_MOVE_COLOR = "#CDD26A"
26
+ SELECT_COLOR = "#F6F669"
27
+ MOVE_DOT_COLOR = "#7B9E57"
28
+
29
+
30
+ class ChessGUI:
31
+ def __init__(self):
32
+ self.root = tk.Tk()
33
+ self.root.title("Parallax-Chess V3")
34
+ self.root.resizable(False, False)
35
+
36
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
37
+ cfg = SmallLMConfig(
38
+ vocab_size=VOCAB_SIZE, d_model=512, n_heads=8, n_kv_heads=4,
39
+ n_layers=8, intermediate_size=2048, max_seq_len=128,
40
+ norm_type="rms", rope_type="neox", n_experts=0,
41
+ n_loops=1, loop_mode="per_layer",
42
+ )
43
+ self.model = ParallaxChessV3(cfg).to(self.device)
44
+ ckpt_path = Path(__file__).parent / "checkpoints" / "parallax_chess_v3"
45
+ candidates = sorted(ckpt_path.glob("step_*.pt"), reverse=True)
46
+ ckpt = torch.load(str(candidates[0]), map_location=self.device)
47
+ self.model.load_state_dict(ckpt["model"])
48
+ self.model.eval()
49
+ print("Loaded %s" % candidates[0].name)
50
+
51
+ self.board = chess.Board()
52
+ self.SQ = 72
53
+ self.selected = None
54
+ self.legal_targets = []
55
+ self.last_move = None
56
+ self.player_color = chess.WHITE
57
+ self.flipped = False
58
+
59
+ self.piece_font = tkfont.Font(family="Segoe UI Symbol", size=36)
60
+ self.coord_font = tkfont.Font(family="Consolas", size=10, weight="bold")
61
+ self.info_font = tkfont.Font(family="Consolas", size=11)
62
+
63
+ top = tk.Frame(self.root)
64
+ top.pack(padx=10, pady=5, fill=tk.X)
65
+ self.status_var = tk.StringVar(value="Your turn (White)")
66
+ tk.Label(top, textvariable=self.status_var, font=self.info_font, width=35).pack(side=tk.LEFT)
67
+ self.eval_var = tk.StringVar(value="Eval: 0.0")
68
+ tk.Label(top, textvariable=self.eval_var, font=self.info_font, width=15).pack(side=tk.LEFT)
69
+
70
+ canvas_size = self.SQ * 8 + 30
71
+ self.canvas = tk.Canvas(self.root, width=canvas_size, height=canvas_size, bg="#312E2B")
72
+ self.canvas.pack(padx=10, pady=5)
73
+ self.canvas.bind("<Button-1>", self.on_click)
74
+
75
+ bottom = tk.Frame(self.root)
76
+ bottom.pack(padx=10, pady=5)
77
+ tk.Button(bottom, text="New Game", font=self.info_font, command=self.new_game).pack(side=tk.LEFT, padx=5)
78
+ tk.Button(bottom, text="Flip Board", font=self.info_font, command=self.flip_board).pack(side=tk.LEFT, padx=5)
79
+
80
+ self.draw_board()
81
+
82
+ def sq_to_rc(self, sq):
83
+ """Convert chess square to display (row, col) based on flip."""
84
+ file = chess.square_file(sq)
85
+ rank = chess.square_rank(sq)
86
+ if self.flipped:
87
+ row = rank
88
+ col = 7 - file
89
+ else:
90
+ row = 7 - rank
91
+ col = file
92
+ return row, col
93
+
94
+ def rc_to_sq(self, row, col):
95
+ """Convert display (row, col) to chess square."""
96
+ if self.flipped:
97
+ file = 7 - col
98
+ rank = row
99
+ else:
100
+ file = col
101
+ rank = 7 - row
102
+ return chess.square(file, rank)
103
+
104
+ def new_game(self):
105
+ self.board = chess.Board()
106
+ self.selected = None
107
+ self.legal_targets = []
108
+ self.last_move = None
109
+ self.status_var.set("Your turn (White)")
110
+ self.eval_var.set("Eval: 0.0")
111
+ self.draw_board()
112
+
113
+ def flip_board(self):
114
+ self.flipped = not self.flipped
115
+ self.draw_board()
116
+
117
+ def draw_board(self):
118
+ self.canvas.delete("all")
119
+ offset = 15
120
+
121
+ for display_row in range(8):
122
+ for display_col in range(8):
123
+ x1 = offset + display_col * self.SQ
124
+ y1 = offset + display_row * self.SQ
125
+ x2 = x1 + self.SQ
126
+ y2 = y1 + self.SQ
127
+
128
+ sq = self.rc_to_sq(display_row, display_col)
129
+
130
+ is_dark = (display_row + display_col) % 2 == 1
131
+ color = DARK if is_dark else LIGHT
132
+
133
+ if self.last_move and sq in (self.last_move.from_square, self.last_move.to_square):
134
+ color = LAST_MOVE_COLOR
135
+ if self.selected == sq:
136
+ color = SELECT_COLOR
137
+
138
+ self.canvas.create_rectangle(x1, y1, x2, y2, fill=color, outline="")
139
+
140
+ piece = self.board.piece_at(sq)
141
+ if piece:
142
+ symbol = PIECE_UNICODE[piece.symbol()]
143
+ fill = "#FFFFFF" if piece.color == chess.WHITE else "#1a1a1a"
144
+ self.canvas.create_text(
145
+ x1 + self.SQ // 2, y1 + self.SQ // 2,
146
+ text=symbol, font=self.piece_font, fill=fill,
147
+ )
148
+
149
+ if sq in self.legal_targets:
150
+ if self.board.piece_at(sq):
151
+ self.canvas.create_oval(
152
+ x1 + 4, y1 + 4, x2 - 4, y2 - 4,
153
+ outline=MOVE_DOT_COLOR, width=3, fill="",
154
+ )
155
+ else:
156
+ r_dot = 8
157
+ cx, cy = x1 + self.SQ // 2, y1 + self.SQ // 2
158
+ self.canvas.create_oval(
159
+ cx - r_dot, cy - r_dot, cx + r_dot, cy + r_dot,
160
+ fill=MOVE_DOT_COLOR, outline="",
161
+ )
162
+
163
+ # File labels (a-h) along bottom
164
+ files = "abcdefgh" if not self.flipped else "hgfedcba"
165
+ for i in range(8):
166
+ x = offset + i * self.SQ + self.SQ // 2
167
+ self.canvas.create_text(x, offset + 8 * self.SQ + 12, text=files[i],
168
+ font=self.coord_font, fill="#AAAAAA")
169
+
170
+ # Rank labels (1-8) along left
171
+ ranks = "87654321" if not self.flipped else "12345678"
172
+ for i in range(8):
173
+ y = offset + i * self.SQ + self.SQ // 2
174
+ self.canvas.create_text(offset - 8, y, text=ranks[i],
175
+ font=self.coord_font, fill="#AAAAAA")
176
+
177
+ def on_click(self, event):
178
+ if self.board.is_game_over():
179
+ self.status_var.set("Game over: %s" % self.board.result())
180
+ return
181
+ if self.board.turn != self.player_color:
182
+ return
183
+
184
+ offset = 15
185
+ display_col = (event.x - offset) // self.SQ
186
+ display_row = (event.y - offset) // self.SQ
187
+ if display_col < 0 or display_col > 7 or display_row < 0 or display_row > 7:
188
+ return
189
+
190
+ sq = self.rc_to_sq(display_row, display_col)
191
+
192
+ if self.selected is not None and sq in self.legal_targets:
193
+ move = chess.Move(self.selected, sq)
194
+ piece = self.board.piece_at(self.selected)
195
+ if piece and piece.piece_type == chess.PAWN:
196
+ if (piece.color == chess.WHITE and chess.square_rank(sq) == 7) or \
197
+ (piece.color == chess.BLACK and chess.square_rank(sq) == 0):
198
+ move = chess.Move(self.selected, sq, promotion=chess.QUEEN)
199
+
200
+ self.board.push(move)
201
+ self.last_move = move
202
+ self.selected = None
203
+ self.legal_targets = []
204
+ self.draw_board()
205
+ self.root.after(50, self.ai_move)
206
+ return
207
+
208
+ piece = self.board.piece_at(sq)
209
+ if piece and piece.color == self.board.turn:
210
+ self.selected = sq
211
+ self.legal_targets = [m.to_square for m in self.board.legal_moves if m.from_square == sq]
212
+ else:
213
+ self.selected = None
214
+ self.legal_targets = []
215
+ self.draw_board()
216
+
217
+ def ai_move(self):
218
+ if self.board.is_game_over():
219
+ self.status_var.set("Game over: %s" % self.board.result())
220
+ return
221
+
222
+ self.status_var.set("V3 thinking...")
223
+ self.root.update()
224
+
225
+ t0 = time.time()
226
+ move = self.model.predict_move(self.board, self.device)
227
+ elapsed = time.time() - t0
228
+
229
+ if move is None:
230
+ self.status_var.set("V3 resigns!")
231
+ return
232
+
233
+ self.board.push(move)
234
+ self.last_move = move
235
+ self.status_var.set("V3 played %s (%.1fs)" % (move.uci(), elapsed))
236
+
237
+ eval_cp = self.model.evaluate(self.board, self.device)
238
+ self.eval_var.set("Eval: %.0f cp" % eval_cp)
239
+
240
+ self.draw_board()
241
+
242
+ if self.board.is_game_over():
243
+ self.status_var.set("Game over: %s" % self.board.result())
244
+
245
+ def run(self):
246
+ self.root.mainloop()
247
+
248
+
249
+ if __name__ == "__main__":
250
+ ChessGUI().run()