File size: 9,696 Bytes
ba69de3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
"""

Minimal UCI engine wrapper for the trained ChessResNet model.



Can be used by chess GUIs (Arena, cutechess, En Croissant, etc.) or

test harnesses that speak the UCI protocol.



Usage:

    python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cuda

    python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cpu

"""
import argparse
import sys
from pathlib import Path

import chess
import torch

sys.path.insert(0, str(Path(__file__).resolve().parent))
from common_chess import encode_board, move_to_action_id, legal_action_ids
from model import ChessResNet, create_model_from_config

# ── Draw-aware move selection helpers ──────────────────────────────────────

PIECE_VALUES = {
    chess.PAWN: 100,
    chess.KNIGHT: 320,
    chess.BISHOP: 330,
    chess.ROOK: 500,
    chess.QUEEN: 900,
    chess.KING: 0,
}


def material_score_for_side(board: chess.Board, side: chess.Color) -> int:
    """Return *side*'s material advantage in centipawns (positive = side ahead)."""
    score = 0
    for piece_type in chess.PIECE_TYPES:
        value = PIECE_VALUES[piece_type]
        score += len(board.pieces(piece_type, side)) * value
        score -= len(board.pieces(piece_type, not side)) * value
    return score


def move_causes_drawish(board: chess.Board, move: chess.Move) -> bool:
    """Check whether *move* immediately leads to a drawish outcome."""
    b = board.copy(stack=True)
    b.push(move)
    if b.is_repetition(3):
        return True
    if b.can_claim_threefold_repetition():
        return True
    if b.is_fifty_moves():
        return True
    if b.can_claim_fifty_moves():
        return True
    if b.is_stalemate():
        return True
    if b.is_insufficient_material():
        return True
    return False


def choose_move_with_draw_awareness(

    board: chess.Board,

    legal_moves: list,

    legal_action_ids: list[int],

    policy_logits: torch.Tensor,

    value_pred,

    topk: int = 12,

) -> chess.Move:
    """Draw-aware top-k rerank of legal moves.



    Parameters

    ----------

    board : current python-chess Board (must retain move stack).

    legal_moves : list of chess.Move, parallel to *legal_action_ids*.

    legal_action_ids : list of int action IDs parallel to *legal_moves*.

    policy_logits : full [N_ACTIONS] torch.Tensor (on any device).

    value_pred : scalar value-head output (can be tensor or float).

    topk : number of top candidates to consider for reranking.



    Returns a legal ``chess.Move``.

    """
    side = board.turn
    root_value = float(value_pred.squeeze().item() if hasattr(value_pred, "item") else value_pred)
    material = material_score_for_side(board, side)

    # Material fallback: override value signal when material gap is large
    if material >= 500:
        root_value = max(root_value, 0.50)
    if material <= -500:
        root_value = min(root_value, -0.50)

    # Sort legal moves by policy logit descending
    scored = [
        (float(policy_logits[aid].item() if hasattr(policy_logits, "item") else policy_logits[aid]), move)
        for aid, move in zip(legal_action_ids, legal_moves)
    ]
    scored.sort(key=lambda x: x[0], reverse=True)
    sorted_moves = [m for _, m in scored]
    original_top1 = sorted_moves[0]

    # ── Advantage: avoid draws ────────────────────────────────────────
    if root_value > 0.35:
        for move in sorted_moves[:topk]:
            if not move_causes_drawish(board, move):
                return move
        return original_top1  # fallback – every top-k move is drawish

    # ── Disadvantage: prefer draws ────────────────────────────────────
    if root_value < -0.35:
        for move in sorted_moves[:topk]:
            if move_causes_drawish(board, move):
                return move
        return original_top1  # fallback – no drawish move in top-k

    # ── Near-equality: stick with policy top1 ─────────────────────────
    return original_top1


