Download model.py from ohnah/micro500: direct link, hf CLI and curl.
- Browser
- Download file 967 Bytes
-
https://huggingface.co/ohnah/micro500/resolve/main/model.py
- Command line
-
hf download hf://ohnah/micro500/model.py
-
curl -L -o model.py https://huggingface.co/ohnah/micro500/resolve/main/model.py
967 Bytes
| import torch | |
| import torch.nn as nn | |
| class MicroLM(nn.Module): | |
| """Character-level recurrent LM with exactly 500 parameters.""" | |
| def __init__(self, V=27, d=4, h=8, pad=0): | |
| super().__init__() | |
| self.emb = nn.Embedding(V, d) | |
| self.ih = nn.Linear(d, h) | |
| self.hh = nn.Linear(h, h, bias=False) | |
| self.proj = nn.Linear(h, d) | |
| self.out_bias = nn.Parameter(torch.zeros(V)) | |
| # dead-weight padding to hit the exact parameter count (unused in forward) | |
| self.pad = nn.Parameter(torch.zeros(pad)) if pad > 0 else None | |
| self.h = h | |
| def forward(self, x): | |
| B, T = x.shape | |
| e = self.emb(x) | |
| hs = torch.zeros(B, self.h, device=x.device) | |
| outs = [] | |
| for t in range(T): | |
| hs = torch.tanh(self.ih(e[:, t]) + self.hh(hs)) | |
| outs.append(hs) | |
| z = self.proj(torch.stack(outs, dim=1)) | |
| return z @ self.emb.weight.T + self.out_bias # tied output layer |