| --- |
| license: mit |
| tags: |
| - mechanistic-interpretability |
| - forward-self-models |
| - activation-prediction |
| --- |
| |
| # Forward Self-Model Checkpoints |
|
|
| Checkpoints from [Forward Self-Models Learn an Empirical Approximation of Neural Network Computation](https://jagilley.github.io/forward-self-models.html). |
|
|
| A *forward self-model* is a small auxiliary network trained to predict a neural network's later-layer activations from its earlier-layer activations, learning an empirical approximation of the computational function that the intervening layers implement. |
|
|
| Code: [github.com/jagilley/forward-self-models](https://github.com/jagilley/forward-self-models) |
|
|
| ## Checkpoints |
|
|
| ### `llama-3.2-1B-layer8-forward-model/` |
|
|
| A 26.2M-parameter forward model (2.1% of Llama) predicting layer 7 → layer 8 of [Llama 3.2 1B](https://huggingface.co/meta-llama/Llama-3.2-1B). Achieves 0.937 cosine similarity with the target activations and 74% KL recovery in causal substitution. |
|
|
| - **Architecture**: 1-layer transformer, 1 attention head, d_head=128, SwiGLU MLP (hidden 4096) |
| - **Training**: MSE on 100M tokens of frozen Llama activations (FineWeb-Edu), lr=1e-4 |
| - **Files**: `fwd_model.pt` (weights), `config.json` (training config and metrics) |
|
|
| ### `toy-gpt-30M-layer1-forward-model/` |
|
|
| A 330K-parameter forward model (~1% of main model) predicting block 0 → block 1 of a 28.9M-parameter GPT-2. Achieves 0.972 cosine similarity and 94% KL recovery. |
|
|
| - **Architecture**: 1-layer transformer, 1 attention head, d_head=64, GELU MLP (hidden 512) |
| - **Main model**: 4-layer, 4-head, 256-dim GPT-2 trained on FineWeb-Edu (10M tokens) |
| - **Files**: `fwd_model.pt` (forward model weights), `main_model.pt` (main GPT weights), `config.json` (training config and metrics) |
| |
| ## Loading |
| |
| ```python |
| import torch |
| from forward_model import TransformerForwardModel |
|
|
| # Llama forward model |
| fwd = TransformerForwardModel( |
| d_model=2048, d_head=128, n_head=1, n_layer=1, |
| mlp_mult=2.0, block_size=2048, use_swiglu=True, |
| ) |
| fwd.load_state_dict(torch.load("llama-3.2-1B-layer8-forward-model/fwd_model.pt", map_location="cpu")) |
| |
| # Toy GPT forward model |
| fwd_toy = TransformerForwardModel( |
| d_model=256, d_head=64, n_head=1, n_layer=1, |
| mlp_mult=2, block_size=128, |
| ) |
| fwd_toy.load_state_dict(torch.load("toy-gpt-30M-layer1-forward-model/fwd_model.pt", map_location="cpu")) |
| ``` |
| |
| `forward_model.py` and `gpt_model.py` are included in the repo root for convenience. |
| |
| |