ParallaxOpen commited on
Commit
83c4df4
·
verified ·
1 Parent(s): 1829c5b

Upload vela_chess_engine.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. vela_chess_engine.py +143 -0
vela_chess_engine.py ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Vela-Chess Engine v2: Neural chess with custom board encoding.
3
+ Pure neural network evaluation — no Stockfish cheating.
4
+ """
5
+ import sys, torch, chess
6
+ from pathlib import Path
7
+
8
+ ROOT = Path(__file__).parent.parent
9
+ sys.path.insert(0, str(ROOT / "src"))
10
+ sys.path.insert(0, str(ROOT / "scripts"))
11
+ sys.path.insert(0, str(ROOT.parent))
12
+ from centauri.models.base import SmallLMConfig
13
+ from centauri.models.torch.small_lm import SmallLM
14
+ from chess_encoding import (
15
+ encode_board, encode_move, decode_move,
16
+ VOCAB_SIZE, BOS_TOKEN, EOS_TOKEN, SEP_TOKEN,
17
+ MOVE_FROM_OFFSET, MOVE_TO_OFFSET, PROMO_OFFSET
18
+ )
19
+
20
+
21
+ class VelaChessV2:
22
+ """Neural chess engine using board-to-move prediction."""
23
+
24
+ def __init__(self, checkpoint_path, device=None):
25
+ if device is None:
26
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
27
+ self.device = device
28
+
29
+ ckpt = torch.load(str(checkpoint_path), map_location="cpu")
30
+ cfg = SmallLMConfig(**ckpt["config"])
31
+ self.model = SmallLM(cfg)
32
+ self.model.load_state_dict(ckpt["model"])
33
+ self.model = self.model.to(self.device).eval()
34
+ self.n_params = sum(p.numel() for p in self.model.parameters())
35
+
36
+ def _get_move_logits(self, board):
37
+ """Get logits for next move token given board position."""
38
+ board_tokens = encode_board(board)
39
+ inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device)
40
+
41
+ with torch.no_grad():
42
+ out = self.model(idx=inp, targets=None)
43
+ logits = out["logits"]
44
+ if logits.dim() == 3:
45
+ logits = logits[:, -1, :]
46
+ return logits.squeeze()
47
+
48
+ def _get_second_token_logits(self, board, first_token):
49
+ """Get logits for second move token given board + first token."""
50
+ board_tokens = encode_board(board)
51
+ move_tokens = [first_token]
52
+ full = board_tokens + move_tokens
53
+ inp = torch.tensor([full], dtype=torch.long).to(self.device)
54
+
55
+ with torch.no_grad():
56
+ out = self.model(idx=inp, targets=None)
57
+ logits = out["logits"]
58
+ if logits.dim() == 3:
59
+ logits = logits[:, -1, :]
60
+ return logits.squeeze()
61
+
62
+ def _score_move(self, board, move, temperature=0.5):
63
+ """Score a legal move using the model (2-step prediction)."""
64
+ board_tokens = encode_board(board)
65
+ move_tokens = encode_move(move)
66
+ full = board_tokens + move_tokens
67
+ inp = torch.tensor([full[:-1]], dtype=torch.long).to(self.device)
68
+ target = torch.tensor([full[-1]], dtype=torch.long).to(self.device)
69
+
70
+ with torch.no_grad():
71
+ out = self.model(idx=inp, targets=None)
72
+ logits = out["logits"]
73
+ if logits.dim() == 3:
74
+ logits = logits[:, -1, :]
75
+ log_probs = torch.log_softmax(logits / temperature, dim=-1)
76
+ score = log_probs[0, target.item()].item()
77
+ return score
78
+
79
+ def choose_move(self, board, temperature=0.5):
80
+ """Choose the best legal move using model scoring."""
81
+ legal_moves = list(board.legal_moves)
82
+ if not legal_moves:
83
+ return None
84
+ if len(legal_moves) == 1:
85
+ return legal_moves[0]
86
+
87
+ scores = {}
88
+ for move in legal_moves:
89
+ scores[move] = self._score_move(board, move, temperature)
90
+
91
+ best_move = max(scores, key=scores.get)
92
+ return best_move
93
+
94
+ def analyze(self, board, top_n=5, temperature=0.5):
95
+ """Analyze position, return top N moves with scores."""
96
+ legal_moves = list(board.legal_moves)
97
+ if not legal_moves:
98
+ return []
99
+
100
+ scores = {}
101
+ for move in legal_moves:
102
+ scores[move] = self._score_move(board, move, temperature)
103
+
104
+ ranked = sorted(scores.items(), key=lambda x: -x[1])
105
+ return ranked[:top_n]
106
+
107
+ def choose_move_generative(self, board, temperature=0.5, num_attempts=10):
108
+ """Choose move by generating tokens autoregressively, then filtering valid moves."""
109
+ board_tokens = encode_board(board)
110
+ inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device)
111
+
112
+ candidates = []
113
+ for _ in range(num_attempts):
114
+ with torch.no_grad():
115
+ # Generate first token
116
+ out = self.model(idx=inp, targets=None)
117
+ logits = out["logits"]
118
+ if logits.dim() == 3:
119
+ logits = logits[:, -1, :]
120
+ probs = torch.softmax(logits / temperature, dim=-1)
121
+ first_token = torch.multinomial(probs, 1).item()
122
+
123
+ # Generate second token
124
+ second_input = torch.cat([inp, torch.tensor([[first_token]], device=self.device)], dim=1)
125
+ out2 = self.model(idx=second_input, targets=None)
126
+ logits2 = out2["logits"]
127
+ if logits2.dim() == 3:
128
+ logits2 = logits2[:, -1, :]
129
+ probs2 = torch.softmax(logits2 / temperature, dim=-1)
130
+ second_token = torch.multinomial(probs2, 1).item()
131
+
132
+ move = decode_move([first_token, second_token], board)
133
+ if move is not None:
134
+ candidates.append(move)
135
+
136
+ if not candidates:
137
+ # Fallback to scoring
138
+ return self.choose_move(board, temperature)
139
+
140
+ # Return most common candidate
141
+ from collections import Counter
142
+ counts = Counter(candidates)
143
+ return counts.most_common(1)[0][0]