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
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
|