Parallax-Chess-Preview / vela_chess_engine.py
ParallaxOpen's picture
Upload vela_chess_engine.py with huggingface_hub
83c4df4 verified
Raw History Blame Contribute Delete
5.53 kB
#!/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]