File size: 2,709 Bytes
25faada
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Loads and runs the Bestemshe ResMLP neural network from Hugging Face.

Needs only: pip install torch huggingface_hub
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from huggingface_hub import hf_hub_download


class ResBlock(nn.Module):
    def __init__(self, w):
        super().__init__()
        self.norm = nn.LayerNorm(w)
        self.fc1 = nn.Linear(w, w)
        self.fc2 = nn.Linear(w, w)

    def forward(self, x):
        return x + self.fc2(F.gelu(self.fc1(self.norm(x))))


class BestemsheNet(nn.Module):
    """Embeddings -> Residual MLP -> Value head (loss/draw/win) + Policy head (5 moves)."""

    def __init__(self, width=1024, blocks=8, emb=32):
        super().__init__()
        self.pit_emb = nn.Embedding(51, emb)   # stones per pit: 0..50
        self.kaz_emb = nn.Embedding(13, emb)   # kazan // 2: 0..12
        self.inp = nn.Linear(12 * emb, width)
        self.body = nn.Sequential(*[ResBlock(width) for _ in range(blocks)])
        self.value_head = nn.Linear(width, 3)
        self.policy_head = nn.Linear(width, 5)

    def forward(self, pits, kaz):
        x = torch.cat([self.pit_emb(pits).flatten(1),
                       self.kaz_emb(kaz).flatten(1)], dim=1)
        h = self.body(F.gelu(self.inp(x)))
        return self.value_head(h), self.policy_head(h)


class BestemsheAI:
    def __init__(self, repo_id="ansarzeinulla/bestemshe-resmlp", filename="latest.pt", device="cpu"):
        self.device = device
        path = hf_hub_download(repo_id=repo_id, filename=filename)
        ckpt = torch.load(path, map_location="cpu")
        args = ckpt["args"]
        self.model = BestemsheNet(args["width"], args["blocks"]).to(device).eval()

        # Strip torch.compile() prefix if present
        state_dict = {k.replace("_orig_mod.", ""): v for k, v in ckpt["model"].items()}
        self.model.load_state_dict(state_dict)
        self.step = ckpt.get("step", 0)

    @torch.no_grad()
    def predict(self, pits, kazan_self, kazan_opp):
        """pits: 10 ints (mover's 5 pits, then opponent's 5). kazan_*: even ints 0..24."""
        pits_t = torch.tensor([pits], dtype=torch.long, device=self.device)
        kaz_t = torch.tensor([[kazan_self // 2, kazan_opp // 2]], dtype=torch.long, device=self.device)

        v_logits, p_logits = self.model(pits_t, kaz_t)
        probs = F.softmax(v_logits[0], dim=0).tolist()
        policy = p_logits[0].tolist()

        return {
            "p_loss": probs[0],
            "p_draw": probs[1],
            "p_win": probs[2],
            "predicted_class": int(v_logits[0].argmax()),
            "best_move": int(torch.tensor(policy).argmax()),  # 0..4
            "policy_logits": policy,
        }