sae-lens 4.0.4__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-4.0.4/LICENSE +21 -0
- sae_lens-4.0.4/PKG-INFO +95 -0
- sae_lens-4.0.4/README.md +55 -0
- sae_lens-4.0.4/pyproject.toml +103 -0
- sae_lens-4.0.4/sae_lens/__init__.py +35 -0
- sae_lens-4.0.4/sae_lens/analysis/__init__.py +0 -0
- sae_lens-4.0.4/sae_lens/analysis/feature_statistics.py +98 -0
- sae_lens-4.0.4/sae_lens/analysis/hooked_sae_transformer.py +314 -0
- sae_lens-4.0.4/sae_lens/analysis/neuronpedia_integration.py +490 -0
- sae_lens-4.0.4/sae_lens/analysis/tsea.py +657 -0
- sae_lens-4.0.4/sae_lens/cache_activations_runner.py +188 -0
- sae_lens-4.0.4/sae_lens/config.py +646 -0
- sae_lens-4.0.4/sae_lens/evals.py +989 -0
- sae_lens-4.0.4/sae_lens/load_model.py +44 -0
- sae_lens-4.0.4/sae_lens/pretokenize_runner.py +192 -0
- sae_lens-4.0.4/sae_lens/pretrained_saes.yaml +12856 -0
- sae_lens-4.0.4/sae_lens/sae.py +818 -0
- sae_lens-4.0.4/sae_lens/sae_training_runner.py +237 -0
- sae_lens-4.0.4/sae_lens/tokenization_and_batching.py +102 -0
- sae_lens-4.0.4/sae_lens/toolkit/__init__.py +0 -0
- sae_lens-4.0.4/sae_lens/toolkit/pretrained_sae_loaders.py +535 -0
- sae_lens-4.0.4/sae_lens/toolkit/pretrained_saes.py +152 -0
- sae_lens-4.0.4/sae_lens/toolkit/pretrained_saes_directory.py +91 -0
- sae_lens-4.0.4/sae_lens/toy_model_runner.py +65 -0
- sae_lens-4.0.4/sae_lens/training/__init__.py +0 -0
- sae_lens-4.0.4/sae_lens/training/activations_store.py +702 -0
- sae_lens-4.0.4/sae_lens/training/geometric_median.py +107 -0
- sae_lens-4.0.4/sae_lens/training/optim.py +159 -0
- sae_lens-4.0.4/sae_lens/training/sae_trainer.py +432 -0
- sae_lens-4.0.4/sae_lens/training/toy_models.py +521 -0
- sae_lens-4.0.4/sae_lens/training/train_toy_sae.py +126 -0
- sae_lens-4.0.4/sae_lens/training/training_sae.py +537 -0
- sae_lens-4.0.4/sae_lens/training/upload_saes_to_huggingface.py +137 -0
sae_lens-4.0.4/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.
|
sae_lens-4.0.4/PKG-INFO
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
Metadata-Version: 2.1
|
|
2
|
+
Name: sae-lens
|
|
3
|
+
Version: 4.0.4
|
|
4
|
+
Summary: Training and Analyzing Sparse Autoencoders (SAEs)
|
|
5
|
+
Home-page: https://jbloomaus.github.io/SAELens
|
|
6
|
+
License: MIT
|
|
7
|
+
Keywords: deep-learning,sparse-autoencoders,mechanistic-interpretability,PyTorch
|
|
8
|
+
Author: Joseph Bloom
|
|
9
|
+
Requires-Python: >=3.10,<4.0
|
|
10
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
11
|
+
Classifier: Programming Language :: Python :: 3
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
16
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
17
|
+
Provides-Extra: mamba
|
|
18
|
+
Requires-Dist: automated-interpretability (>=0.0.5,<1.0.0)
|
|
19
|
+
Requires-Dist: babe (>=0.0.7,<0.0.8)
|
|
20
|
+
Requires-Dist: datasets (>=2.17.1,<3.0.0)
|
|
21
|
+
Requires-Dist: mamba-lens (>=0.0.4,<0.0.5) ; extra == "mamba"
|
|
22
|
+
Requires-Dist: matplotlib (>=3.8.3,<4.0.0)
|
|
23
|
+
Requires-Dist: matplotlib-inline (>=0.1.6,<0.2.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: pytest-profiling (>=1.7.0,<2.0.0)
|
|
28
|
+
Requires-Dist: python-dotenv (>=1.0.1,<2.0.0)
|
|
29
|
+
Requires-Dist: pyyaml (>=6.0.1,<7.0.0)
|
|
30
|
+
Requires-Dist: pyzmq (==26.0.0)
|
|
31
|
+
Requires-Dist: safetensors (>=0.4.2,<0.5.0)
|
|
32
|
+
Requires-Dist: transformer-lens (>=2.0.0,<3.0.0)
|
|
33
|
+
Requires-Dist: transformers (>=4.38.1,<5.0.0)
|
|
34
|
+
Requires-Dist: typer (>=0.12.3,<0.13.0)
|
|
35
|
+
Requires-Dist: typing-extensions (>=4.10.0,<5.0.0)
|
|
36
|
+
Requires-Dist: zstandard (>=0.22.0,<0.23.0)
|
|
37
|
+
Project-URL: Repository, https://github.com/jbloomAus/SAELens
|
|
38
|
+
Description-Content-Type: text/markdown
|
|
39
|
+
|
|
40
|
+
<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">
|
|
41
|
+
|
|
42
|
+
# SAE Lens
|
|
43
|
+
[](https://pypi.org/project/sae-lens/)
|
|
44
|
+
[](https://opensource.org/licenses/MIT)
|
|
45
|
+
[](https://github.com/jbloomAus/SAELens/actions/workflows/build.yml)
|
|
46
|
+
[](https://github.com/jbloomAus/SAELens/actions/workflows/deploy_docs.yml)
|
|
47
|
+
[](https://codecov.io/gh/jbloomAus/SAELens)
|
|
48
|
+
|
|
49
|
+
SAELens exists to help researchers:
|
|
50
|
+
- Train sparse autoencoders.
|
|
51
|
+
- Analyse sparse autoencoders / research mechanistic interpretability.
|
|
52
|
+
- Generate insights which make it easier to create safe and aligned AI systems.
|
|
53
|
+
|
|
54
|
+
Please refer to the [documentation](https://jbloomaus.github.io/SAELens/) for information on how to:
|
|
55
|
+
- Download and Analyse pre-trained sparse autoencoders.
|
|
56
|
+
- Train your own sparse autoencoders.
|
|
57
|
+
- Generate feature dashboards with the [SAE-Vis Library](https://github.com/callummcdougall/sae_vis/tree/main).
|
|
58
|
+
|
|
59
|
+
SAE Lens is the result of many contributors working collectively to improve humanity's understanding of neural networks, many of whom are motivated by a desire to [safeguard humanity from risks posed by artificial intelligence](https://80000hours.org/problem-profiles/artificial-intelligence/).
|
|
60
|
+
|
|
61
|
+
This library is maintained by [Joseph Bloom](https://www.jbloomaus.com/) and [David Chanin](https://github.com/chanind).
|
|
62
|
+
|
|
63
|
+
## Loading Pre-trained SAEs.
|
|
64
|
+
|
|
65
|
+
Pre-trained SAEs for various models can be imported via SAE Lens. See this [page](https://jbloomaus.github.io/SAELens/sae_table/) in the readme for a list of all SAEs.
|
|
66
|
+
## Tutorials
|
|
67
|
+
|
|
68
|
+
- [SAE Lens + Neuronpedia](tutorials/tutorial_2_0.ipynb)[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/tutorial_2_0.ipynb)
|
|
69
|
+
- [Loading and Analysing Pre-Trained Sparse Autoencoders](tutorials/basic_loading_and_analysing.ipynb)
|
|
70
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/basic_loading_and_analysing.ipynb)
|
|
71
|
+
- [Understanding SAE Features with the Logit Lens](tutorials/logits_lens_with_features.ipynb)
|
|
72
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/logits_lens_with_features.ipynb)
|
|
73
|
+
- [Training a Sparse Autoencoder](tutorials/training_a_sparse_autoencoder.ipynb)
|
|
74
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/training_a_sparse_autoencoder.ipynb)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
## Join the Slack!
|
|
78
|
+
|
|
79
|
+
Feel free to join the [Open Source Mechanistic Interpretability Slack](https://join.slack.com/t/opensourcemechanistic/shared_invite/zt-2k0id7mv8-CsIgPLmmHd03RPJmLUcapw) for support!
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
## Citation
|
|
83
|
+
|
|
84
|
+
Please cite the package as follows:
|
|
85
|
+
|
|
86
|
+
```
|
|
87
|
+
@misc{bloom2024saetrainingcodebase,
|
|
88
|
+
title = {SAELens},
|
|
89
|
+
author = {Joseph Bloom, Curt Tigges and David Chanin},
|
|
90
|
+
year = {2024},
|
|
91
|
+
howpublished = {\url{https://github.com/jbloomAus/SAELens}},
|
|
92
|
+
}
|
|
93
|
+
```
|
|
94
|
+
|
|
95
|
+
|
sae_lens-4.0.4/README.md
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
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
|
+
[](https://pypi.org/project/sae-lens/)
|
|
5
|
+
[](https://opensource.org/licenses/MIT)
|
|
6
|
+
[](https://github.com/jbloomAus/SAELens/actions/workflows/build.yml)
|
|
7
|
+
[](https://github.com/jbloomAus/SAELens/actions/workflows/deploy_docs.yml)
|
|
8
|
+
[](https://codecov.io/gh/jbloomAus/SAELens)
|
|
9
|
+
|
|
10
|
+
SAELens exists to help researchers:
|
|
11
|
+
- Train sparse autoencoders.
|
|
12
|
+
- Analyse sparse autoencoders / research mechanistic interpretability.
|
|
13
|
+
- Generate insights which make it easier to create safe and aligned AI systems.
|
|
14
|
+
|
|
15
|
+
Please refer to the [documentation](https://jbloomaus.github.io/SAELens/) for information on how to:
|
|
16
|
+
- Download and Analyse pre-trained sparse autoencoders.
|
|
17
|
+
- Train your own sparse autoencoders.
|
|
18
|
+
- Generate feature dashboards with the [SAE-Vis Library](https://github.com/callummcdougall/sae_vis/tree/main).
|
|
19
|
+
|
|
20
|
+
SAE Lens is the result of many contributors working collectively to improve humanity's understanding of neural networks, many of whom are motivated by a desire to [safeguard humanity from risks posed by artificial intelligence](https://80000hours.org/problem-profiles/artificial-intelligence/).
|
|
21
|
+
|
|
22
|
+
This library is maintained by [Joseph Bloom](https://www.jbloomaus.com/) and [David Chanin](https://github.com/chanind).
|
|
23
|
+
|
|
24
|
+
## Loading Pre-trained SAEs.
|
|
25
|
+
|
|
26
|
+
Pre-trained SAEs for various models can be imported via SAE Lens. See this [page](https://jbloomaus.github.io/SAELens/sae_table/) in the readme for a list of all SAEs.
|
|
27
|
+
## Tutorials
|
|
28
|
+
|
|
29
|
+
- [SAE Lens + Neuronpedia](tutorials/tutorial_2_0.ipynb)[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/tutorial_2_0.ipynb)
|
|
30
|
+
- [Loading and Analysing Pre-Trained Sparse Autoencoders](tutorials/basic_loading_and_analysing.ipynb)
|
|
31
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/basic_loading_and_analysing.ipynb)
|
|
32
|
+
- [Understanding SAE Features with the Logit Lens](tutorials/logits_lens_with_features.ipynb)
|
|
33
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/logits_lens_with_features.ipynb)
|
|
34
|
+
- [Training a Sparse Autoencoder](tutorials/training_a_sparse_autoencoder.ipynb)
|
|
35
|
+
[](https://githubtocolab.com/jbloomAus/SAELens/blob/main/tutorials/training_a_sparse_autoencoder.ipynb)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
## Join the Slack!
|
|
39
|
+
|
|
40
|
+
Feel free to join the [Open Source Mechanistic Interpretability Slack](https://join.slack.com/t/opensourcemechanistic/shared_invite/zt-2k0id7mv8-CsIgPLmmHd03RPJmLUcapw) for support!
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
## Citation
|
|
44
|
+
|
|
45
|
+
Please cite the package as follows:
|
|
46
|
+
|
|
47
|
+
```
|
|
48
|
+
@misc{bloom2024saetrainingcodebase,
|
|
49
|
+
title = {SAELens},
|
|
50
|
+
author = {Joseph Bloom, Curt Tigges and David Chanin},
|
|
51
|
+
year = {2024},
|
|
52
|
+
howpublished = {\url{https://github.com/jbloomAus/SAELens}},
|
|
53
|
+
}
|
|
54
|
+
```
|
|
55
|
+
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
[tool.poetry]
|
|
2
|
+
name = "sae-lens"
|
|
3
|
+
version = "4.0.4"
|
|
4
|
+
description = "Training and Analyzing Sparse Autoencoders (SAEs)"
|
|
5
|
+
authors = ["Joseph Bloom"]
|
|
6
|
+
readme = "README.md"
|
|
7
|
+
packages = [{include = "sae_lens"}]
|
|
8
|
+
include = ["pretrained_saes.yaml"]
|
|
9
|
+
repository = "https://github.com/jbloomAus/SAELens"
|
|
10
|
+
homepage = "https://jbloomaus.github.io/SAELens"
|
|
11
|
+
license = "MIT"
|
|
12
|
+
keywords = [
|
|
13
|
+
"deep-learning",
|
|
14
|
+
"sparse-autoencoders",
|
|
15
|
+
"mechanistic-interpretability",
|
|
16
|
+
"PyTorch",
|
|
17
|
+
]
|
|
18
|
+
classifiers = ["Topic :: Scientific/Engineering :: Artificial Intelligence"]
|
|
19
|
+
|
|
20
|
+
[tool.poetry.dependencies]
|
|
21
|
+
python = "^3.10"
|
|
22
|
+
transformer-lens = "^2.0.0"
|
|
23
|
+
transformers = "^4.38.1"
|
|
24
|
+
plotly = "^5.19.0"
|
|
25
|
+
plotly-express = "^0.4.1"
|
|
26
|
+
matplotlib = "^3.8.3"
|
|
27
|
+
matplotlib-inline = "^0.1.6"
|
|
28
|
+
datasets = "^2.17.1"
|
|
29
|
+
babe = "^0.0.7"
|
|
30
|
+
nltk = "^3.8.1"
|
|
31
|
+
safetensors = "^0.4.2"
|
|
32
|
+
typer = "^0.12.3"
|
|
33
|
+
mamba-lens = { version = "^0.0.4", optional = true }
|
|
34
|
+
pyzmq = "26.0.0"
|
|
35
|
+
automated-interpretability = ">=0.0.5,<1.0.0"
|
|
36
|
+
python-dotenv = "^1.0.1"
|
|
37
|
+
pyyaml = "^6.0.1"
|
|
38
|
+
pytest-profiling = "^1.7.0"
|
|
39
|
+
zstandard = "^0.22.0"
|
|
40
|
+
typing-extensions = "^4.10.0"
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
[tool.poetry.group.dev.dependencies]
|
|
44
|
+
black = { version = "24.4.0", extras = ["jupyter"] }
|
|
45
|
+
pytest = "^8.0.2"
|
|
46
|
+
pytest-cov = "^4.1.0"
|
|
47
|
+
pre-commit = "^3.6.2"
|
|
48
|
+
flake8 = "7.0.0"
|
|
49
|
+
isort = "5.13.2"
|
|
50
|
+
pyright = "1.1.365"
|
|
51
|
+
mamba-lens = "^0.0.4"
|
|
52
|
+
ansible-lint = { version = "^24.2.3", markers = "platform_system != 'Windows'" }
|
|
53
|
+
botocore = "^1.34.101"
|
|
54
|
+
boto3 = "^1.34.101"
|
|
55
|
+
docstr-coverage = "^2.3.2"
|
|
56
|
+
mkdocs = "^1.6.1"
|
|
57
|
+
mkdocs-material = "^9.5.34"
|
|
58
|
+
mkdocs-autorefs = "^1.1.0"
|
|
59
|
+
mkdocs-section-index = "^0.3.9"
|
|
60
|
+
mkdocstrings = "^0.25.2"
|
|
61
|
+
mkdocstrings-python = "^1.10.9"
|
|
62
|
+
tabulate = "^0.9.0"
|
|
63
|
+
|
|
64
|
+
[tool.poetry.extras]
|
|
65
|
+
mamba = ["mamba-lens"]
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
[tool.isort]
|
|
69
|
+
profile = "black"
|
|
70
|
+
src_paths = ["sae_lens", "tests"]
|
|
71
|
+
|
|
72
|
+
[tool.pyright]
|
|
73
|
+
typeCheckingMode = "strict"
|
|
74
|
+
reportMissingTypeStubs = "none"
|
|
75
|
+
reportUnknownMemberType = "none"
|
|
76
|
+
reportUnknownArgumentType = "none"
|
|
77
|
+
reportUnknownVariableType = "none"
|
|
78
|
+
reportUntypedFunctionDecorator = "none"
|
|
79
|
+
reportUnnecessaryIsInstance = "none"
|
|
80
|
+
reportUnnecessaryComparison = "none"
|
|
81
|
+
reportConstantRedefinition = "none"
|
|
82
|
+
reportUnknownLambdaType = "none"
|
|
83
|
+
reportPrivateUsage = "none"
|
|
84
|
+
reportDeprecated = "none"
|
|
85
|
+
reportPrivateImportUsage = "none"
|
|
86
|
+
ignore = [
|
|
87
|
+
"**/wandb/**"
|
|
88
|
+
]
|
|
89
|
+
|
|
90
|
+
[build-system]
|
|
91
|
+
requires = ["poetry-core"]
|
|
92
|
+
build-backend = "poetry.core.masonry.api"
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
[tool.semantic_release]
|
|
96
|
+
version_variables = [
|
|
97
|
+
"sae_lens/__init__.py:__version__",
|
|
98
|
+
]
|
|
99
|
+
version_toml = [
|
|
100
|
+
"pyproject.toml:tool.poetry.version",
|
|
101
|
+
]
|
|
102
|
+
branch = "main"
|
|
103
|
+
build_command = "pip install poetry && poetry build"
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
__version__ = "4.0.4"
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
from .analysis.hooked_sae_transformer import HookedSAETransformer
|
|
5
|
+
from .cache_activations_runner import CacheActivationsRunner
|
|
6
|
+
from .config import (
|
|
7
|
+
CacheActivationsRunnerConfig,
|
|
8
|
+
LanguageModelSAERunnerConfig,
|
|
9
|
+
PretokenizeRunnerConfig,
|
|
10
|
+
)
|
|
11
|
+
from .evals import run_evals
|
|
12
|
+
from .pretokenize_runner import PretokenizeRunner, pretokenize_runner
|
|
13
|
+
from .sae import SAE, SAEConfig
|
|
14
|
+
from .sae_training_runner import SAETrainingRunner
|
|
15
|
+
from .training.activations_store import ActivationsStore
|
|
16
|
+
from .training.training_sae import TrainingSAE, TrainingSAEConfig
|
|
17
|
+
from .training.upload_saes_to_huggingface import upload_saes_to_huggingface
|
|
18
|
+
|
|
19
|
+
__all__ = [
|
|
20
|
+
"SAE",
|
|
21
|
+
"SAEConfig",
|
|
22
|
+
"TrainingSAE",
|
|
23
|
+
"TrainingSAEConfig",
|
|
24
|
+
"HookedSAETransformer",
|
|
25
|
+
"ActivationsStore",
|
|
26
|
+
"LanguageModelSAERunnerConfig",
|
|
27
|
+
"SAETrainingRunner",
|
|
28
|
+
"CacheActivationsRunnerConfig",
|
|
29
|
+
"CacheActivationsRunner",
|
|
30
|
+
"PretokenizeRunnerConfig",
|
|
31
|
+
"PretokenizeRunner",
|
|
32
|
+
"pretokenize_runner",
|
|
33
|
+
"run_evals",
|
|
34
|
+
"upload_saes_to_huggingface",
|
|
35
|
+
]
|
|
File without changes
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
import pandas as pd
|
|
2
|
+
import torch
|
|
3
|
+
from tqdm import tqdm
|
|
4
|
+
from transformer_lens import HookedTransformer
|
|
5
|
+
|
|
6
|
+
from sae_lens.sae import SAE
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@torch.no_grad()
|
|
10
|
+
def get_feature_property_df(sae: SAE, feature_sparsity: torch.Tensor):
|
|
11
|
+
"""
|
|
12
|
+
feature_property_df = get_feature_property_df(sae, log_feature_density.cpu())
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
W_dec_normalized = (
|
|
16
|
+
sae.W_dec.cpu()
|
|
17
|
+
) # / sparse_autoencoder.W_dec.cpu().norm(dim=-1, keepdim=True)
|
|
18
|
+
W_enc_normalized = (sae.W_enc.cpu() / sae.W_enc.cpu().norm(dim=-1, keepdim=True)).T
|
|
19
|
+
|
|
20
|
+
d_e_projection = (W_dec_normalized * W_enc_normalized).sum(-1)
|
|
21
|
+
b_dec_projection = sae.b_dec.cpu() @ W_dec_normalized.T
|
|
22
|
+
|
|
23
|
+
temp_df = pd.DataFrame(
|
|
24
|
+
{
|
|
25
|
+
"log_feature_sparsity": feature_sparsity + 1e-10,
|
|
26
|
+
"d_e_projection": d_e_projection,
|
|
27
|
+
# "d_e_projection_normalized": d_e_projection_normalized,
|
|
28
|
+
"b_enc": sae.b_enc.detach().cpu(),
|
|
29
|
+
"b_dec_projection": b_dec_projection,
|
|
30
|
+
"feature": list(range(sae.cfg.d_sae)), # type: ignore
|
|
31
|
+
"dead_neuron": (feature_sparsity < -9).cpu(),
|
|
32
|
+
}
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
return temp_df
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@torch.no_grad()
|
|
39
|
+
def get_stats_df(projection: torch.Tensor):
|
|
40
|
+
"""
|
|
41
|
+
Returns a dataframe with the mean, std, skewness and kurtosis of the projection
|
|
42
|
+
"""
|
|
43
|
+
mean = projection.mean(dim=1, keepdim=True)
|
|
44
|
+
diffs = projection - mean
|
|
45
|
+
var = (diffs**2).mean(dim=1, keepdim=True)
|
|
46
|
+
std = torch.pow(var, 0.5)
|
|
47
|
+
zscores = diffs / std
|
|
48
|
+
skews = torch.mean(torch.pow(zscores, 3.0), dim=1)
|
|
49
|
+
kurtosis = torch.mean(torch.pow(zscores, 4.0), dim=1)
|
|
50
|
+
|
|
51
|
+
stats_df = pd.DataFrame(
|
|
52
|
+
{
|
|
53
|
+
"feature": range(len(skews)),
|
|
54
|
+
"mean": mean.numpy().squeeze(),
|
|
55
|
+
"std": std.numpy().squeeze(),
|
|
56
|
+
"skewness": skews.numpy(),
|
|
57
|
+
"kurtosis": kurtosis.numpy(),
|
|
58
|
+
}
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
return stats_df
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@torch.no_grad()
|
|
65
|
+
def get_all_stats_dfs(
|
|
66
|
+
gpt2_small_sparse_autoencoders: dict[str, SAE], # [hook_point, sae]
|
|
67
|
+
gpt2_small_sae_sparsities: dict[str, torch.Tensor], # [hook_point, sae]
|
|
68
|
+
model: HookedTransformer,
|
|
69
|
+
cosine_sim: bool = False,
|
|
70
|
+
):
|
|
71
|
+
stats_dfs = []
|
|
72
|
+
pbar = tqdm(gpt2_small_sparse_autoencoders.keys())
|
|
73
|
+
for key in pbar:
|
|
74
|
+
layer = int(key.split(".")[1])
|
|
75
|
+
sparse_autoencoder = gpt2_small_sparse_autoencoders[key]
|
|
76
|
+
pbar.set_description(f"Processing layer {sparse_autoencoder.cfg.hook_name}")
|
|
77
|
+
W_U_stats_df_dec, _ = get_W_U_W_dec_stats_df(
|
|
78
|
+
sparse_autoencoder.W_dec.cpu(), model, cosine_sim
|
|
79
|
+
)
|
|
80
|
+
log_feature_sparsity = gpt2_small_sae_sparsities[key].detach().cpu()
|
|
81
|
+
W_U_stats_df_dec["log_feature_sparsity"] = log_feature_sparsity
|
|
82
|
+
W_U_stats_df_dec["layer"] = layer + (1 if "post" in key else 0)
|
|
83
|
+
stats_dfs.append(W_U_stats_df_dec)
|
|
84
|
+
|
|
85
|
+
W_U_stats_df_dec_all_layers = pd.concat(stats_dfs, axis=0)
|
|
86
|
+
return W_U_stats_df_dec_all_layers
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
@torch.no_grad()
|
|
90
|
+
def get_W_U_W_dec_stats_df(
|
|
91
|
+
W_dec: torch.Tensor, model: HookedTransformer, cosine_sim: bool = False
|
|
92
|
+
) -> tuple[pd.DataFrame, torch.Tensor]:
|
|
93
|
+
W_U = model.W_U.detach().cpu()
|
|
94
|
+
if cosine_sim:
|
|
95
|
+
W_U = W_U / W_U.norm(dim=0, keepdim=True)
|
|
96
|
+
dec_projection_onto_W_U = W_dec @ W_U
|
|
97
|
+
W_U_stats_df = get_stats_df(dec_projection_onto_W_U)
|
|
98
|
+
return W_U_stats_df, dec_projection_onto_W_U
|