Download python/model.py from candypunk/NanoJev-Web: direct link, hf CLI and curl.
- Browser
- Download file 3.27 kB
-
https://huggingface.co/candypunk/NanoJev-Web/resolve/main/python/model.py
- Command line
-
hf download hf://candypunk/NanoJev-Web/python/model.py
-
curl -L -o model.py https://huggingface.co/candypunk/NanoJev-Web/resolve/main/python/model.py
3.27 kB
| """Decision network extracted from NanoJev; no trainer or game imports.""" | |
| import torch | |
| from torch import nn | |
| import torch.nn.functional as F | |
| class DecisionModel(nn.Module): | |
| def __init__(self, backbone, set_head): | |
| super().__init__() | |
| self.backbone = backbone | |
| hidden = backbone.config.hidden_size | |
| self.norm = nn.LayerNorm(hidden) | |
| self.scalar = nn.Linear(hidden, 1) # Nonzero random initialization avoids a dead first step. | |
| nn.init.normal_(self.scalar.weight, std=0.02) | |
| nn.init.zeros_(self.scalar.bias) | |
| self.set_head = set_head | |
| if set_head == 'attention': | |
| self.set_project = nn.Linear(hidden + 1, 128) | |
| self.set_attention = nn.MultiheadAttention(128, 4, dropout=0.0, batch_first=True) | |
| self.set_output = nn.Linear(128, 1) | |
| # Only the final residual projection starts at zero; its upstream layers are nonzero. | |
| nn.init.zeros_(self.set_output.weight) | |
| nn.init.zeros_(self.set_output.bias) | |
| def forward(self, examples, pad_token): | |
| paths = [ids for ex in examples for ids in ex['leaf_tokens']] | |
| device = self.scalar.weight.device | |
| lengths = torch.tensor([len(ids) for ids in paths], device=device) | |
| width = int(lengths.max()) | |
| tokens = torch.full((len(paths), width), pad_token, dtype=torch.long, device=device) | |
| for i, ids in enumerate(paths): | |
| tokens[i, :len(ids)] = torch.tensor(ids, device=device) | |
| attention = torch.arange(width, device=device)[None, :] < lengths[:, None] | |
| hidden = self.backbone(input_ids=tokens, attention_mask=attention, | |
| use_cache=False).last_hidden_state | |
| leaves = hidden[torch.arange(len(paths), device=device), lengths-1] | |
| kmax = max(len(ex['candidate_ids']) for ex in examples) | |
| h = leaves.new_zeros((len(examples), kmax, leaves.shape[-1])) | |
| valid = torch.zeros((len(examples), kmax), dtype=torch.bool, device=device) | |
| offset = 0 | |
| for i, ex in enumerate(examples): | |
| n = len(ex['leaf_tokens']) | |
| h[i, :n] = leaves[offset:offset+n] | |
| valid[i, :len(ex['candidate_ids'])] = True | |
| offset += n | |
| h = self.norm(h) | |
| z = self.scalar(h).squeeze(-1).float() | |
| choice = torch.tensor([i for i, ex in enumerate(examples) if ex['type'] == 'choice'], device=device) | |
| if self.set_head == 'attention' and len(choice): | |
| log_k = valid[choice].sum(-1).float().log()[:, None, None].expand(-1, kmax, 1) | |
| u = self.set_project(torch.cat([h[choice], log_k.to(h.dtype)], dim=-1)) | |
| mixed, _ = self.set_attention(u, u, u, key_padding_mask=~valid[choice], need_weights=False) | |
| delta = self.set_output(torch.tanh(u + mixed)).squeeze(-1).float() | |
| z = z.index_add(0, choice, delta) | |
| # Boolean has one semantic path and one scalar, representing logits [0,z]. | |
| out = [] | |
| for i, ex in enumerate(examples): | |
| if ex['type'] == 'boolean': | |
| out.append(F.pad(torch.stack([z[i, 0] * 0, z[i, 0]]), (0, kmax-2))) | |
| else: | |
| out.append(z[i]) | |
| return torch.stack(out).masked_fill(~valid, -1e9), valid | |