ansarzeinulla's picture
proto
25faada
Raw History Blame Contribute Delete
2.71 kB
"""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,
}