Spaces:
Running on Zero
Running on Zero
Download model.py from ansarzeinulla/Bestemshe-God-Algorithm: direct link, hf CLI and curl.
- Browser
- Download file 2.71 kB
-
https://huggingface.co/spaces/ansarzeinulla/Bestemshe-God-Algorithm/resolve/main/model.py
- Command line
-
hf download hf://spaces/ansarzeinulla/Bestemshe-God-Algorithm/model.py
-
curl -L -o model.py https://huggingface.co/spaces/ansarzeinulla/Bestemshe-God-Algorithm/resolve/main/model.py
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) | |
| 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, | |
| } | |