jspr's picture
Minor cleanup
86cd27c verified
|
Raw
History Blame Contribute Delete
2.43 kB
---
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.