Download forward_model.py from jspr/forward-self-model-checkpoints: direct link, hf CLI and curl.
- Browser
- Download file 4.01 kB
-
https://huggingface.co/jspr/forward-self-model-checkpoints/resolve/main/forward_model.py
- Command line
-
hf download hf://jspr/forward-self-model-checkpoints/forward_model.py
-
curl -L -o forward_model.py https://huggingface.co/jspr/forward-self-model-checkpoints/resolve/main/forward_model.py
4.01 kB
| """Forward models for predicting transformer activations. | |
| ForwardModel: per-position MLP. Structurally blind to cross-position effects. | |
| TransformerForwardModel: small transformer with capacity bottleneck. | |
| """ | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| class ForwardModel(nn.Module): | |
| def __init__(self, d_model: int, hidden_mult: int = 2): | |
| super().__init__() | |
| hidden = d_model * hidden_mult | |
| self.net = nn.Sequential( | |
| nn.Linear(d_model, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, d_model), | |
| ) | |
| n_params = sum(p.numel() for p in self.parameters()) | |
| print(f"ForwardModel: {n_params/1e3:.1f}K parameters " | |
| f"(d_model={d_model}, hidden={hidden})") | |
| def forward(self, x): | |
| return self.net(x) | |
| class ForwardBlock(nn.Module): | |
| def __init__(self, d_model: int, d_head: int, n_head: int, mlp_mult: float, | |
| use_swiglu: bool = False): | |
| super().__init__() | |
| self.d_head = d_head | |
| self.n_head = n_head | |
| self.use_swiglu = use_swiglu | |
| self.ln1 = nn.LayerNorm(d_model) | |
| self.q_proj = nn.Linear(d_model, d_head * n_head) | |
| self.k_proj = nn.Linear(d_model, d_head * n_head) | |
| self.v_proj = nn.Linear(d_model, d_head * n_head) | |
| self.out_proj = nn.Linear(d_head * n_head, d_model) | |
| self.ln2 = nn.LayerNorm(d_model) | |
| mlp_hidden = int(d_model * mlp_mult) | |
| if use_swiglu: | |
| self.gate_proj = nn.Linear(d_model, mlp_hidden) | |
| self.up_proj = nn.Linear(d_model, mlp_hidden) | |
| self.down_proj = nn.Linear(mlp_hidden, d_model) | |
| else: | |
| self.mlp = nn.Sequential( | |
| nn.Linear(d_model, mlp_hidden), | |
| nn.GELU(), | |
| nn.Linear(mlp_hidden, d_model), | |
| ) | |
| def forward(self, x, causal_mask): | |
| B, T, C = x.size() | |
| h = self.ln1(x) | |
| q = self.q_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2) | |
| k = self.k_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2) | |
| v = self.v_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2) | |
| att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.d_head)) | |
| att = att.masked_fill(causal_mask[:, :, :T, :T] == 0, float("-inf")) | |
| att = F.softmax(att, dim=-1) | |
| y = att @ v | |
| y = y.transpose(1, 2).contiguous().view(B, T, self.d_head * self.n_head) | |
| x = x + self.out_proj(y) | |
| h2 = self.ln2(x) | |
| if self.use_swiglu: | |
| x = x + self.down_proj(F.silu(self.gate_proj(h2)) * self.up_proj(h2)) | |
| else: | |
| x = x + self.mlp(h2) | |
| return x | |
| class TransformerForwardModel(nn.Module): | |
| def __init__(self, d_model: int, d_head: int = 64, n_head: int = 1, | |
| n_layer: int = 1, mlp_mult: float = 2, block_size: int = 128, | |
| causal: bool = True, use_swiglu: bool = False): | |
| super().__init__() | |
| self.d_model = d_model | |
| if causal: | |
| mask = torch.tril(torch.ones(block_size, block_size)) | |
| else: | |
| mask = torch.ones(block_size, block_size) | |
| self.register_buffer("causal_mask", mask.view(1, 1, block_size, block_size)) | |
| self.blocks = nn.ModuleList([ | |
| ForwardBlock(d_model, d_head, n_head, mlp_mult, use_swiglu=use_swiglu) | |
| for _ in range(n_layer) | |
| ]) | |
| mlp_hidden = int(d_model * mlp_mult) | |
| mlp_type = "SwiGLU" if use_swiglu else "GELU" | |
| n_params = sum(p.numel() for p in self.parameters()) | |
| print(f"TransformerForwardModel: {n_params/1e3:.1f}K parameters " | |
| f"(d_model={d_model}, d_head={d_head}, n_head={n_head}, " | |
| f"n_layer={n_layer}, mlp_hidden={mlp_hidden}, mlp={mlp_type}" | |
| f"{', bidirectional' if not causal else ''})") | |
| def forward(self, x): | |
| for block in self.blocks: | |
| x = block(x, self.causal_mask) | |
| return x | |