File size: 5,528 Bytes
83c4df4 | 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 | #!/usr/bin/env python3
"""Vela-Chess Engine v2: Neural chess with custom board encoding.
Pure neural network evaluation — no Stockfish cheating.
"""
import sys, torch, chess
from pathlib import Path
ROOT = Path(__file__).parent.parent
sys.path.insert(0, str(ROOT / "src"))
sys.path.insert(0, str(ROOT / "scripts"))
sys.path.insert(0, str(ROOT.parent))
from centauri.models.base import SmallLMConfig
from centauri.models.torch.small_lm import SmallLM
from chess_encoding import (
encode_board, encode_move, decode_move,
VOCAB_SIZE, BOS_TOKEN, EOS_TOKEN, SEP_TOKEN,
MOVE_FROM_OFFSET, MOVE_TO_OFFSET, PROMO_OFFSET
)
class VelaChessV2:
"""Neural chess engine using board-to-move prediction."""
def __init__(self, checkpoint_path, device=None):
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.device = device
ckpt = torch.load(str(checkpoint_path), map_location="cpu")
cfg = SmallLMConfig(**ckpt["config"])
self.model = SmallLM(cfg)
self.model.load_state_dict(ckpt["model"])
self.model = self.model.to(self.device).eval()
self.n_params = sum(p.numel() for p in self.model.parameters())
def _get_move_logits(self, board):
"""Get logits for next move token given board position."""
board_tokens = encode_board(board)
inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device)
with torch.no_grad():
out = self.model(idx=inp, targets=None)
logits = out["logits"]
if logits.dim() == 3:
logits = logits[:, -1, :]
return logits.squeeze()
def _get_second_token_logits(self, board, first_token):
"""Get logits for second move token given board + first token."""
board_tokens = encode_board(board)
move_tokens = [first_token]
full = board_tokens + move_tokens
inp = torch.tensor([full], dtype=torch.long).to(self.device)
with torch.no_grad():
out = self.model(idx=inp, targets=None)
logits = out["logits"]
if logits.dim() == 3:
logits = logits[:, -1, :]
return logits.squeeze()
def _score_move(self, board, move, temperature=0.5):
"""Score a legal move using the model (2-step prediction)."""
board_tokens = encode_board(board)
move_tokens = encode_move(move)
full = board_tokens + move_tokens
inp = torch.tensor([full[:-1]], dtype=torch.long).to(self.device)
target = torch.tensor([full[-1]], dtype=torch.long).to(self.device)
with torch.no_grad():
out = self.model(idx=inp, targets=None)
logits = out["logits"]
if logits.dim() == 3:
logits = logits[:, -1, :]
log_probs = torch.log_softmax(logits / temperature, dim=-1)
score = log_probs[0, target.item()].item()
return score
def choose_move(self, board, temperature=0.5):
"""Choose the best legal move using model scoring."""
legal_moves = list(board.legal_moves)
if not legal_moves:
return None
if len(legal_moves) == 1:
return legal_moves[0]
scores = {}
for move in legal_moves:
scores[move] = self._score_move(board, move, temperature)
best_move = max(scores, key=scores.get)
return best_move
def analyze(self, board, top_n=5, temperature=0.5):
"""Analyze position, return top N moves with scores."""
legal_moves = list(board.legal_moves)
if not legal_moves:
return []
scores = {}
for move in legal_moves:
scores[move] = self._score_move(board, move, temperature)
ranked = sorted(scores.items(), key=lambda x: -x[1])
return ranked[:top_n]
def choose_move_generative(self, board, temperature=0.5, num_attempts=10):
"""Choose move by generating tokens autoregressively, then filtering valid moves."""
board_tokens = encode_board(board)
inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device)
candidates = []
for _ in range(num_attempts):
with torch.no_grad():
# Generate first token
out = self.model(idx=inp, targets=None)
logits = out["logits"]
if logits.dim() == 3:
logits = logits[:, -1, :]
probs = torch.softmax(logits / temperature, dim=-1)
first_token = torch.multinomial(probs, 1).item()
# Generate second token
second_input = torch.cat([inp, torch.tensor([[first_token]], device=self.device)], dim=1)
out2 = self.model(idx=second_input, targets=None)
logits2 = out2["logits"]
if logits2.dim() == 3:
logits2 = logits2[:, -1, :]
probs2 = torch.softmax(logits2 / temperature, dim=-1)
second_token = torch.multinomial(probs2, 1).item()
move = decode_move([first_token, second_token], board)
if move is not None:
candidates.append(move)
if not candidates:
# Fallback to scoring
return self.choose_move(board, temperature)
# Return most common candidate
from collections import Counter
counts = Counter(candidates)
return counts.most_common(1)[0][0]
|