class UCIEngine:
    """Minimal UCI chess engine using a trained ChessResNet model."""

    def __init__(self, ckpt_path: str, device: str = "cuda"):
        self.device = device if torch.cuda.is_available() and device == "cuda" else "cpu"
        self.board = chess.Board()
        self.model = self._load_model(ckpt_path)
        self.model.eval()
        self._stop_requested = False

    def _load_model(self, ckpt_path: str) -> ChessResNet:
        ckpt = torch.load(ckpt_path, map_location=self.device, weights_only=True)
        model_config = ckpt.get("model_config", {})
        if not model_config:
            model_config = {
                "channels": ckpt.get("args", {}).get("channels", 256),
                "blocks": ckpt.get("args", {}).get("blocks", 20),
                "num_actions": 20480,
            }
        model = create_model_from_config(model_config)
        model.load_state_dict(ckpt["model"])
        model.to(self.device)
        return model

    def uci_new_game(self):
        """Reset board for a new game."""
        self.board.reset()
        self._stop_requested = False

    def set_position(self, fen: str | None = None, moves: list[str] | None = None):
        """Set up position from FEN and optional move list."""
        if fen:
            self.board.set_fen(fen)
        else:
            self.board.reset()
        if moves:
            for m in moves:
                self.board.push(chess.Move.from_uci(m))

    def get_best_move(self, movetime_ms: int = 1000) -> tuple[str, float]:
        """

        Return (bestmove_uci, top_logit) by evaluating the current board.



        Only considers legal moves.  This is a single-forward-pass evaluator;

        it does NOT do MCTS or search.

        """
        # Encode board
        planes = encode_board(self.board)
        inp = torch.from_numpy(planes).unsqueeze(0).float().to(self.device)  # [1,18,8,8]

        with torch.no_grad():
            with torch.autocast(device_type=self.device, enabled=(self.device == "cuda")):
                policy_logits, value = self.model(inp)

        policy_logits = policy_logits.squeeze(0)  # [20480]

        # Get legal moves and their action IDs
        action_ids, moves = legal_action_ids(self.board)

        if not moves:
            return "0000", float("-inf")

        # Draw-aware top-k rerank (avoids threefold-repetition, 50-move, etc.)
        best_move_obj = choose_move_with_draw_awareness(
            self.board, moves, action_ids, policy_logits, value, topk=12
        )
        best_move = best_move_obj.uci()
        best_logit = float(policy_logits[move_to_action_id(best_move_obj, self.board.turn)])

        return best_move, best_logit

    def handle_go(self, tokens: list[str]):
        """Process 'go' command and output bestmove."""
        movetime_ms = 1000
        if "movetime" in tokens:
            idx = tokens.index("movetime") + 1
            if idx < len(tokens):
                movetime_ms = int(tokens[idx])

        best_move, _ = self.get_best_move(movetime_ms)
        print(f"bestmove {best_move}", flush=True)

    def handle_position(self, tokens: list[str]):
        """Process 'position' command."""
        fen = None
        moves = []

        if "startpos" in tokens:
            pass  # use starting position (board.reset() already done or standard)
        elif "fen" in tokens:
            # Collect FEN string up to "moves" keyword
            idx = tokens.index("fen") + 1
            fen_parts = []
            while idx < len(tokens) and tokens[idx] != "moves":
                fen_parts.append(tokens[idx])
                idx += 1
            fen = " ".join(fen_parts)

        if "moves" in tokens:
            idx = tokens.index("moves") + 1
            moves = tokens[idx:]

        self.set_position(fen, moves)

    def run(self):
        """Main UCI loop: read commands from stdin, respond to stdout."""
        while True:
            line = sys.stdin.readline()
            if not line:
                break
            line = line.strip()
            if not line:
                continue

            parts = line.split()
            cmd = parts[0]

            if cmd == "uci":
                print("id name JoeyStage1StockfishDistill", flush=True)
                print("id author Joey", flush=True)
                print("uciok", flush=True)
            elif cmd == "isready":
                print("readyok", flush=True)
            elif cmd == "ucinewgame":
                self.uci_new_game()
            elif cmd == "position":
                self.handle_position(parts[1:])
            elif cmd == "go":
                self.handle_go(parts[1:])
            elif cmd == "stop":
                self._stop_requested = True
            elif cmd == "quit":
                break


def main():
    parser = argparse.ArgumentParser(description="UCI chess engine")
    parser.add_argument("--ckpt", default='', help="Path to checkpoint .pt file")
    parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
    args = parser.parse_args()

    engine = UCIEngine(args.ckpt, args.device)
    engine.run()


if __name__ == "__main__":
    main()