"""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, }