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]