File size: 2,883 Bytes
a5f2e46
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from pathlib import Path
import json
import torch
import chess
from safetensors import safe_open
from chesstransformer.models.transformer.position2move import Position2MoveModel
from chesstransformer.models.tokenizer import PostionTokenizer, MoveTokenizer

data_foler = Path(__file__).parents[1] / "data"


class Position2MoveBot:
    def __init__(
        self,
        model_path: str = str((data_foler / "models/position2moveV2.1/best_model/model.safetensors").resolve()),
        device: str = "cpu",
    ):
        self.device = device
        self.position_tokenizer = PostionTokenizer()
        self.move_tokenizer = MoveTokenizer()

        config_path = Path(model_path).parent / "config.json"

        with open(config_path, "r") as f:
            self.config = json.load(f)

        # Load model
        self.model = Position2MoveModel(**self.config).to(device)

        with safe_open(model_path, framework="pt", device=device) as f:
            state_dict = {k: f.get_tensor(k) for k in f.keys()}

            # Handle compiled model prefix (_orig_mod.)
            if any(k.startswith("_orig_mod.") for k in state_dict.keys()):
                state_dict = {k.replace("_orig_mod.", ""): v for k, v in state_dict.items()}

            self.model.load_state_dict(state_dict)

        self.model.eval()

    @torch.no_grad()
    def predict(self, board: chess.Board):
        tokens_ids = self.position_tokenizer.encode(board)
        torch_input = torch.tensor(tokens_ids).unsqueeze(0).long().to(self.device)  # Add batch dimension

        is_white = board.turn
        is_white = torch.tensor([is_white]).bool().to(self.device)  # Add batch dimension

        logits = self.model(torch_input, is_white)

        legal_moves = [m.uci() for m in board.legal_moves]

        mask = torch.full((logits.size(-1),), float("-inf")).to(self.device)
        for move in legal_moves:
            move_id = self.move_tokenizer.encode(move)
            mask[move_id] = 0.0

        masked_logits = logits[0] + mask  # Assuming batch size of 1
        probs = torch.softmax(masked_logits, dim=-1)
        # draw the move from the distribution
        predicted_move_id = torch.multinomial(probs * 0.2, num_samples=1).item()
        proba = probs[predicted_move_id].item()

        predicted_move = self.move_tokenizer.decode(predicted_move_id)
        return predicted_move, proba


if __name__ == "__main__":
    bot = Position2MoveBot()

    board = chess.Board()

    print("Initial board:")
    print(board)

    move, proba = bot.predict(board)
    print(f"Bot suggests move: {move} with probability {proba:.4f}")

    board.push_uci(move)
    print("Board after bot move:")
    print(board)

    move, proba = bot.predict(board)
    print(f"Bot suggests move: {move} with probability {proba:.4f}")

    board.push_uci(move)
    print("Board after bot move:")
    print(board)