Skip to content

Commit 82e1068

Browse files
authored
Docs/additional hooked transformer cleanup (#1818)
* Removed remaining references to HookedTransformer * Fragile MPS test fix
1 parent d88c7d3 commit 82e1068

33 files changed

Lines changed: 177 additions & 157 deletions

‎.claude/commands/add-model-support.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ Each step names the doc to read **when you reach that step** — don't load all
4646
4747
6. **Add the HF repo entry** to [`data/supported_models.json`](../../transformer_lens/tools/model_registry/data/supported_models.json) per [§Adding the HF repo to the registry](../../transformer_lens/model_bridge/supported_architectures/AGENTS.md#adding-the-hf-repo-to-the-registry). Ask the user about adding canonical sibling variants from `CANONICAL_AUTHORS_BY_ARCH[<HFArchClass>]`.
4848
49-
7. **Verify** end-to-end: `/verify-model $ARGUMENTS`. Read both `status` AND per-phase scores. `STATUS_VERIFIED` means hard gates passed (see [§Phase-score thresholds](../../transformer_lens/tools/model_registry/AGENTS.md#phase-score-thresholds)) — but P4's 50% bar is intentionally lenient. P4 well below 100% on a small parity-test model + `status==1` → suspect missing `preprocess_weights` fold or wrong `default_prepend_bos`; investigate before step 8.
49+
7. **Verify** end-to-end: `/verify-model $ARGUMENTS`. Read both `status` AND per-phase scores. `STATUS_VERIFIED` means hard gates passed (see [§Phase-score thresholds](../../transformer_lens/tools/model_registry/AGENTS.md#phase-score-thresholds)) — but P4's ~54.5% bar (`p4_pass_threshold()`, not a fixed number) is intentionally lenient. P4 well below 100% on a small parity-test model + `status==1` → suspect missing `preprocess_weights` fold or wrong `default_prepend_bos`; investigate before step 8.
5050
5151
8. **Write tests** per [§Required tests](../../transformer_lens/model_bridge/supported_architectures/AGENTS.md#required-tests) (unit + integration). Copy the closest sibling.
5252

‎.claude/commands/verify-model.md‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@ Full reference: [tools/model_registry/AGENTS.md §Flag reference](../../transfor
3737
- `--max-memory <gb>` — skip if param estimate exceeds; e.g. `16` on a 24 GB GPU leaves headroom for activations
3838
- `--phases 1 2 3` — restrict (P4 is slowest; restrict when debugging P1 forward parity)
3939
- `--dry-run` — see above; always first
40-
- `--no-hf-reference` / `--no-ht-reference` — skip HF / HT comparison (faster, lower confidence)
40+
- `--no-hf-reference` — skip the HF reference comparison (faster, lower confidence; the run can only reach `STATUS_PROVISIONAL`, never verified)
4141
- `--reverify` — re-test `status==1`
4242
- `--retry-failed` — re-test `status==3` (read existing `note` first)
4343

@@ -52,9 +52,10 @@ Hard thresholds (`_MIN_PHASE_SCORES` in `verify_models.py`):
5252
| 1 | 100% | — | `STATUS_FAILED` |
5353
| 2 | 75% | `logits_equivalence`, `loss_equivalence` | `STATUS_FAILED` |
5454
| 3 | 75% | `logits_equivalence`, `loss_equivalence` | `STATUS_FAILED` |
55-
| 4 | 50% | — | **Non-gating** — adds `"low text quality"` to `note`; never fails. |
55+
| 4 | ~54.5% (`p4_pass_threshold()`, derived from the judge bake-off noise floor — not a fixed number) | — | **Non-gating** — adds `"low text quality"` to `note`; never fails. |
5656
| 7 | 75% | `multimodal_forward` | `STATUS_FAILED`. NULL = fail. |
57-
| 8 | 75% | `audio_forward` | `STATUS_FAILED`. NULL = fail. |
57+
| 8 | 75% | `audio_forward`, `audio_text_forward` | `STATUS_FAILED`. NULL = fail. |
58+
| 9 | 75% | `vision_forward`, `vision_cache` | `STATUS_FAILED`. NULL = fail. |
5859

5960
`STATUS_VERIFIED` means hard gates passed. `note` carries quality flags or failure details.
6061

‎AGENTS.md‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ make unit-test # fast, no model loads
5858
make integration-test # cross-component
5959
make acceptance-test # end-to-end
6060
make docstring-test # doctest + doctest-plus
61-
make notebook-test # slow; subset run in CI
61+
make notebook-test # slow, local-only; CI runs its own per-notebook matrix (notebook-checks in checks.yml), not this target
6262
make test-pr # unit + docstring + acceptance + integration (PR-review surface)
6363
make test # everything (long; includes benchmarks + notebooks)
6464

@@ -213,9 +213,9 @@ Load-bearing pins live in [pyproject.toml](pyproject.toml):
213213

214214
| Pin | Where | Why it matters |
215215
|---|---|---|
216-
| `transformers>=5.4.0` | `[project] dependencies` | The Bridge adapter contract is written against HF module layouts; every minor HF release can break adapter component-mappings. Bumping is a real test pass. |
216+
| `transformers>=5.9.0` | `[project] dependencies` | The Bridge adapter contract is written against HF module layouts; every minor HF release can break adapter component-mappings. Bumping is a real test pass. |
217217
| `torch>=2.6` | `[project] dependencies` | Hook system relies on PyTorch's forward / backward hook semantics; major torch bumps occasionally change ordering. |
218-
| `accelerate>=0.23.0` | `[project] dependencies` | Required for Llama-family loading. |
218+
| `accelerate>=1.1.0` | `[project] dependencies` | Required for Llama-family loading and `device_map` disk offload (`align_module_device`). |
219219
| `numpy>=1.24` / `>=1.26` | `[project] dependencies` (python-version-conditional) | Doctest float formatting can drift across NumPy versions. |
220220
| `isort==5.8.0` | `[dependency-groups] dev` (exact) | Format check pins to exactly this version; a bump flips the formatting of every file. |
221221

‎CLAUDE.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323

2424
- [AGENTS.md §10](AGENTS.md#10-hard-rules) — hard rules; load-bearing.
2525
- [AGENTS.md §2](AGENTS.md#2-two-systems-live-in-this-repo) — the single model system (the `Hooked*` model classes were removed in 4.0).
26-
- [tests/QUARANTINES.md](tests/QUARANTINES.md) — check before debugging any failing test. The macOS-arm64 KV-cache skip is the most common time-sink.
26+
- [tests/QUARANTINES.md](tests/QUARANTINES.md) — check before debugging any failing test. A skip with a matching reason is not your bug; a failure not on that list is real.
2727
- [debugging_numerical_divergence.md](docs/source/content/debugging_numerical_divergence.md) — Bridge-vs-HF logit drift bisection.
2828
- [compatibility_mode.md](docs/source/content/compatibility_mode.md) — `bridge.enable_compatibility_mode()` contract; read before adding tests that use it.
2929
- [devtools/adapter_builder/](devtools/adapter_builder/README.md) — autonomous adapter builder (contributor tooling; agent-teams mode needs Max, solo mode runs on any tier); manual path is `/add-model-support`.

‎README.md‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -177,14 +177,14 @@ cover:
177177
`dt_proj`, `out_proj` for Mamba-1; `in_proj`, `conv1d`, `inner_norm`, `out_proj` for
178178
Mamba-2)
179179
* Stateful generation with cache-aware decode steps
180-
* The `compute_effective_attention` utility (in
181-
`transformer_lens.model_bridge.supported_architectures.mamba2`) that materializes
182-
Mamba-2's SSD-derived attention matrix for comparison with transformer attention
183-
patterns
180+
* `ActivationCache.compute_ssm_effective_attention(layer=...)`, which materializes the
181+
SSD-derived attention matrix of a Mamba-2 (or any other SSM / hybrid) layer for
182+
comparison with transformer attention patterns; the older
183+
`mamba2.compute_effective_attention` remains only as a deprecated alias
184184

185185
Verification lives in the integration tests at
186186
`tests/integration/model_bridge/test_mamba_adapter.py` and
187-
`tests/integration/model_bridge/test_mamba2_adapter.py` (81 tests total), and the
187+
`tests/integration/model_bridge/test_mamba2_adapter.py`, and the
188188
`verify_models` benchmark suite now covers the SSM and hybrid families. Mamba-1,
189189
Mamba-2, gated-delta-net (Qwen3.5 / Qwen3-Next), NemotronH, and GraniteMoeHybrid
190190
all declare `applicable_phases = [1, 2, 3, 4]`, so their forward parity (P1, vs raw

‎demos/Exploratory_Analysis_Demo.ipynb‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1936,7 +1936,7 @@
19361936
"cell_type": "markdown",
19371937
"metadata": {},
19381938
"source": [
1939-
"The above suggests that it would be a useful bit of infrastructure to have a \"wiki\" for the heads of a model, giving their scores according to some metrics re head functions, like the ones we've seen here. TransformerLens makes this easy to make, as just changing the name input to `HookedTransformer.from_pretrained` gives a different model but in the same architecture, so the same code should work. If you want to make this, I'd love to see it! \n",
1939+
"The above suggests that it would be a useful bit of infrastructure to have a \"wiki\" for the heads of a model, giving their scores according to some metrics re head functions, like the ones we've seen here. TransformerLens makes this easy to make, as just changing the name input to `TransformerBridge.boot_transformers` gives a different model but in the same architecture, so the same code should work. If you want to make this, I'd love to see it! \n",
19401940
"\n",
19411941
"As a proof of concept, [I made a mosaic of all induction heads across the 40 models then in TransformerLens](https://www.neelnanda.io/mosaic).\n",
19421942
"\n",
@@ -2094,7 +2094,7 @@
20942094
],
20952095
"metadata": {
20962096
"kernelspec": {
2097-
"display_name": "transformer-lens",
2097+
"display_name": "transformer-lens (3.12.12)",
20982098
"language": "python",
20992099
"name": "python3"
21002100
},

‎demos/Inspect_Bridge_Demo.ipynb‎

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
"# Inspect Driver — Demo\n",
99
"\n",
1010
"Turn a model served through [`inspect_ai`](https://inspect.aisi.org.uk) into a\n",
11-
"TransformerLens `HookedTransformer`: `run_with_cache`, named hook points, and\n",
11+
"TransformerLens `TransformerBridge`: `run_with_cache`, named hook points, and\n",
1212
"interventions all work over the Inspect boundary.\n",
1313
"\n",
1414
"`boot_inspect` boots the model behind an `inspect_ai` provider and wraps it in a\n",
@@ -679,7 +679,15 @@
679679
}
680680
},
681681
"outputs": [],
682-
"source": "# (9) Tool-aware generation: for a tool-capable instruct model the provider renders tool\n# schemas into the chat template and parses tool calls from the output (so Inspect agent\n# loops run). gpt2 has no tool template, so here we just show the tool-call parser:\nfrom transformer_lens.model_bridge.sources.inspect.transformers_provider import _parse_tool_calls\n\ncalls = _parse_tool_calls('<tool_call>{\"name\": \"add\", \"arguments\": {\"a\": 2, \"b\": 2}}</tool_call>')\nprint(\"parsed tool calls:\", [(c.function, c.arguments) for c in calls])"
682+
"source": [
683+
"# (9) Tool-aware generation: for a tool-capable instruct model the provider renders tool\n",
684+
"# schemas into the chat template and parses tool calls from the output (so Inspect agent\n",
685+
"# loops run). gpt2 has no tool template, so here we just show the tool-call parser:\n",
686+
"from transformer_lens.model_bridge.sources.inspect.transformers_provider import _parse_tool_calls\n",
687+
"\n",
688+
"calls = _parse_tool_calls('<tool_call>{\"name\": \"add\", \"arguments\": {\"a\": 2, \"b\": 2}}</tool_call>')\n",
689+
"print(\"parsed tool calls:\", [(c.function, c.arguments) for c in calls])"
690+
]
683691
},
684692
{
685693
"cell_type": "markdown",
@@ -706,7 +714,7 @@
706714
],
707715
"metadata": {
708716
"kernelspec": {
709-
"display_name": "transformer-lens",
717+
"display_name": "transformer-lens (3.12.12)",
710718
"language": "python",
711719
"name": "python3"
712720
},
@@ -3112,4 +3120,4 @@
31123120
},
31133121
"nbformat": 4,
31143122
"nbformat_minor": 5
3115-
}
3123+
}

‎demos/LIT_Integration_Demo.ipynb‎

Lines changed: 6 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -122,8 +122,8 @@
122122
"\n",
123123
"# LIT integration imports\n",
124124
"from transformer_lens.lit import (\n",
125-
" HookedTransformerLIT,\n",
126-
" HookedTransformerLITConfig,\n",
125+
" TransformerLensLIT,\n",
126+
" TransformerLensLITConfig,\n",
127127
" SimpleTextDataset,\n",
128128
" PromptCompletionDataset,\n",
129129
" IOIDataset,\n",
@@ -155,36 +155,10 @@
155155
},
156156
{
157157
"cell_type": "code",
158-
"execution_count": 3,
158+
"execution_count": null,
159159
"id": "18cbabbd",
160160
"metadata": {},
161-
"outputs": [
162-
{
163-
"name": "stdout",
164-
"output_type": "stream",
165-
"text": [
166-
"Loading gpt2-small...\n"
167-
]
168-
},
169-
{
170-
"name": "stderr",
171-
"output_type": "stream",
172-
"text": [
173-
"`torch_dtype` is deprecated! Use `dtype` instead!\n"
174-
]
175-
},
176-
{
177-
"name": "stdout",
178-
"output_type": "stream",
179-
"text": [
180-
"Loaded pretrained model gpt2-small into HookedTransformer\n",
181-
"Loaded model: gpt2\n",
182-
" Layers: 12\n",
183-
" Heads: 12\n",
184-
" d_model: 768\n"
185-
]
186-
}
187-
],
161+
"outputs": [],
188162
"source": [
189163
"# Load GPT-2 (124M parameters)\n",
190164
"# Other options: \"gpt2-medium\", \"gpt2-large\", \"gpt2-xl\", \"EleutherAI/pythia-70m\", etc.\n",
@@ -231,7 +205,7 @@
231205
],
232206
"source": [
233207
"# Configure the wrapper\n",
234-
"config = HookedTransformerLITConfig(\n",
208+
"config = TransformerLensLITConfig(\n",
235209
" max_seq_length=256, # Maximum input length\n",
236210
" batch_size=4, # Batch size for inference\n",
237211
" top_k=10, # Number of top predictions to show\n",
@@ -243,7 +217,7 @@
243217
")\n",
244218
"\n",
245219
"# Create the wrapper\n",
246-
"lit_model = HookedTransformerLIT(model, config=config)\n",
220+
"lit_model = TransformerLensLIT(model, config=config)\n",
247221
"\n",
248222
"print(f\"Created LIT wrapper: {lit_model.description()}\")\n",
249223
"print(f\"\\nInput spec keys: {list(lit_model.input_spec().keys())}\")\n",

‎demos/Main_Demo.ipynb‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -995,7 +995,7 @@
995995
"\n",
996996
"LayerNorm is a normalization technique used by transformers, analogous to BatchNorm but more friendly to massive parallelisation. No one *really* knows why it works, but it seems to improve model numerical stability. Unlike BatchNorm, LayerNorm actually changes the functional form of the model, which makes it a massive pain for interpretability! \n",
997997
"\n",
998-
"Folding LayerNorm is a technique to make it lower overhead to deal with, and the flags `center_writing_weights` and `fold_ln` in `HookedTransformer.from_pretrained` apply this automatically (they default to True). These simplify the internal structure without changing the weights.\n",
998+
"Folding LayerNorm is a technique to make it lower overhead to deal with, and the flags `fold_ln` and `center_writing_weights` in `model.enable_compatibility_mode()` apply this automatically (they default to True). These simplify the internal structure without changing the weights.\n",
999999
"\n",
10001000
"Intuitively, LayerNorm acts on each residual stream vector (ie for each batch element and token position) independently, sets their mean to 0 (centering) and standard deviation to 1 (normalizing) (*across* the residual stream dimension - very weird!), and then applies a learned elementwise scaling and translation to each vector.\n",
10011001
"\n",
@@ -1807,7 +1807,7 @@
18071807
"source": [
18081808
"## Loading Pre-Trained Checkpoints\n",
18091809
"\n",
1810-
"**Note:** `TransformerBridge.boot_transformers` mainly works with HuggingFace revisions, and is only compatible with checkpoints for a few model families. For other models with saved checkpoints, keep using `HookedTransformer.from_pretrained`.\n",
1810+
"**Note:** `TransformerBridge.boot_transformers` mainly works with HuggingFace revisions, and is only compatible with checkpoints for a few model families. Legacy TransformerLens-format repos (`NeelNanda/*`, `ArthurConmy/*`, `Baidicoot/*`) have no HF `model_type`, so `boot_transformers` cannot load them; use `TransformerBridge.boot_tl_legacy(model_name, checkpoint_index=...)` (or `checkpoint_value=...`) for those, as the code below does.\n",
18111811
"\n",
18121812
"There are a lot of interesting questions combining mechanistic interpretability and training dynamics - analysing model capabilities and the underlying circuits that make them possible, and how these change as we train the model. \n",
18131813
"\n",

‎devtools/adapter_builder/CLAUDE.md‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -107,8 +107,8 @@ Nothing is copied into the TransformerLens worktree except `.claude/agents/` (re
107107
- `.env` — see `.env.example` for all variables; gitignored, never commit it
108108
- `HF_TOKEN` — HuggingFace API token (falls back to the repo root `.env` if unset here)
109109
- `DEFAULT_TARGET_REPO` — optional override; defaults to the containing repo
110-
- `DEFAULT_BASE_BRANCH` — default branch (dev-4.x)
111-
- `DEFAULT_MAX_MEMORY_GB` — memory limit for verification (96)
110+
- `DEFAULT_BASE_BRANCH` — base branch for worktrees (`dev` in `.env.example`; the launch scripts fall back to the stale, fully-merged `dev-4.x` if it is unset, so set it)
111+
- `DEFAULT_MAX_MEMORY_GB` — memory limit for verification (`48` in `.env.example`; script fallback `96`)
112112
- `WORKTREE_BASE` — where agent pair worktrees live; optional, defaults to `<parent-of-TransformerLens>/worktrees`
113113
- `NOTIFICATION_WEBHOOK_URL` — Slack webhook for notifications
114114
- `NOTIFICATION_NUMBER` — iMessage fallback

0 commit comments

Comments
 (0)