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
logogram/sites.py ADDED
@@ -0,0 +1,249 @@
1
+ """Expand a spec's scope into concrete sites, and lay them out as a grid.
2
+
3
+ Every result is a grid of rows x columns (layer x head, layer x position, layer x component, or
4
+ one row per chosen site). The model map and the heatmaps draw from this layout.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from dataclasses import dataclass
10
+ from typing import Any
11
+
12
+ from logogram.backends.base import ModelInfo
13
+ from logogram.prompts import PreparedPrompt, common_labels
14
+ from logogram.spec import (
15
+ AllPositions,
16
+ FeaturesScope,
17
+ HeadsScope,
18
+ IndexPosition,
19
+ LabelPosition,
20
+ LastPosition,
21
+ LayerComponentsScope,
22
+ LayerPositionScope,
23
+ Site,
24
+ SitesScope,
25
+ Spec,
26
+ )
27
+
28
+ COMPONENT_LABELS = {
29
+ "resid_pre": "resid pre",
30
+ "resid_mid": "resid mid",
31
+ "resid_post": "resid post",
32
+ "attn_out": "attn",
33
+ "mlp_out": "mlp",
34
+ }
35
+
36
+
37
+ class ScopeError(ValueError):
38
+ pass
39
+
40
+
41
+ @dataclass
42
+ class ResolvedSite:
43
+ index: int
44
+ site: Site
45
+ row: int
46
+ col: int
47
+ label: str
48
+ # A variant of the same site, for methods that measure one site several ways (a steering
49
+ # strength, or its random control). ``variant_key`` names it, for example "×2".
50
+ variant: dict[str, Any] | None = None
51
+ variant_key: str | None = None
52
+
53
+ @property
54
+ def kind(self) -> str:
55
+ return self.site.kind
56
+
57
+ @property
58
+ def layer(self) -> int:
59
+ return self.site.layer
60
+
61
+ @property
62
+ def head(self) -> int | None:
63
+ return self.site.head
64
+
65
+ def position_key(self) -> str:
66
+ return position_key(self.site.position)
67
+
68
+ def to_dict(self) -> dict[str, Any]:
69
+ return {
70
+ "index": self.index,
71
+ "kind": self.kind,
72
+ "layer": self.layer,
73
+ "head": self.head,
74
+ "feature": self.site.feature,
75
+ "position": self.site.position.model_dump(),
76
+ "position_key": self.position_key(),
77
+ "row": self.row,
78
+ "col": self.col,
79
+ "label": self.label,
80
+ "variant": self.variant,
81
+ "variant_key": self.variant_key,
82
+ }
83
+
84
+
85
+ def position_key(position: AllPositions | LastPosition | IndexPosition | LabelPosition) -> str:
86
+ if isinstance(position, AllPositions):
87
+ return "all"
88
+ if isinstance(position, LastPosition):
89
+ return "last"
90
+ if isinstance(position, IndexPosition):
91
+ return str(position.index)
92
+ return position.label
93
+
94
+
95
+ def site_label(site: Site) -> str:
96
+ pos = "" if isinstance(site.position, AllPositions) else f" @ {position_key(site.position)}"
97
+ if site.kind == "head":
98
+ return f"L{site.layer} H{site.head}{pos}"
99
+ if site.kind == "sae_feature":
100
+ return f"L{site.layer} F{site.feature}{pos}"
101
+ return f"L{site.layer} {COMPONENT_LABELS[site.kind]}{pos}"
102
+
103
+
104
+ def resolve_position(
105
+ position: AllPositions | LastPosition | IndexPosition | LabelPosition, prompt: PreparedPrompt
106
+ ) -> int | None:
107
+ """The token index a position refers to in this prompt, or None for all positions."""
108
+ if isinstance(position, AllPositions):
109
+ return None
110
+ if isinstance(position, LastPosition):
111
+ return prompt.length - 1
112
+ if isinstance(position, IndexPosition):
113
+ idx = position.index if position.index >= 0 else prompt.length + position.index
114
+ if not 0 <= idx < prompt.length:
115
+ raise ScopeError(
116
+ f"Token index {position.index} is outside prompt {prompt.index}, which has "
117
+ f"{prompt.length} tokens."
118
+ )
119
+ return idx
120
+ if position.label not in prompt.labels:
121
+ raise ScopeError(
122
+ f"Prompt {prompt.index} has no position named {position.label!r}. Add it to the "
123
+ "dataset's positions, or choose another position."
124
+ )
125
+ return prompt.labels[position.label]
126
+
127
+
128
+ def _check_site(site: Site, info: ModelInfo) -> None:
129
+ if site.layer >= info.n_layers:
130
+ raise ScopeError(
131
+ f"Layer {site.layer} doesn't exist; this model has {info.n_layers} layers."
132
+ )
133
+ if site.kind == "sae_feature":
134
+ return # checked against the SAE, which knows its layer and features
135
+ if site.kind not in info.site_kinds:
136
+ raise ScopeError(f"This model has no {site.kind} site in TransformerLens.")
137
+ if site.kind == "head" and site.head is not None and site.head >= info.n_heads:
138
+ raise ScopeError(f"Head {site.head} doesn't exist; this model has {info.n_heads} heads.")
139
+
140
+
141
+ def expand_scope(
142
+ spec: Spec, info: ModelInfo, prompts: list[PreparedPrompt]
143
+ ) -> tuple[list[ResolvedSite], dict[str, Any]]:
144
+ scope = spec.scope
145
+ sites: list[ResolvedSite] = []
146
+
147
+ def add(site: Site, row: int, col: int) -> None:
148
+ sites.append(ResolvedSite(len(sites), site, row, col, site_label(site)))
149
+
150
+ if isinstance(scope, HeadsScope):
151
+ for layer in range(info.n_layers):
152
+ for head in range(info.n_heads):
153
+ add(Site(kind="head", layer=layer, head=head, position=scope.position), layer, head)
154
+ layout = {
155
+ "kind": "heads",
156
+ "row_title": "Layer",
157
+ "col_title": "Head",
158
+ "rows": [{"key": str(r), "label": str(r)} for r in range(info.n_layers)],
159
+ "cols": [{"key": str(c), "label": str(c)} for c in range(info.n_heads)],
160
+ }
161
+ elif isinstance(scope, LayerPositionScope):
162
+ if scope.site not in info.site_kinds:
163
+ raise ScopeError(f"This model has no {scope.site} site in TransformerLens.")
164
+ cols: list[dict[str, Any]] = []
165
+ positions: list[AllPositions | LastPosition | IndexPosition | LabelPosition] = []
166
+ if scope.positions == "each":
167
+ lengths = sorted({p.length for p in prompts})
168
+ if len(lengths) > 1:
169
+ raise ScopeError(
170
+ f"Prompts have different token lengths ({lengths[0]}–{lengths[-1]}), so "
171
+ "positions don't line up. Use labelled positions instead, or a dataset "
172
+ "whose prompts share one template length."
173
+ )
174
+ first = prompts[0]
175
+ differs = set(first.differing_positions())
176
+ for j in range(first.length):
177
+ positions.append(IndexPosition(index=j))
178
+ cols.append(
179
+ {
180
+ "key": str(j),
181
+ "label": first.clean.tokens[j],
182
+ "position": j,
183
+ "clean": first.clean.tokens[j],
184
+ "corrupt": first.corrupt.tokens[j],
185
+ "differs": j in differs,
186
+ }
187
+ )
188
+ else:
189
+ labels = common_labels(prompts)
190
+ if not labels:
191
+ raise ScopeError(
192
+ "The dataset has no named positions shared by every prompt. Generate an "
193
+ "IOI dataset, or add positions to your JSONL."
194
+ )
195
+ for label in labels:
196
+ positions.append(LabelPosition(label=label))
197
+ cols.append({"key": label, "label": label})
198
+ for layer in range(info.n_layers):
199
+ for j, pos in enumerate(positions):
200
+ add(Site(kind=scope.site, layer=layer, position=pos), layer, j)
201
+ layout = {
202
+ "kind": "layer_position",
203
+ "site": scope.site,
204
+ "row_title": "Layer",
205
+ "col_title": "Position",
206
+ "rows": [{"key": str(r), "label": str(r)} for r in range(info.n_layers)],
207
+ "cols": cols,
208
+ }
209
+ elif isinstance(scope, LayerComponentsScope):
210
+ for kind in scope.components:
211
+ if kind not in info.site_kinds:
212
+ raise ScopeError(f"This model has no {kind} site in TransformerLens.")
213
+ for layer in range(info.n_layers):
214
+ for j, kind in enumerate(scope.components):
215
+ add(Site(kind=kind, layer=layer, position=scope.position), layer, j)
216
+ layout = {
217
+ "kind": "layer_components",
218
+ "row_title": "Layer",
219
+ "col_title": "Component",
220
+ "rows": [{"key": str(r), "label": str(r)} for r in range(info.n_layers)],
221
+ "cols": [{"key": k, "label": COMPONENT_LABELS[k]} for k in scope.components],
222
+ }
223
+ elif isinstance(scope, SitesScope):
224
+ for i, site in enumerate(scope.sites):
225
+ add(site, i, 0)
226
+ layout = {
227
+ "kind": "sites",
228
+ "row_title": "Site",
229
+ "col_title": "",
230
+ "rows": [{"key": str(i), "label": site_label(s)} for i, s in enumerate(scope.sites)],
231
+ "cols": [{"key": "effect", "label": "effect"}],
232
+ }
233
+ elif isinstance(scope, FeaturesScope):
234
+ raise ScopeError(
235
+ "Sweeping every SAE feature needs attribution patching, which estimates them all at "
236
+ "once. To patch features for real, choose them as sites."
237
+ )
238
+ else: # pragma: no cover - exhaustive
239
+ raise ScopeError(f"Unknown scope {scope!r}")
240
+
241
+ checked: set[str] = set()
242
+ for rs in sites:
243
+ _check_site(rs.site, info)
244
+ key = rs.site.position.model_dump_json()
245
+ if key not in checked:
246
+ checked.add(key)
247
+ for prompt in prompts:
248
+ resolve_position(rs.site.position, prompt)
249
+ return sites, layout