logogram 0.1.0__py3-none-any.whl
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.
- logogram/__init__.py +6 -0
- logogram/__main__.py +5 -0
- logogram/analysis.py +419 -0
- logogram/atp.py +120 -0
- logogram/backends/__init__.py +5 -0
- logogram/backends/base.py +202 -0
- logogram/backends/hub.py +375 -0
- logogram/backends/saes.py +277 -0
- logogram/backends/transformer_lens.py +872 -0
- logogram/cli.py +496 -0
- logogram/compare.py +177 -0
- logogram/datasets.py +159 -0
- logogram/direct.py +193 -0
- logogram/engine.py +550 -0
- logogram/examples/ioi-gpt2/.gitignore +3 -0
- logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
- logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
- logogram/examples/ioi-gpt2/project.json +6 -0
- logogram/exports.py +33 -0
- logogram/features.py +368 -0
- logogram/fileio.py +63 -0
- logogram/ioi.py +220 -0
- logogram/paths.py +204 -0
- logogram/project.py +444 -0
- logogram/prompts.py +204 -0
- logogram/research.py +84 -0
- logogram/results.py +240 -0
- logogram/runner.py +396 -0
- logogram/runs.py +98 -0
- logogram/sae.py +161 -0
- logogram/schema.py +302 -0
- logogram/server/__init__.py +1 -0
- logogram/server/app.py +1083 -0
- logogram/server/models.py +426 -0
- logogram/server/security.py +212 -0
- logogram/server/state.py +585 -0
- logogram/sites.py +249 -0
- logogram/spec.py +518 -0
- logogram/stats.py +171 -0
- logogram/steering.py +258 -0
- logogram/system.py +379 -0
- logogram/updates.py +194 -0
- logogram/verify.py +39 -0
- logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
- logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
- logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
- logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
- logogram/web_dist/favicon.svg +1 -0
- logogram/web_dist/index.html +15 -0
- logogram-0.1.0.dist-info/METADATA +550 -0
- logogram-0.1.0.dist-info/RECORD +54 -0
- logogram-0.1.0.dist-info/WHEEL +4 -0
- logogram-0.1.0.dist-info/entry_points.txt +2 -0
- logogram-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,277 @@
|
|
|
1
|
+
"""Published sparse autoencoders: find, download and read their files.
|
|
2
|
+
|
|
3
|
+
Two safetensors formats are read:
|
|
4
|
+
|
|
5
|
+
* SAELens (``cfg.json`` + ``sae_weights.safetensors``), as in the GPT-2 small residual SAEs. The
|
|
6
|
+
config names the TransformerLens hook the SAE reads, such as ``blocks.8.hook_resid_pre``.
|
|
7
|
+
* EleutherAI's sparsify (``cfg.json`` + ``sae.safetensors``), as for Pythia, SmolLM2 and Llama.
|
|
8
|
+
The folder names the module whose output the SAE reads, such as ``layers.3`` (the residual
|
|
9
|
+
stream after layer 3) or ``layers.3.mlp`` (the MLP output).
|
|
10
|
+
|
|
11
|
+
Hook and module names are library details, so they are mapped to abstract sites here. Parameters
|
|
12
|
+
are read only from safetensors; Gemma Scope's NumPy archives are refused rather than loaded with a
|
|
13
|
+
different reader. Anything a format does that Logogram can't reproduce exactly (a learned scaling,
|
|
14
|
+
a gated encoder, a transcoder) is refused with the reason.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
import json
|
|
20
|
+
import re
|
|
21
|
+
from dataclasses import dataclass
|
|
22
|
+
from pathlib import Path
|
|
23
|
+
from typing import Any
|
|
24
|
+
|
|
25
|
+
import torch
|
|
26
|
+
|
|
27
|
+
from logogram.backends.base import BackendError
|
|
28
|
+
|
|
29
|
+
WEIGHT_FILES = ("sae_weights.safetensors", "sae.safetensors")
|
|
30
|
+
|
|
31
|
+
_TL_HOOK = re.compile(r"blocks\.(\d+)\.hook_(resid_pre|resid_post|attn_out|mlp_out)")
|
|
32
|
+
_MODULE = re.compile(r"(?:.*\.)?layers\.(\d+)(?:\.(mlp|attention|self_attn|attn))?")
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@dataclass
|
|
36
|
+
class SAEParams:
|
|
37
|
+
"""An SAE as read from its files: what it reads, and how it encodes and decodes."""
|
|
38
|
+
|
|
39
|
+
site: str # abstract site: resid_pre, resid_post, attn_out or mlp_out
|
|
40
|
+
layer: int
|
|
41
|
+
d_in: int
|
|
42
|
+
d_sae: int
|
|
43
|
+
W_enc: torch.Tensor # [d_in, d_sae]
|
|
44
|
+
b_enc: torch.Tensor # [d_sae]
|
|
45
|
+
W_dec: torch.Tensor # [d_sae, d_in]
|
|
46
|
+
b_dec: torch.Tensor # [d_in]
|
|
47
|
+
activation: str # relu, topk or jumprelu
|
|
48
|
+
k: int | None
|
|
49
|
+
threshold: torch.Tensor | None # [d_sae], JumpReLU only
|
|
50
|
+
subtract_b_dec: bool # subtract b_dec from the input before encoding
|
|
51
|
+
normalize: str # "none", or "layer_norm": each input standardized, each output restored
|
|
52
|
+
format: str # saelens or eleuther
|
|
53
|
+
note: str # what the SAE says about itself, for the app (model name, training hook)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def site_of_hook(name: str) -> tuple[str, int]:
|
|
57
|
+
match = _TL_HOOK.fullmatch(name)
|
|
58
|
+
if match is None:
|
|
59
|
+
raise BackendError(
|
|
60
|
+
f"This SAE reads {name}, which Logogram can't map to a residual stream, attention "
|
|
61
|
+
"output or MLP output. SAEs on single heads or on the stream between attention and "
|
|
62
|
+
"MLP aren't supported yet."
|
|
63
|
+
)
|
|
64
|
+
return match.group(2), int(match.group(1))
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def site_of_module(folder: str) -> tuple[str, int]:
|
|
68
|
+
match = _MODULE.fullmatch(folder.strip("/").split("/")[-1])
|
|
69
|
+
if match is None:
|
|
70
|
+
raise BackendError(
|
|
71
|
+
f"This SAE reads the module {folder!r}, which Logogram can't map to a layer's "
|
|
72
|
+
"output, attention output or MLP output."
|
|
73
|
+
)
|
|
74
|
+
part = match.group(2)
|
|
75
|
+
kind = "resid_post" if part is None else "mlp_out" if part == "mlp" else "attn_out"
|
|
76
|
+
return kind, int(match.group(1))
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def read_sae(folder: Path, path: str = "") -> SAEParams:
|
|
80
|
+
"""Read an SAE from a local folder holding its config and safetensors weights."""
|
|
81
|
+
from safetensors.torch import load_file
|
|
82
|
+
|
|
83
|
+
if (folder / "params.npz").exists():
|
|
84
|
+
raise BackendError(
|
|
85
|
+
"This SAE is published as a NumPy archive (params.npz), as Gemma Scope is. Logogram "
|
|
86
|
+
"reads SAE parameters only from safetensors files."
|
|
87
|
+
)
|
|
88
|
+
cfg_path = folder / "cfg.json"
|
|
89
|
+
weights = next((folder / name for name in WEIGHT_FILES if (folder / name).is_file()), None)
|
|
90
|
+
if not cfg_path.is_file() or weights is None:
|
|
91
|
+
raise BackendError(
|
|
92
|
+
"This folder has no SAE in a format Logogram reads: it needs cfg.json with "
|
|
93
|
+
"sae_weights.safetensors (SAELens) or sae.safetensors (EleutherAI)."
|
|
94
|
+
)
|
|
95
|
+
try:
|
|
96
|
+
cfg: dict[str, Any] = json.loads(cfg_path.read_text(encoding="utf-8"))
|
|
97
|
+
except (OSError, ValueError) as exc:
|
|
98
|
+
raise BackendError(f"The SAE's cfg.json can't be read: {exc}") from exc
|
|
99
|
+
tensors = load_file(str(weights), device="cpu")
|
|
100
|
+
if weights.name == "sae_weights.safetensors":
|
|
101
|
+
return _saelens(cfg, tensors)
|
|
102
|
+
return _eleuther(cfg, tensors, path or folder.name)
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _need(tensors: dict[str, torch.Tensor], *names: str) -> list[torch.Tensor]:
|
|
106
|
+
missing = [n for n in names if n not in tensors]
|
|
107
|
+
if missing:
|
|
108
|
+
raise BackendError(f"The SAE's weights are missing {', '.join(missing)}.")
|
|
109
|
+
return [tensors[n].float() for n in names]
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def _saelens(cfg: dict[str, Any], tensors: dict[str, torch.Tensor]) -> SAEParams:
|
|
113
|
+
hook = cfg.get("hook_name") or cfg.get("hook_point")
|
|
114
|
+
if not isinstance(hook, str):
|
|
115
|
+
raise BackendError("The SAE's cfg.json doesn't say which hook it reads.")
|
|
116
|
+
if cfg.get("hook_head_index", cfg.get("hook_point_head_index")) is not None:
|
|
117
|
+
raise BackendError("SAEs on a single attention head aren't supported yet.")
|
|
118
|
+
site, layer = site_of_hook(hook)
|
|
119
|
+
architecture = cfg.get("architecture", "standard")
|
|
120
|
+
if architecture not in ("standard", "jumprelu", "topk"):
|
|
121
|
+
raise BackendError(f"SAELens {architecture} SAEs aren't supported yet.")
|
|
122
|
+
normalize = cfg.get("normalize_activations") or "none"
|
|
123
|
+
if normalize not in ("none", "layer_norm"):
|
|
124
|
+
raise BackendError(
|
|
125
|
+
f"This SAE rescales its inputs ({normalize}) with a factor measured on its training "
|
|
126
|
+
"data, which Logogram can't reproduce exactly."
|
|
127
|
+
)
|
|
128
|
+
W_enc, b_enc, W_dec, b_dec = _need(tensors, "W_enc", "b_enc", "W_dec", "b_dec")
|
|
129
|
+
activation, k, threshold = "relu", None, None
|
|
130
|
+
fn = cfg.get("activation_fn_str") or cfg.get("activation_fn") or "relu"
|
|
131
|
+
if architecture == "jumprelu":
|
|
132
|
+
activation = "jumprelu"
|
|
133
|
+
(threshold,) = _need(tensors, "threshold")
|
|
134
|
+
elif fn == "topk" or architecture == "topk":
|
|
135
|
+
activation = "topk"
|
|
136
|
+
kwargs = cfg.get("activation_fn_kwargs") or {}
|
|
137
|
+
k = int(kwargs.get("k") or cfg.get("k") or 0)
|
|
138
|
+
if k <= 0:
|
|
139
|
+
raise BackendError("This TopK SAE doesn't say how many features it keeps (k).")
|
|
140
|
+
elif fn != "relu":
|
|
141
|
+
raise BackendError(f"SAEs with a {fn} activation aren't supported yet.")
|
|
142
|
+
note = " · ".join(str(v) for v in (cfg.get("model_name"), hook) if v)
|
|
143
|
+
return SAEParams(
|
|
144
|
+
site=site,
|
|
145
|
+
layer=layer,
|
|
146
|
+
d_in=int(W_enc.shape[0]),
|
|
147
|
+
d_sae=int(W_enc.shape[1]),
|
|
148
|
+
W_enc=W_enc,
|
|
149
|
+
b_enc=b_enc,
|
|
150
|
+
W_dec=W_dec,
|
|
151
|
+
b_dec=b_dec,
|
|
152
|
+
activation=activation,
|
|
153
|
+
k=k,
|
|
154
|
+
threshold=threshold,
|
|
155
|
+
subtract_b_dec=bool(cfg.get("apply_b_dec_to_input", True)),
|
|
156
|
+
normalize=normalize,
|
|
157
|
+
format="saelens",
|
|
158
|
+
note=note,
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def _eleuther(cfg: dict[str, Any], tensors: dict[str, torch.Tensor], folder: str) -> SAEParams:
|
|
163
|
+
for flag, what in (
|
|
164
|
+
("transcode", "Transcoders"),
|
|
165
|
+
("skip_connection", "SAEs with a skip connection"),
|
|
166
|
+
("signed", "Signed SAEs"),
|
|
167
|
+
):
|
|
168
|
+
if cfg.get(flag):
|
|
169
|
+
raise BackendError(f"{what} aren't supported yet.")
|
|
170
|
+
site, layer = site_of_module(folder)
|
|
171
|
+
encoder, b_enc, W_dec, b_dec = _need(
|
|
172
|
+
tensors, "encoder.weight", "encoder.bias", "W_dec", "b_dec"
|
|
173
|
+
)
|
|
174
|
+
k = int(cfg.get("k") or 0)
|
|
175
|
+
if k <= 0:
|
|
176
|
+
raise BackendError("This TopK SAE doesn't say how many features it keeps (k).")
|
|
177
|
+
return SAEParams(
|
|
178
|
+
site=site,
|
|
179
|
+
layer=layer,
|
|
180
|
+
d_in=int(encoder.shape[1]),
|
|
181
|
+
d_sae=int(encoder.shape[0]),
|
|
182
|
+
W_enc=encoder.T.contiguous(),
|
|
183
|
+
b_enc=b_enc,
|
|
184
|
+
W_dec=W_dec,
|
|
185
|
+
b_dec=b_dec,
|
|
186
|
+
activation="topk",
|
|
187
|
+
k=k,
|
|
188
|
+
threshold=None,
|
|
189
|
+
subtract_b_dec=True,
|
|
190
|
+
normalize="none",
|
|
191
|
+
format="eleuther",
|
|
192
|
+
note=f"EleutherAI sparsify · {folder}",
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
# -- the Hub -----------------------------------------------------------------------------------
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
def list_saes(repo: str, revision: str | None = None) -> tuple[str, list[str]]:
|
|
200
|
+
"""The exact revision of an SAE repository and the folders in it that hold an SAE."""
|
|
201
|
+
from huggingface_hub import HfApi
|
|
202
|
+
|
|
203
|
+
from logogram.backends.hub import _friendly_hub_error
|
|
204
|
+
|
|
205
|
+
try:
|
|
206
|
+
info = HfApi().model_info(repo, revision=revision)
|
|
207
|
+
except Exception as exc: # noqa: BLE001
|
|
208
|
+
raise _friendly_hub_error(repo, exc) from exc
|
|
209
|
+
names = [s.rfilename for s in info.siblings or []]
|
|
210
|
+
folders = sorted(
|
|
211
|
+
{
|
|
212
|
+
n.rsplit("/", 1)[0] if "/" in n else ""
|
|
213
|
+
for n in names
|
|
214
|
+
if n.rsplit("/", 1)[-1] in WEIGHT_FILES
|
|
215
|
+
},
|
|
216
|
+
key=_natural,
|
|
217
|
+
)
|
|
218
|
+
return str(info.sha), folders
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
def _natural(text: str) -> list[Any]:
|
|
222
|
+
return [int(part) if part.isdigit() else part for part in re.split(r"(\d+)", text)]
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def download_sae(
|
|
226
|
+
repo: str,
|
|
227
|
+
path: str,
|
|
228
|
+
revision: str | None,
|
|
229
|
+
progress: Any = None,
|
|
230
|
+
cancel: Any = None,
|
|
231
|
+
) -> tuple[str, Path]:
|
|
232
|
+
"""Download one SAE (its config and safetensors weights) at an exact revision."""
|
|
233
|
+
from huggingface_hub import HfApi
|
|
234
|
+
|
|
235
|
+
from logogram.backends import hub
|
|
236
|
+
|
|
237
|
+
prefix = f"{path.strip('/')}/" if path.strip("/") else ""
|
|
238
|
+
try:
|
|
239
|
+
info = HfApi().model_info(repo, revision=revision, files_metadata=True)
|
|
240
|
+
except Exception as exc: # noqa: BLE001
|
|
241
|
+
cached = _cached(repo, prefix, revision)
|
|
242
|
+
if cached is not None and hub._is_offline_error(exc):
|
|
243
|
+
return cached
|
|
244
|
+
raise hub._friendly_hub_error(repo, exc) from exc
|
|
245
|
+
sizes = {s.rfilename: int(s.size or 0) for s in info.siblings or []}
|
|
246
|
+
if f"{prefix}params.npz" in sizes:
|
|
247
|
+
raise BackendError(
|
|
248
|
+
"This SAE is published as a NumPy archive (params.npz), as Gemma Scope is. Logogram "
|
|
249
|
+
"reads SAE parameters only from safetensors files."
|
|
250
|
+
)
|
|
251
|
+
weights = next((f"{prefix}{w}" for w in WEIGHT_FILES if f"{prefix}{w}" in sizes), None)
|
|
252
|
+
if weights is None or f"{prefix}cfg.json" not in sizes:
|
|
253
|
+
raise BackendError(
|
|
254
|
+
f"{repo} has no SAE at {path or 'its top level'}: Logogram needs cfg.json with "
|
|
255
|
+
"sae_weights.safetensors or sae.safetensors there."
|
|
256
|
+
)
|
|
257
|
+
files = hub.RepoFiles(
|
|
258
|
+
revision=str(info.sha),
|
|
259
|
+
files=[(f"{prefix}cfg.json", sizes[f"{prefix}cfg.json"]), (weights, sizes[weights])],
|
|
260
|
+
n_params=None,
|
|
261
|
+
gated=bool(info.gated),
|
|
262
|
+
)
|
|
263
|
+
folder = hub.download(repo, files, progress, cancel=cancel)
|
|
264
|
+
return str(info.sha), folder
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def _cached(repo: str, prefix: str, revision: str | None) -> tuple[str, Path] | None:
|
|
268
|
+
from huggingface_hub import try_to_load_from_cache
|
|
269
|
+
|
|
270
|
+
found = try_to_load_from_cache(repo, f"{prefix}cfg.json", revision=revision or "main")
|
|
271
|
+
if not isinstance(found, str):
|
|
272
|
+
return None
|
|
273
|
+
folder = Path(found).parent
|
|
274
|
+
if not any((folder / w).is_file() for w in WEIGHT_FILES):
|
|
275
|
+
return None
|
|
276
|
+
snapshot = folder if not prefix else Path(found).parents[prefix.count("/")]
|
|
277
|
+
return snapshot.name, folder
|