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.
Files changed (54) hide show
  1. logogram/__init__.py +6 -0
  2. logogram/__main__.py +5 -0
  3. logogram/analysis.py +419 -0
  4. logogram/atp.py +120 -0
  5. logogram/backends/__init__.py +5 -0
  6. logogram/backends/base.py +202 -0
  7. logogram/backends/hub.py +375 -0
  8. logogram/backends/saes.py +277 -0
  9. logogram/backends/transformer_lens.py +872 -0
  10. logogram/cli.py +496 -0
  11. logogram/compare.py +177 -0
  12. logogram/datasets.py +159 -0
  13. logogram/direct.py +193 -0
  14. logogram/engine.py +550 -0
  15. logogram/examples/ioi-gpt2/.gitignore +3 -0
  16. logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
  17. logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
  18. logogram/examples/ioi-gpt2/project.json +6 -0
  19. logogram/exports.py +33 -0
  20. logogram/features.py +368 -0
  21. logogram/fileio.py +63 -0
  22. logogram/ioi.py +220 -0
  23. logogram/paths.py +204 -0
  24. logogram/project.py +444 -0
  25. logogram/prompts.py +204 -0
  26. logogram/research.py +84 -0
  27. logogram/results.py +240 -0
  28. logogram/runner.py +396 -0
  29. logogram/runs.py +98 -0
  30. logogram/sae.py +161 -0
  31. logogram/schema.py +302 -0
  32. logogram/server/__init__.py +1 -0
  33. logogram/server/app.py +1083 -0
  34. logogram/server/models.py +426 -0
  35. logogram/server/security.py +212 -0
  36. logogram/server/state.py +585 -0
  37. logogram/sites.py +249 -0
  38. logogram/spec.py +518 -0
  39. logogram/stats.py +171 -0
  40. logogram/steering.py +258 -0
  41. logogram/system.py +379 -0
  42. logogram/updates.py +194 -0
  43. logogram/verify.py +39 -0
  44. logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
  45. logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
  46. logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
  47. logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
  48. logogram/web_dist/favicon.svg +1 -0
  49. logogram/web_dist/index.html +15 -0
  50. logogram-0.1.0.dist-info/METADATA +550 -0
  51. logogram-0.1.0.dist-info/RECORD +54 -0
  52. logogram-0.1.0.dist-info/WHEEL +4 -0
  53. logogram-0.1.0.dist-info/entry_points.txt +2 -0
  54. 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