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