Views
No views yet
output = W_left @ (x ⊙ (W_right @ x)) + bias
| Layer | Final FVU | Variance Explained | Notes |
|---|---|---|---|
| 0 | 0.0075 | 99.2% | Best performance |
| 1-2 | 0.167-0.174 | 82.6-83.2% | Hardest to approximate |
| 3-22 | 0.037-0.066 | 93.4-96.3% | Consistent performance |
| 23 | 0.0259 | 97.4% | Second-best |
.
├── layer_0/
│ ├── transcoder_weights_l0_bilinear_muon_3000b.pt
│ └── config.yaml
├── layer_1/
│ ├── transcoder_weights_l1_bilinear_muon_3000b.pt
│ └── config.yaml
...
├── layer_23/
│ ├── transcoder_weights_l23_bilinear_muon_3000b.pt
│ └── config.yaml
├── figures/
│ ├── all_layers_comparison.png
│ ├── training_curves_overlaid_layers_0_5.png
│ ├── training_curves_overlaid_layers_6_11.png
│ ├── training_curves_overlaid_layers_12_17.png
│ └── training_curves_overlaid_layers_18_23.png
└── README.md1import torch
2from transformers import AutoModelForCausalLM, AutoTokenizer
3
4# Load base model
5model = AutoModelForCausalLM.from_pretrained("EleutherAI/pythia-410m")
6tokenizer = AutoTokenizer.from_pretrained("EleutherAI/pythia-410m")
7
8# Load transcoder for layer 3
9layer_idx = 3
10checkpoint = torch.load(f"layer_{layer_idx}/transcoder_weights_l{layer_idx}_bilinear_muon_3000b.pt")
11
12# Extract configuration
13config = checkpoint['config']
14print(f"Input dim: {config.n_inputs}")
15print(f"Hidden dim: {config.n_hidden}")
16print(f"Output dim: {config.n_outputs}")
17
18# Reconstruct model (example - you'll need the Bilinear class)
19class Bilinear(torch.nn.Module):
20 def __init__(self, n_inputs, n_hidden, n_outputs, bias=True):
21 super().__init__()
22 self.W_left = torch.nn.Linear(n_hidden, n_outputs, bias=bias)
23 self.W_right = torch.nn.Linear(n_inputs, n_hidden, bias=False)
24
25 def forward(self, x):
26 right = self.W_right(x)
27 hadamard = x.unsqueeze(-1) * right.unsqueeze(-2)
28 return self.W_left(hadamard.sum(dim=-2))
29
30transcoder = Bilinear(config.n_inputs, config.n_hidden, config.n_outputs, config.bias)
31transcoder.load_state_dict(checkpoint['model_state_dict'])
32transcoder.eval()
33
34# Use transcoder to approximate MLP
35with torch.no_grad():
36 # Get MLP input from layer 3
37 inputs = tokenizer("Hello world", return_tensors="pt")
38 outputs = model(**inputs, output_hidden_states=True)
39 mlp_input = outputs.hidden_states[layer_idx] # Before MLP
40
41 # Approximate MLP output with transcoder
42 transcoded_output = transcoder(mlp_input).pt file) contains:model_state_dict: Model weightsoptimizer_state_dict: Optimizer stateconfig: Configuration object with dimensionsmse_losses: List of MSE losses per batchvariance_explained: List of variance explained per batchfvu_values: List of FVU values per batchlayer_idx: Layer index (0-23)d_model: Model dimension (1024)1@misc{pythia410m-bilinear-transcoders,
2 title={Bilinear MLP Transcoders for Pythia-410m},
3 author={[Your Name]},
4 year={2025},
5 publisher={Hugging Face},
6 url={https://huggingface.co/[your-username]/pythia-410m-bilinear-transcoders}
7}