sae-lens 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
sae_lens-0.1.0/LICENSE ADDED
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2023 Joseph Bloom
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,263 @@
1
+ Metadata-Version: 2.1
2
+ Name: sae-lens
3
+ Version: 0.1.0
4
+ Summary: Training and Analyzing Sparse Autoencoders (SAEs)
5
+ Author: Joseph Bloom
6
+ Requires-Python: >=3.10,<4.0
7
+ Classifier: Programming Language :: Python :: 3
8
+ Classifier: Programming Language :: Python :: 3.10
9
+ Classifier: Programming Language :: Python :: 3.11
10
+ Classifier: Programming Language :: Python :: 3.12
11
+ Requires-Dist: babe (>=0.0.7,<0.0.8)
12
+ Requires-Dist: datasets (>=2.17.1,<3.0.0)
13
+ Requires-Dist: ipykernel (>=6.29.2,<7.0.0)
14
+ Requires-Dist: jupyter (>=1.0.0,<2.0.0)
15
+ Requires-Dist: matplotlib (>=3.8.3,<4.0.0)
16
+ Requires-Dist: matplotlib-inline (>=0.1.6,<0.2.0)
17
+ Requires-Dist: mkdocs (>=1.5.3,<2.0.0)
18
+ Requires-Dist: mkdocs-autorefs (>=1.0.1,<2.0.0)
19
+ Requires-Dist: mkdocs-material (>=9.5.15,<10.0.0)
20
+ Requires-Dist: mkdocs-section-index (>=0.3.8,<0.4.0)
21
+ Requires-Dist: mkdocstrings (>=0.24.1,<0.25.0)
22
+ Requires-Dist: mkdocstrings-python (>=1.9.0,<2.0.0)
23
+ Requires-Dist: nbformat (>=5.9.2,<6.0.0)
24
+ Requires-Dist: nltk (>=3.8.1,<4.0.0)
25
+ Requires-Dist: plotly (>=5.19.0,<6.0.0)
26
+ Requires-Dist: plotly-express (>=0.4.1,<0.5.0)
27
+ Requires-Dist: sae-vis (==0.2.6)
28
+ Requires-Dist: transformer-lens (>=1.14.0,<2.0.0)
29
+ Requires-Dist: transformers (>=4.38.1,<5.0.0)
30
+ Description-Content-Type: text/markdown
31
+
32
+ <img width="1308" alt="Screenshot 2024-03-21 at 3 08 28 pm" src="https://github.com/jbloomAus/mats_sae_training/assets/69127271/209012ec-a779-4036-b4be-7b7739ea87f6">
33
+
34
+ # SAE Lens
35
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
36
+ [![build](https://github.com/jbloomAus/mats_sae_training/actions/workflows/build.yml/badge.svg)](https://github.com/jbloomAus/mats_sae_training/actions/workflows/build.yml)
37
+ [![Deploy Docs](https://github.com/jbloomAus/mats_sae_training/actions/workflows/deploy_docs.yml/badge.svg)](https://github.com/jbloomAus/mats_sae_training/actions/workflows/deploy_docs.yml)
38
+ [![codecov](https://codecov.io/gh/jbloomAus/mats_sae_training/graph/badge.svg?token=N83NGH8CGE)](https://codecov.io/gh/jbloomAus/mats_sae_training)
39
+
40
+ SAELens exists to help researchers:
41
+ - Train sparse autoencoders.
42
+ - Analyse sparse autoencoders / research mechanistic interpretability.
43
+ - Generate insights which make it easier to create safe and aligned AI systems.
44
+
45
+ ## Quick Start
46
+
47
+ ### Set Up
48
+
49
+ This project uses [Poetry](https://python-poetry.org/) for dependency management. Ensure Poetry is installed, then to install the dependencies, run:
50
+
51
+ ```
52
+ poetry install
53
+ ```
54
+
55
+ ### Loading Sparse Autoencoders from Huggingface
56
+
57
+ [Previously trained sparse autoencoders](https://huggingface.co/jbloom/GPT2-Small-SAEs) can be loaded from huggingface with close to single line of code. For more details and performance metrics for these sparse autoencoder, read my [blog post](https://www.alignmentforum.org/posts/f9EgfLSurAiqRJySD/open-source-sparse-autoencoders-for-all-residual-stream).
58
+
59
+ ```python
60
+ import torch
61
+ from sae_lens import LMSparseAutoencoderSessionloader
62
+ from huggingface_hub import hf_hub_download
63
+
64
+ layer = 8 # pick a layer you want.
65
+ REPO_ID = "jbloom/GPT2-Small-SAEs"
66
+ FILENAME = f"final_sparse_autoencoder_gpt2-small_blocks.{layer}.hook_resid_pre_24576.pt"
67
+ path = hf_hub_download(repo_id=REPO_ID, filename=FILENAME)
68
+ model, sparse_autoencoder, activation_store = LMSparseAutoencoderSessionloader.load_session_from_pretrained(
69
+ path = path
70
+ )
71
+ sparse_autoencoder.eval()
72
+ ```
73
+
74
+ You can also load the feature sparsity from huggingface.
75
+
76
+ ```python
77
+ FILENAME = f"final_sparse_autoencoder_gpt2-small_blocks.{layer}.hook_resid_pre_24576_log_feature_sparsity.pt"
78
+ path = hf_hub_download(repo_id=REPO_ID, filename=FILENAME)
79
+ log_feature_sparsity = torch.load(path, map_location=sparse_autoencoder.cfg.device)
80
+
81
+ ```
82
+ ### Background
83
+
84
+ We highly recommend this [tutorial](https://www.lesswrong.com/posts/LnHowHgmrMbWtpkxx/intro-to-superposition-and-sparse-autoencoders-colab).
85
+
86
+
87
+ ## High Level
88
+
89
+ ### Motivation
90
+
91
+ - **Accelerate SAE Research**: Support fast experimentation to understand SAEs and improve SAE training so we can train SAEs on larger and more diverse models.
92
+ - **Make Research like Play**: Support research into language model internals via SAEs. Good tooling can make research tremendously exciting and enjoyable. Balancing modifiability and reliability with ease of understanding / access is the name of the game here.
93
+ - **Build an awesome community**: Mechanistic Interpretability already has an awesome community but as that community grows, it makes sense that there will be niches. I'd love to build a great community around Sparse Autoencoders.
94
+
95
+ ### Goals
96
+
97
+ #### **SAE Training**: SAE Training features will fit into a number of categories including:
98
+ - **Making it easy to train SAEs**: Training SAEs is hard for a number of reasons and so making it easy for people to train SAEs with relatively little expertise seems like the main way this codebase will create value.
99
+ - **Training SAEs on more models**: Supporting training of SAEs on more models, architectures, different activations within those models.
100
+ - **Being better at training SAEs**: Enabling methodological changes which may improve SAE performance as measured by reconstruction loss, Cross Entropy Loss when using reconstructed activation, L1 loss, L0 and interpretability of features as well as improving speed of training or reducing the compute resources required to train SAEs.
101
+ - **Being better at measuring SAE Performance**: How do we know when SAEs are doing what we want them to? Improving training metrics should allow better decisions about which methods to use and which hyperparameters choices we make.
102
+ - **Training SAE variants**: People are already training “Transcoders” which map from one activation to another (such as before / after an MLP layer). These can be easily supported with a few changes. Other variants will come in time and
103
+
104
+ #### **Analysis with SAEs**: Using SAEs to understand neural network internals is an exciting, but complicated task.
105
+ - **Feature-wise Interpretability**: This looks something like "for each feature, have as much knowledge about it as possible". Part of this will feature dashboard improvements, or supporting better integrations with Neuronpedia.
106
+ - **Mechanistic Interpretability**: This comprises the more traditional kinds of Mechanistic Interpretability which TransformerLens supports and should be supported by this codebase. Making it easy to patch, ablate or otherwise intervene on features so as to find circuits will likely speed up lots of researchers.
107
+
108
+ ### Other Stuff
109
+
110
+ I think there are lots of other types of analysis that could be done in the future with SAE features. I've already explored many different types of statistical tests which can reveal interesting properties of features. There are also things like saliency mapping and attribution techniques which it would be nice to support.
111
+ - Accessibility and Code Quality: The codebase won’t be used if it doesn’t work and it also won’t get used if it’s too hard to understand, modify or read.
112
+ Making the code accessible: This involves tasks like turning the code base into a python package.
113
+ - Knowing how the code is supposed to work: Is the code well-documented? This will require docstrings, tutorials and links to related work and publications. Getting aligned on what the code does is critical to sharing a resource like this.
114
+ - Knowing the code works as intended: All code should be tested. Unit tests and acceptance tests are both important.
115
+ - Knowing the code is actually performant: This will ensure code works as intended. However deep learning introduces lots of complexity which makes actually running benchmarks essential to having confidence in the code.
116
+
117
+
118
+ ## Code Overview
119
+
120
+ The codebase contains 2 folders worth caring about:
121
+
122
+ - training: The main body of the code is here. Everything required for training SAEs.
123
+ - analysis: This code is mainly house the feature visualizer code we use to generate dashboards. It was written by Callum McDougal but I've ported it here with permission and edited it to work with a few different activation types.
124
+
125
+ Some other folders:
126
+
127
+ - tutorials: These aren't well maintained but I'll aim to clean them up soon.
128
+ - tests: When first developing the codebase, I was writing more tests. I have no idea whether they are currently working!
129
+
130
+ I've been commiting my research code to the `Research` folder but am not expecting other people use or look at that.
131
+
132
+
133
+ ### Training your own Sparse Autoencoder
134
+
135
+ Sparse Autoencoders can be intimidating at first but it's fairly simple to train one once you know what each part of the config does. I've created a config class which you instantiate and pass to the runner which will complete your training run and log it's progress to wandb.
136
+
137
+ Let's go through the major components of the config:
138
+ - Data: SAE's autoencode model activations. We need to specify the model, the part of the models activations we want to autoencode and the dataset the model is operating on when generating those activations. We now automatically detect if that dataset is tokenized and most huggingface datasets should be fine. One slightly annoying detail is that you need to know the dimensionality of those activations when contructing your SAE but you can get that in the transformerlens [docs](https://neelnanda-io.github.io/TransformerLens/generated/model_properties_table.html). Any language model in the table from those docs should work.
139
+ - SAE Parameters: Your expansion factor will determine the size of your SAE and the decoder bias initialization method should always be geometric_median or mean. Mean is faster but theoretically sub-optimal. I use another package to get the geometric median and it can be quite slow.
140
+ - Training Parameters: These are most critical. The right L1 coefficient (coefficient in the activation sparsity inducing term in the loss) changes with your learning rate but a good bet would be to use LR 4e-4 and L1 8e-5 for GPT2 small. These will vary for other models and playing around with them / short runs can be helpful. Training batch size of 4096 is standard and I'm not really sure whether there's benefit to playing with it. In theory a larger context size (one accurate to whatever the model was trained with) seems good but it's computationally cheaper to use 128. Learning rate warm up is important to avoid dead neurons.
141
+ - Activation Store Parameters: The activation store shuffles activations from forward passes over samples from your data. The larger it is, the better shuffling you'll get. In theory more shuffling is good. The total training tokens is a very important parameter. The more the better, but you'll often see good results having trained on a few hundred million tokens. Store batch batch size is a function of your gpu and how many forward passes of your model you want to do simultaneously when collecting activations.
142
+ - Dead Neurons / Sparsity Metrics: The config around resampling was more important when we were using resampling to avoid dead neurons (see Anthropic's post on this), but using ghost gradients, the resampling protcol is much simpler. I'd always set ghost grad to True and feature sampling method to None. The feature sampling window effects the dashboard statistics tracking feature occurence and the dead feature window tracks how many forward passes a neuron must not activate before we apply ghost grads to it.
143
+ - WANDB: Fairly straightfoward. Don't set log frequency too high or your dashboard will be slow!
144
+ - Device: I can run this code on my macbook with "mps" but mostly do runs with cuda.
145
+ - Dtype: Float16 maybe could work but I had some funky results and have left it at float32 for the time being.
146
+ - Checkpoints: I'd collected checkpoints on runs you care about but turn them off when tuning since it can be slow.
147
+
148
+
149
+ ```python
150
+ import torch
151
+ import os
152
+ import sys
153
+
154
+ os.environ["TOKENIZERS_PARALLELISM"] = "false"
155
+ os.environ["WANDB__SERVICE_WAIT"] = "300"
156
+
157
+ from sae_lens.training.config import LanguageModelSAERunnerConfig
158
+ from sae_lens.training.lm_runner import language_model_sae_runner
159
+
160
+ cfg = LanguageModelSAERunnerConfig(
161
+
162
+ # Data Generating Function (Model + Training Distibuion)
163
+ model_name = "gpt2-small",
164
+ hook_point = "blocks.2.hook_resid_pre",
165
+ hook_point_layer = 2,
166
+ d_in = 768,
167
+ dataset_path = "Skylion007/openwebtext",
168
+ is_dataset_tokenized=False,
169
+
170
+ # SAE Parameters
171
+ expansion_factor = 64,
172
+ b_dec_init_method = "geometric_median",
173
+
174
+ # Training Parameters
175
+ lr = 0.0004,
176
+ l1_coefficient = 0.00008,
177
+ lr_scheduler_name="constantwithwarmup",
178
+ train_batch_size = 4096,
179
+ context_size = 128,
180
+ lr_warm_up_steps=5000,
181
+
182
+ # Activation Store Parameters
183
+ n_batches_in_buffer = 128,
184
+ total_training_tokens = 1_000_000 * 300,
185
+ store_batch_size = 32,
186
+
187
+ # Dead Neurons and Sparsity
188
+ use_ghost_grads=True,
189
+ feature_sampling_window = 1000,
190
+ dead_feature_window=5000,
191
+ dead_feature_threshold = 1e-6,
192
+
193
+ # WANDB
194
+ log_to_wandb = True,
195
+ wandb_project= "mats_sae_training_gpt2",
196
+ wandb_entity = None,
197
+ wandb_log_frequency=100,
198
+
199
+ # Misc
200
+ device = "cuda",
201
+ seed = 42,
202
+ n_checkpoints = 10,
203
+ checkpoint_path = "checkpoints",
204
+ dtype = torch.float32,
205
+ )
206
+
207
+ sparse_autoencoder = language_model_sae_runner(cfg)
208
+
209
+ ```
210
+
211
+
212
+ ## Loading a Pretrained Language Model
213
+
214
+ Once your SAE is trained, the final SAE weights will be saved to wandb and are loadable via the session loader. The session loader will return:
215
+ - The model your SAE was trained on (presumably you're interested in studying this. It's always a HookedTransformer)
216
+ - Your SAE.
217
+ - An activations loader: from which you can get randomly sampled activations or batches of tokens from the dataset you used to train the SAE. (more on this in the tutorial)
218
+
219
+ ```python
220
+ from sae_lens import LMSparseAutoencoderSessionloader
221
+
222
+ path ="path/to/sparse_autoencoder.pt"
223
+ model, sparse_autoencoder, activations_loader = LMSparseAutoencoderSessionloader.load_session_from_pretrained(
224
+ path
225
+ )
226
+
227
+ ```
228
+ ## Tutorials
229
+
230
+ I wrote a tutorial to show users how to do some basic exploration of their SAE.
231
+ - `evaluating_your_sae.ipynb`: A quick/dirty notebook showing how to check L0 and Prediction loss with your SAE, as well as showing how to generate interactive dashboards using Callum's reporduction of [Anthropics interface](https://transformer-circuits.pub/2023/monosemantic-features#setup-interface).
232
+ - `logits_lens_with_features.ipynb`: A notebook showing how to reproduce the analysis from this [LessWrong post](https://www.lesswrong.com/posts/qykrYY6rXXM7EEs8Q/understanding-sae-features-with-the-logit-lens).
233
+
234
+ ## Example Dashboard
235
+
236
+ WandB Dashboards provide lots of useful insights while training SAE's. Here's a screenshot from one training run.
237
+
238
+ ![screenshot](content/dashboard_screenshot.png)
239
+
240
+
241
+ ## Example Output
242
+
243
+ Here's one feature we found in the residual stream of Layer 10 of GPT-2 Small:
244
+
245
+ ![alt text](content/readme_screenshot_predict_pronoun_feature.png). Open `gpt2_resid_pre10_predict_pronoun_feature.html` in your browser to interact with the dashboard (WIP).
246
+
247
+ Note, probably this feature could split into more mono-semantic features in a larger SAE that had been trained for longer. (this was was only about 49152 features trained on 10M tokens from OpenWebText).
248
+
249
+
250
+ ## Citations and References:
251
+
252
+ Research:
253
+ - [Towards Monosemanticy](https://transformer-circuits.pub/2023/monosemantic-features)
254
+ - [Sparse Autoencoders Find Highly Interpretable Features in Language Model](https://arxiv.org/abs/2309.08600)
255
+
256
+
257
+
258
+ Reference Implementations:
259
+ - [Neel Nanda](https://github.com/neelnanda-io/1L-Sparse-Autoencoder)
260
+ - [AI-Safety-Foundation](https://github.com/ai-safety-foundation/sparse_autoencoder).
261
+ - [Arthur Conmy](https://github.com/ArthurConmy/sae).
262
+ - [Callum McDougall](https://github.com/callummcdougall/sae-exercises-mats/tree/main)
263
+
@@ -0,0 +1,231 @@
1
+ <img width="1308" alt="Screenshot 2024-03-21 at 3 08 28 pm" src="https://github.com/jbloomAus/mats_sae_training/assets/69127271/209012ec-a779-4036-b4be-7b7739ea87f6">
2
+
3
+ # SAE Lens
4
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
5
+ [![build](https://github.com/jbloomAus/mats_sae_training/actions/workflows/build.yml/badge.svg)](https://github.com/jbloomAus/mats_sae_training/actions/workflows/build.yml)
6
+ [![Deploy Docs](https://github.com/jbloomAus/mats_sae_training/actions/workflows/deploy_docs.yml/badge.svg)](https://github.com/jbloomAus/mats_sae_training/actions/workflows/deploy_docs.yml)
7
+ [![codecov](https://codecov.io/gh/jbloomAus/mats_sae_training/graph/badge.svg?token=N83NGH8CGE)](https://codecov.io/gh/jbloomAus/mats_sae_training)
8
+
9
+ SAELens exists to help researchers:
10
+ - Train sparse autoencoders.
11
+ - Analyse sparse autoencoders / research mechanistic interpretability.
12
+ - Generate insights which make it easier to create safe and aligned AI systems.
13
+
14
+ ## Quick Start
15
+
16
+ ### Set Up
17
+
18
+ This project uses [Poetry](https://python-poetry.org/) for dependency management. Ensure Poetry is installed, then to install the dependencies, run:
19
+
20
+ ```
21
+ poetry install
22
+ ```
23
+
24
+ ### Loading Sparse Autoencoders from Huggingface
25
+
26
+ [Previously trained sparse autoencoders](https://huggingface.co/jbloom/GPT2-Small-SAEs) can be loaded from huggingface with close to single line of code. For more details and performance metrics for these sparse autoencoder, read my [blog post](https://www.alignmentforum.org/posts/f9EgfLSurAiqRJySD/open-source-sparse-autoencoders-for-all-residual-stream).
27
+
28
+ ```python
29
+ import torch
30
+ from sae_lens import LMSparseAutoencoderSessionloader
31
+ from huggingface_hub import hf_hub_download
32
+
33
+ layer = 8 # pick a layer you want.
34
+ REPO_ID = "jbloom/GPT2-Small-SAEs"
35
+ FILENAME = f"final_sparse_autoencoder_gpt2-small_blocks.{layer}.hook_resid_pre_24576.pt"
36
+ path = hf_hub_download(repo_id=REPO_ID, filename=FILENAME)
37
+ model, sparse_autoencoder, activation_store = LMSparseAutoencoderSessionloader.load_session_from_pretrained(
38
+ path = path
39
+ )
40
+ sparse_autoencoder.eval()
41
+ ```
42
+
43
+ You can also load the feature sparsity from huggingface.
44
+
45
+ ```python
46
+ FILENAME = f"final_sparse_autoencoder_gpt2-small_blocks.{layer}.hook_resid_pre_24576_log_feature_sparsity.pt"
47
+ path = hf_hub_download(repo_id=REPO_ID, filename=FILENAME)
48
+ log_feature_sparsity = torch.load(path, map_location=sparse_autoencoder.cfg.device)
49
+
50
+ ```
51
+ ### Background
52
+
53
+ We highly recommend this [tutorial](https://www.lesswrong.com/posts/LnHowHgmrMbWtpkxx/intro-to-superposition-and-sparse-autoencoders-colab).
54
+
55
+
56
+ ## High Level
57
+
58
+ ### Motivation
59
+
60
+ - **Accelerate SAE Research**: Support fast experimentation to understand SAEs and improve SAE training so we can train SAEs on larger and more diverse models.
61
+ - **Make Research like Play**: Support research into language model internals via SAEs. Good tooling can make research tremendously exciting and enjoyable. Balancing modifiability and reliability with ease of understanding / access is the name of the game here.
62
+ - **Build an awesome community**: Mechanistic Interpretability already has an awesome community but as that community grows, it makes sense that there will be niches. I'd love to build a great community around Sparse Autoencoders.
63
+
64
+ ### Goals
65
+
66
+ #### **SAE Training**: SAE Training features will fit into a number of categories including:
67
+ - **Making it easy to train SAEs**: Training SAEs is hard for a number of reasons and so making it easy for people to train SAEs with relatively little expertise seems like the main way this codebase will create value.
68
+ - **Training SAEs on more models**: Supporting training of SAEs on more models, architectures, different activations within those models.
69
+ - **Being better at training SAEs**: Enabling methodological changes which may improve SAE performance as measured by reconstruction loss, Cross Entropy Loss when using reconstructed activation, L1 loss, L0 and interpretability of features as well as improving speed of training or reducing the compute resources required to train SAEs.
70
+ - **Being better at measuring SAE Performance**: How do we know when SAEs are doing what we want them to? Improving training metrics should allow better decisions about which methods to use and which hyperparameters choices we make.
71
+ - **Training SAE variants**: People are already training “Transcoders” which map from one activation to another (such as before / after an MLP layer). These can be easily supported with a few changes. Other variants will come in time and
72
+
73
+ #### **Analysis with SAEs**: Using SAEs to understand neural network internals is an exciting, but complicated task.
74
+ - **Feature-wise Interpretability**: This looks something like "for each feature, have as much knowledge about it as possible". Part of this will feature dashboard improvements, or supporting better integrations with Neuronpedia.
75
+ - **Mechanistic Interpretability**: This comprises the more traditional kinds of Mechanistic Interpretability which TransformerLens supports and should be supported by this codebase. Making it easy to patch, ablate or otherwise intervene on features so as to find circuits will likely speed up lots of researchers.
76
+
77
+ ### Other Stuff
78
+
79
+ I think there are lots of other types of analysis that could be done in the future with SAE features. I've already explored many different types of statistical tests which can reveal interesting properties of features. There are also things like saliency mapping and attribution techniques which it would be nice to support.
80
+ - Accessibility and Code Quality: The codebase won’t be used if it doesn’t work and it also won’t get used if it’s too hard to understand, modify or read.
81
+ Making the code accessible: This involves tasks like turning the code base into a python package.
82
+ - Knowing how the code is supposed to work: Is the code well-documented? This will require docstrings, tutorials and links to related work and publications. Getting aligned on what the code does is critical to sharing a resource like this.
83
+ - Knowing the code works as intended: All code should be tested. Unit tests and acceptance tests are both important.
84
+ - Knowing the code is actually performant: This will ensure code works as intended. However deep learning introduces lots of complexity which makes actually running benchmarks essential to having confidence in the code.
85
+
86
+
87
+ ## Code Overview
88
+
89
+ The codebase contains 2 folders worth caring about:
90
+
91
+ - training: The main body of the code is here. Everything required for training SAEs.
92
+ - analysis: This code is mainly house the feature visualizer code we use to generate dashboards. It was written by Callum McDougal but I've ported it here with permission and edited it to work with a few different activation types.
93
+
94
+ Some other folders:
95
+
96
+ - tutorials: These aren't well maintained but I'll aim to clean them up soon.
97
+ - tests: When first developing the codebase, I was writing more tests. I have no idea whether they are currently working!
98
+
99
+ I've been commiting my research code to the `Research` folder but am not expecting other people use or look at that.
100
+
101
+
102
+ ### Training your own Sparse Autoencoder
103
+
104
+ Sparse Autoencoders can be intimidating at first but it's fairly simple to train one once you know what each part of the config does. I've created a config class which you instantiate and pass to the runner which will complete your training run and log it's progress to wandb.
105
+
106
+ Let's go through the major components of the config:
107
+ - Data: SAE's autoencode model activations. We need to specify the model, the part of the models activations we want to autoencode and the dataset the model is operating on when generating those activations. We now automatically detect if that dataset is tokenized and most huggingface datasets should be fine. One slightly annoying detail is that you need to know the dimensionality of those activations when contructing your SAE but you can get that in the transformerlens [docs](https://neelnanda-io.github.io/TransformerLens/generated/model_properties_table.html). Any language model in the table from those docs should work.
108
+ - SAE Parameters: Your expansion factor will determine the size of your SAE and the decoder bias initialization method should always be geometric_median or mean. Mean is faster but theoretically sub-optimal. I use another package to get the geometric median and it can be quite slow.
109
+ - Training Parameters: These are most critical. The right L1 coefficient (coefficient in the activation sparsity inducing term in the loss) changes with your learning rate but a good bet would be to use LR 4e-4 and L1 8e-5 for GPT2 small. These will vary for other models and playing around with them / short runs can be helpful. Training batch size of 4096 is standard and I'm not really sure whether there's benefit to playing with it. In theory a larger context size (one accurate to whatever the model was trained with) seems good but it's computationally cheaper to use 128. Learning rate warm up is important to avoid dead neurons.
110
+ - Activation Store Parameters: The activation store shuffles activations from forward passes over samples from your data. The larger it is, the better shuffling you'll get. In theory more shuffling is good. The total training tokens is a very important parameter. The more the better, but you'll often see good results having trained on a few hundred million tokens. Store batch batch size is a function of your gpu and how many forward passes of your model you want to do simultaneously when collecting activations.
111
+ - Dead Neurons / Sparsity Metrics: The config around resampling was more important when we were using resampling to avoid dead neurons (see Anthropic's post on this), but using ghost gradients, the resampling protcol is much simpler. I'd always set ghost grad to True and feature sampling method to None. The feature sampling window effects the dashboard statistics tracking feature occurence and the dead feature window tracks how many forward passes a neuron must not activate before we apply ghost grads to it.
112
+ - WANDB: Fairly straightfoward. Don't set log frequency too high or your dashboard will be slow!
113
+ - Device: I can run this code on my macbook with "mps" but mostly do runs with cuda.
114
+ - Dtype: Float16 maybe could work but I had some funky results and have left it at float32 for the time being.
115
+ - Checkpoints: I'd collected checkpoints on runs you care about but turn them off when tuning since it can be slow.
116
+
117
+
118
+ ```python
119
+ import torch
120
+ import os
121
+ import sys
122
+
123
+ os.environ["TOKENIZERS_PARALLELISM"] = "false"
124
+ os.environ["WANDB__SERVICE_WAIT"] = "300"
125
+
126
+ from sae_lens.training.config import LanguageModelSAERunnerConfig
127
+ from sae_lens.training.lm_runner import language_model_sae_runner
128
+
129
+ cfg = LanguageModelSAERunnerConfig(
130
+
131
+ # Data Generating Function (Model + Training Distibuion)
132
+ model_name = "gpt2-small",
133
+ hook_point = "blocks.2.hook_resid_pre",
134
+ hook_point_layer = 2,
135
+ d_in = 768,
136
+ dataset_path = "Skylion007/openwebtext",
137
+ is_dataset_tokenized=False,
138
+
139
+ # SAE Parameters
140
+ expansion_factor = 64,
141
+ b_dec_init_method = "geometric_median",
142
+
143
+ # Training Parameters
144
+ lr = 0.0004,
145
+ l1_coefficient = 0.00008,
146
+ lr_scheduler_name="constantwithwarmup",
147
+ train_batch_size = 4096,
148
+ context_size = 128,
149
+ lr_warm_up_steps=5000,
150
+
151
+ # Activation Store Parameters
152
+ n_batches_in_buffer = 128,
153
+ total_training_tokens = 1_000_000 * 300,
154
+ store_batch_size = 32,
155
+
156
+ # Dead Neurons and Sparsity
157
+ use_ghost_grads=True,
158
+ feature_sampling_window = 1000,
159
+ dead_feature_window=5000,
160
+ dead_feature_threshold = 1e-6,
161
+
162
+ # WANDB
163
+ log_to_wandb = True,
164
+ wandb_project= "mats_sae_training_gpt2",
165
+ wandb_entity = None,
166
+ wandb_log_frequency=100,
167
+
168
+ # Misc
169
+ device = "cuda",
170
+ seed = 42,
171
+ n_checkpoints = 10,
172
+ checkpoint_path = "checkpoints",
173
+ dtype = torch.float32,
174
+ )
175
+
176
+ sparse_autoencoder = language_model_sae_runner(cfg)
177
+
178
+ ```
179
+
180
+
181
+ ## Loading a Pretrained Language Model
182
+
183
+ Once your SAE is trained, the final SAE weights will be saved to wandb and are loadable via the session loader. The session loader will return:
184
+ - The model your SAE was trained on (presumably you're interested in studying this. It's always a HookedTransformer)
185
+ - Your SAE.
186
+ - An activations loader: from which you can get randomly sampled activations or batches of tokens from the dataset you used to train the SAE. (more on this in the tutorial)
187
+
188
+ ```python
189
+ from sae_lens import LMSparseAutoencoderSessionloader
190
+
191
+ path ="path/to/sparse_autoencoder.pt"
192
+ model, sparse_autoencoder, activations_loader = LMSparseAutoencoderSessionloader.load_session_from_pretrained(
193
+ path
194
+ )
195
+
196
+ ```
197
+ ## Tutorials
198
+
199
+ I wrote a tutorial to show users how to do some basic exploration of their SAE.
200
+ - `evaluating_your_sae.ipynb`: A quick/dirty notebook showing how to check L0 and Prediction loss with your SAE, as well as showing how to generate interactive dashboards using Callum's reporduction of [Anthropics interface](https://transformer-circuits.pub/2023/monosemantic-features#setup-interface).
201
+ - `logits_lens_with_features.ipynb`: A notebook showing how to reproduce the analysis from this [LessWrong post](https://www.lesswrong.com/posts/qykrYY6rXXM7EEs8Q/understanding-sae-features-with-the-logit-lens).
202
+
203
+ ## Example Dashboard
204
+
205
+ WandB Dashboards provide lots of useful insights while training SAE's. Here's a screenshot from one training run.
206
+
207
+ ![screenshot](content/dashboard_screenshot.png)
208
+
209
+
210
+ ## Example Output
211
+
212
+ Here's one feature we found in the residual stream of Layer 10 of GPT-2 Small:
213
+
214
+ ![alt text](content/readme_screenshot_predict_pronoun_feature.png). Open `gpt2_resid_pre10_predict_pronoun_feature.html` in your browser to interact with the dashboard (WIP).
215
+
216
+ Note, probably this feature could split into more mono-semantic features in a larger SAE that had been trained for longer. (this was was only about 49152 features trained on 10M tokens from OpenWebText).
217
+
218
+
219
+ ## Citations and References:
220
+
221
+ Research:
222
+ - [Towards Monosemanticy](https://transformer-circuits.pub/2023/monosemantic-features)
223
+ - [Sparse Autoencoders Find Highly Interpretable Features in Language Model](https://arxiv.org/abs/2309.08600)
224
+
225
+
226
+
227
+ Reference Implementations:
228
+ - [Neel Nanda](https://github.com/neelnanda-io/1L-Sparse-Autoencoder)
229
+ - [AI-Safety-Foundation](https://github.com/ai-safety-foundation/sparse_autoencoder).
230
+ - [Arthur Conmy](https://github.com/ArthurConmy/sae).
231
+ - [Callum McDougall](https://github.com/callummcdougall/sae-exercises-mats/tree/main)
@@ -0,0 +1,68 @@
1
+ [tool.poetry]
2
+ name = "sae-lens"
3
+ version = "0.1.0"
4
+ description = "Training and Analyzing Sparse Autoencoders (SAEs)"
5
+ authors = ["Joseph Bloom"]
6
+ readme = "README.md"
7
+ packages = [{include = "sae_lens"}]
8
+
9
+ [tool.poetry.dependencies]
10
+ python = "^3.10"
11
+ transformer-lens = "^1.14.0"
12
+ transformers = "^4.38.1"
13
+ jupyter = "^1.0.0"
14
+ plotly = "^5.19.0"
15
+ plotly-express = "^0.4.1"
16
+ nbformat = "^5.9.2"
17
+ ipykernel = "^6.29.2"
18
+ matplotlib = "^3.8.3"
19
+ matplotlib-inline = "^0.1.6"
20
+ datasets = "^2.17.1"
21
+ babe = "^0.0.7"
22
+ nltk = "^3.8.1"
23
+ sae-vis = "0.2.6"
24
+ mkdocs = "^1.5.3"
25
+ mkdocs-material = "^9.5.15"
26
+ mkdocs-autorefs = "^1.0.1"
27
+ mkdocs-section-index = "^0.3.8"
28
+ mkdocstrings = "^0.24.1"
29
+ mkdocstrings-python = "^1.9.0"
30
+
31
+
32
+ [tool.poetry.group.dev.dependencies]
33
+ black = "^24.2.0"
34
+ pytest = "^8.0.2"
35
+ pytest-cov = "^4.1.0"
36
+ pre-commit = "^3.6.2"
37
+ flake8 = "^7.0.0"
38
+ isort = "^5.13.2"
39
+ pyright = "^1.1.351"
40
+
41
+ [tool.isort]
42
+ profile = "black"
43
+
44
+ [tool.pyright]
45
+ typeCheckingMode = "strict"
46
+ reportMissingTypeStubs = "none"
47
+ reportUnknownMemberType = "none"
48
+ reportUnknownArgumentType = "none"
49
+ reportUnknownVariableType = "none"
50
+ reportUntypedFunctionDecorator = "none"
51
+ reportUnnecessaryIsInstance = "none"
52
+ reportUnnecessaryComparison = "none"
53
+ reportConstantRedefinition = "none"
54
+ reportUnknownLambdaType = "none"
55
+ reportPrivateUsage = "none"
56
+
57
+ [build-system]
58
+ requires = ["poetry-core"]
59
+ build-backend = "poetry.core.masonry.api"
60
+
61
+
62
+ [tool.semantic_release]
63
+ version_variables = [
64
+ "sae_lens/__init__.py:__version__",
65
+ "pyproject.toml:version",
66
+ ]
67
+ branch = "main"
68
+ build_command = "pip install poetry && poetry build"
@@ -0,0 +1,24 @@
1
+ __version__ = "0.1.0"
2
+
3
+ from .training.activations_store import ActivationsStore
4
+ from .training.cache_activations_runner import cache_activations_runner
5
+ from .training.config import CacheActivationsRunnerConfig, LanguageModelSAERunnerConfig
6
+ from .training.evals import run_evals
7
+ from .training.lm_runner import language_model_sae_runner
8
+ from .training.sae_group import SAEGroup
9
+ from .training.session_loader import LMSparseAutoencoderSessionloader
10
+ from .training.sparse_autoencoder import SparseAutoencoder
11
+ from .training.train_sae_on_language_model import train_sae_group_on_language_model
12
+
13
+ __all__ = [
14
+ "LanguageModelSAERunnerConfig",
15
+ "CacheActivationsRunnerConfig",
16
+ "LMSparseAutoencoderSessionloader",
17
+ "SparseAutoencoder",
18
+ "SAEGroup",
19
+ "run_evals",
20
+ "language_model_sae_runner",
21
+ "cache_activations_runner",
22
+ "ActivationsStore",
23
+ "train_sae_group_on_language_model",
24
+ ]
File without changes