logogram 0.1.1__tar.gz → 0.1.2__tar.gz

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 (184) hide show
  1. {logogram-0.1.1 → logogram-0.1.2}/CHANGELOG.md +34 -0
  2. {logogram-0.1.1 → logogram-0.1.2}/PKG-INFO +24 -10
  3. {logogram-0.1.1 → logogram-0.1.2}/README.md +22 -9
  4. {logogram-0.1.1 → logogram-0.1.2}/pyproject.toml +2 -1
  5. logogram-0.1.2/scripts/validate_real_weights.py +459 -0
  6. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/__init__.py +1 -1
  7. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/atp.py +10 -0
  8. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/backends/transformer_lens.py +3 -5
  9. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/cli.py +10 -23
  10. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/results.py +26 -1
  11. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/server/app.py +23 -4
  12. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/steering.py +29 -0
  13. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/system.py +24 -4
  14. logogram-0.1.1/src/logogram/web_dist/assets/index-CC9_IJY6.js → logogram-0.1.2/src/logogram/web_dist/assets/index-Cw0WBZBB.js +6 -6
  15. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/web_dist/index.html +1 -1
  16. {logogram-0.1.1 → logogram-0.1.2}/tests/test_atp.py +20 -0
  17. {logogram-0.1.1 → logogram-0.1.2}/tests/test_devices.py +14 -0
  18. {logogram-0.1.1 → logogram-0.1.2}/tests/test_hardening.py +39 -0
  19. {logogram-0.1.1 → logogram-0.1.2}/tests/test_steering.py +26 -0
  20. logogram-0.1.2/tests/test_validation_script.py +45 -0
  21. {logogram-0.1.1 → logogram-0.1.2}/uv.lock +50 -2
  22. {logogram-0.1.1 → logogram-0.1.2}/web/src/api/client.ts +3 -5
  23. {logogram-0.1.1 → logogram-0.1.2}/web/src/api/types.ts +4 -0
  24. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ModelDialog.tsx +14 -0
  25. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/RobustnessDialog.tsx +15 -3
  26. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/SystemCheck.tsx +1 -1
  27. {logogram-0.1.1 → logogram-0.1.2}/web/src/store/app.ts +4 -4
  28. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/BaselineView.tsx +13 -3
  29. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/PromptsView.tsx +30 -0
  30. {logogram-0.1.1 → logogram-0.1.2}/.gitignore +0 -0
  31. {logogram-0.1.1 → logogram-0.1.2}/AGENTS.md +0 -0
  32. {logogram-0.1.1 → logogram-0.1.2}/LICENSE +0 -0
  33. {logogram-0.1.1 → logogram-0.1.2}/scripts/check_privacy.py +0 -0
  34. {logogram-0.1.1 → logogram-0.1.2}/scripts/smoke_wheel.py +0 -0
  35. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/__main__.py +0 -0
  36. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/analysis.py +0 -0
  37. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/backends/__init__.py +0 -0
  38. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/backends/base.py +0 -0
  39. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/backends/hub.py +0 -0
  40. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/backends/saes.py +0 -0
  41. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/compare.py +0 -0
  42. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/datasets.py +0 -0
  43. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/direct.py +0 -0
  44. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/engine.py +0 -0
  45. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/examples/ioi-gpt2/.gitignore +0 -0
  46. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/examples/ioi-gpt2/datasets/ioi.jsonl +0 -0
  47. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +0 -0
  48. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/examples/ioi-gpt2/project.json +0 -0
  49. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/exports.py +0 -0
  50. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/features.py +0 -0
  51. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/fileio.py +0 -0
  52. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/ioi.py +0 -0
  53. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/paths.py +0 -0
  54. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/project.py +0 -0
  55. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/prompts.py +0 -0
  56. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/research.py +0 -0
  57. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/runner.py +0 -0
  58. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/runs.py +0 -0
  59. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/sae.py +0 -0
  60. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/schema.py +0 -0
  61. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/server/__init__.py +0 -0
  62. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/server/models.py +0 -0
  63. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/server/security.py +0 -0
  64. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/server/state.py +0 -0
  65. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/sites.py +0 -0
  66. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/spec.py +0 -0
  67. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/stats.py +0 -0
  68. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/updates.py +0 -0
  69. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/verify.py +0 -0
  70. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/web_dist/assets/index-Dt6ecB66.css +0 -0
  71. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
  72. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
  73. {logogram-0.1.1 → logogram-0.1.2}/src/logogram/web_dist/favicon.svg +0 -0
  74. {logogram-0.1.1 → logogram-0.1.2}/tests/browser_server.py +0 -0
  75. {logogram-0.1.1 → logogram-0.1.2}/tests/conftest.py +0 -0
  76. {logogram-0.1.1 → logogram-0.1.2}/tests/test_architectures.py +0 -0
  77. {logogram-0.1.1 → logogram-0.1.2}/tests/test_cli.py +0 -0
  78. {logogram-0.1.1 → logogram-0.1.2}/tests/test_direct.py +0 -0
  79. {logogram-0.1.1 → logogram-0.1.2}/tests/test_engine.py +0 -0
  80. {logogram-0.1.1 → logogram-0.1.2}/tests/test_explorer.py +0 -0
  81. {logogram-0.1.1 → logogram-0.1.2}/tests/test_loading.py +0 -0
  82. {logogram-0.1.1 → logogram-0.1.2}/tests/test_paths.py +0 -0
  83. {logogram-0.1.1 → logogram-0.1.2}/tests/test_privacy.py +0 -0
  84. {logogram-0.1.1 → logogram-0.1.2}/tests/test_release_fixes.py +0 -0
  85. {logogram-0.1.1 → logogram-0.1.2}/tests/test_sae.py +0 -0
  86. {logogram-0.1.1 → logogram-0.1.2}/tests/test_sanity.py +0 -0
  87. {logogram-0.1.1 → logogram-0.1.2}/tests/test_server.py +0 -0
  88. {logogram-0.1.1 → logogram-0.1.2}/tests/test_units.py +0 -0
  89. {logogram-0.1.1 → logogram-0.1.2}/tests/test_updates.py +0 -0
  90. {logogram-0.1.1 → logogram-0.1.2}/web/index.html +0 -0
  91. {logogram-0.1.1 → logogram-0.1.2}/web/package-lock.json +0 -0
  92. {logogram-0.1.1 → logogram-0.1.2}/web/package.json +0 -0
  93. {logogram-0.1.1 → logogram-0.1.2}/web/playwright.config.ts +0 -0
  94. {logogram-0.1.1 → logogram-0.1.2}/web/public/favicon.svg +0 -0
  95. {logogram-0.1.1 → logogram-0.1.2}/web/src/App.module.css +0 -0
  96. {logogram-0.1.1 → logogram-0.1.2}/web/src/App.tsx +0 -0
  97. {logogram-0.1.1 → logogram-0.1.2}/web/src/api/events.ts +0 -0
  98. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/CommandPalette.module.css +0 -0
  99. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/CommandPalette.tsx +0 -0
  100. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/CopyCommand.module.css +0 -0
  101. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/CopyCommand.tsx +0 -0
  102. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Distribution.module.css +0 -0
  103. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Distribution.tsx +0 -0
  104. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/FolderPicker.module.css +0 -0
  105. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/FolderPicker.tsx +0 -0
  106. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Header.module.css +0 -0
  107. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Header.tsx +0 -0
  108. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Heatmap/Heatmap.module.css +0 -0
  109. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Heatmap/Heatmap.tsx +0 -0
  110. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Heatmap/ScaleBar.tsx +0 -0
  111. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/History.module.css +0 -0
  112. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/History.tsx +0 -0
  113. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Inspector.module.css +0 -0
  114. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Inspector.tsx +0 -0
  115. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Logogram.tsx +0 -0
  116. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/LogogramDial.module.css +0 -0
  117. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/LogogramDial.tsx +0 -0
  118. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ModelDialog.module.css +0 -0
  119. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ModelMap/Legend.tsx +0 -0
  120. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ModelMap/ModelMap.module.css +0 -0
  121. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ModelMap/ModelMap.tsx +0 -0
  122. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Notices.module.css +0 -0
  123. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Notices.tsx +0 -0
  124. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Splitter.module.css +0 -0
  125. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/Splitter.tsx +0 -0
  126. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/StatusLine.module.css +0 -0
  127. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/StatusLine.tsx +0 -0
  128. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/TokenStrip.module.css +0 -0
  129. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/TokenStrip.tsx +0 -0
  130. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/TopBar.module.css +0 -0
  131. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/TopBar.tsx +0 -0
  132. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/UpdateNotice.module.css +0 -0
  133. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/UpdateNotice.tsx +0 -0
  134. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ui/Icon.tsx +0 -0
  135. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ui/index.tsx +0 -0
  136. {logogram-0.1.1 → logogram-0.1.2}/web/src/components/ui/ui.module.css +0 -0
  137. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/analysis.ts +0 -0
  138. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/canvas.ts +0 -0
  139. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/color.ts +0 -0
  140. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/export.ts +0 -0
  141. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/format.ts +0 -0
  142. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/heatmapNavigation.ts +0 -0
  143. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/hooks.ts +0 -0
  144. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/keys.ts +0 -0
  145. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/logogram.ts +0 -0
  146. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/sites.ts +0 -0
  147. {logogram-0.1.1 → logogram-0.1.2}/web/src/lib/spec.ts +0 -0
  148. {logogram-0.1.1 → logogram-0.1.2}/web/src/main.tsx +0 -0
  149. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/Projects.tsx +0 -0
  150. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/Screens.module.css +0 -0
  151. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/Welcome.tsx +0 -0
  152. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/Workbench.module.css +0 -0
  153. {logogram-0.1.1 → logogram-0.1.2}/web/src/screens/Workbench.tsx +0 -0
  154. {logogram-0.1.1 → logogram-0.1.2}/web/src/styles/global.css +0 -0
  155. {logogram-0.1.1 → logogram-0.1.2}/web/src/styles/tokens.css +0 -0
  156. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/AttentionView.module.css +0 -0
  157. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/AttentionView.tsx +0 -0
  158. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/BaselineView.module.css +0 -0
  159. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/CompareView.module.css +0 -0
  160. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/CompareView.tsx +0 -0
  161. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ExperimentView.module.css +0 -0
  162. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ExperimentView.tsx +0 -0
  163. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ExploreView.module.css +0 -0
  164. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ExploreView.tsx +0 -0
  165. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/FeaturesView.module.css +0 -0
  166. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/FeaturesView.tsx +0 -0
  167. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/Forest.module.css +0 -0
  168. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/Forest.tsx +0 -0
  169. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/HeadComparisonView.module.css +0 -0
  170. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/HeadComparisonView.tsx +0 -0
  171. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/PredictionsView.module.css +0 -0
  172. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/PredictionsView.tsx +0 -0
  173. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ResearchView.module.css +0 -0
  174. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ResearchView.tsx +0 -0
  175. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ResultsView.module.css +0 -0
  176. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/ResultsView.tsx +0 -0
  177. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/SpecView.module.css +0 -0
  178. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/SpecView.tsx +0 -0
  179. {logogram-0.1.1 → logogram-0.1.2}/web/src/views/views.module.css +0 -0
  180. {logogram-0.1.1 → logogram-0.1.2}/web/tests/context.unit.ts +0 -0
  181. {logogram-0.1.1 → logogram-0.1.2}/web/tests/logogram.unit.ts +0 -0
  182. {logogram-0.1.1 → logogram-0.1.2}/web/tests/workbench.browser.ts +0 -0
  183. {logogram-0.1.1 → logogram-0.1.2}/web/tsconfig.json +0 -0
  184. {logogram-0.1.1 → logogram-0.1.2}/web/vite.config.ts +0 -0
@@ -2,6 +2,40 @@
2
2
 
3
3
  What changed in each version of Logogram.
4
4
 
5
+ ## 0.1.2 (2026-10-08)
6
+
7
+ Fixes and checks that came out of validating every method on real weights.
8
+
9
+ ### Changed
10
+
11
+ - On Apple Silicon, **Automatic** runs models on the CPU. TransformerLens reports that Apple's MPS
12
+ can give silently wrong results, and Logogram hasn't been checked on it yet, so MPS runs only when
13
+ you choose it; the system check and the model dialog say why.
14
+
15
+ ### Added
16
+
17
+ - **Check robustness** reruns a float16 or bfloat16 run in float32 and compares the two.
18
+ - When processing a model's weights is what doesn't fit, the model dialog offers to turn it off.
19
+ - When a generated IOI dataset has names the loaded model splits into several tokens, one click
20
+ makes a new one with names it doesn't.
21
+ - Steering results say when the direction does no more than its random control at any site and
22
+ strength.
23
+ - Attribution patching of residual stream sites warns that it can miss effects there, and that
24
+ **Check robustness** patches every site.
25
+ - `scripts/validate_real_weights.py` runs every method on a real model and checks identities,
26
+ agreement and bit-identical reruns, on one device or two. It is how a release, or Apple's MPS, is
27
+ checked.
28
+
29
+ ### Fixed
30
+
31
+ - The baseline says that the clean prompts prefer the distractor, instead of preferring the answer
32
+ by a negative amount.
33
+
34
+ ### Development
35
+
36
+ - Tests run on Python 3.14 too, and the Linux runners are pinned to Ubuntu 24.04.
37
+ - The test client uses httpx2, as Starlette recommends.
38
+
5
39
  ## 0.1.1 (2026-10-08)
6
40
 
7
41
  Every method has now run end to end on real weights: GPT-2 small with a SAELens and an OpenAI
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: logogram
3
- Version: 0.1.1
3
+ Version: 0.1.2
4
4
  Summary: A local-first workbench for causal experiments inside language models.
5
5
  Project-URL: Homepage, https://github.com/Jeevash23/logogram
6
6
  Project-URL: Source, https://github.com/Jeevash23/logogram
@@ -17,6 +17,7 @@ Classifier: Programming Language :: Python :: 3
17
17
  Classifier: Programming Language :: Python :: 3.11
18
18
  Classifier: Programming Language :: Python :: 3.12
19
19
  Classifier: Programming Language :: Python :: 3.13
20
+ Classifier: Programming Language :: Python :: 3.14
20
21
  Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
21
22
  Requires-Python: >=3.11
22
23
  Requires-Dist: anyio>=4
@@ -73,7 +74,7 @@ PyTorch is chosen per machine:
73
74
  |---|---|---|
74
75
  | Linux with an NVIDIA GPU | `uv tool install logogram` | The default Linux wheels include CUDA. |
75
76
  | Windows with an NVIDIA GPU | `uv tool install --torch-backend=auto logogram` | Picks the CUDA build that matches your driver. |
76
- | Apple Silicon | `uv tool install logogram` | Uses Metal (MPS). Use a native arm64 Python, not one under Rosetta. |
77
+ | Apple Silicon | `uv tool install logogram` | Runs on the CPU unless you choose MPS (see Models). Use a native arm64 Python, not one under Rosetta. |
77
78
  | CPU only | `uv tool install --torch-backend=cpu logogram` | Smaller download. GPT-2 small runs well on a CPU. |
78
79
 
79
80
  Then check the setup:
@@ -114,8 +115,9 @@ This starts a local server on 127.0.0.1, prints its address and opens your brows
114
115
  4. Press **A** on a head to see its attention pattern, or right-click any cell for **Patch here**,
115
116
  **Ablate here** and **Compare across runs**.
116
117
  5. Press **Check robustness** to rerun the sweep with a different baseline, direction or donor
117
- count. Logogram reports the rank correlation, the overlap of the top components and the
118
- components whose conclusion changed, and flags them on the map.
118
+ count, or, for a run in float16 or bfloat16, in float32. Logogram reports the rank
119
+ correlation, the overlap of the top components and the components whose conclusion changed,
120
+ and flags them on the map.
119
121
 
120
122
  Keyboard: arrow keys move across the map, **Ctrl/⌘ K** opens the command palette (type `L9H9` to
121
123
  jump to a head), **1–7** switch views, **[** and **]** step through prompts, **P**, **B**, **A**
@@ -311,7 +313,7 @@ its effect at every strength with intervals. **Check robustness** offers another
311
313
  or the other direction. A mean difference steers only when the pairs differ the same way. In IOI
312
314
  prompts that mix the ABBA and BABA orders, the difference at the last token flips with the order,
313
315
  so the mean cancels and steering does no more than the control; generate prompts of one order to
314
- steer.
316
+ steer. When no site and strength does more than the control, the results say so.
315
317
 
316
318
  **SAE features.** A sparse autoencoder (SAE) rewrites one of the model's activations as a few active
317
319
  features out of thousands, each a direction in the model, plus an error it misses. In
@@ -490,7 +492,11 @@ times faster on a GPU, but they round. In the IOI example on GPT-2 small and Qwe
490
492
  moved effects by at most 0.003 and bfloat16 by up to 0.03: the strongest sites stayed in place,
491
493
  but effects smaller than about 0.01 changed order. On Pythia-70m, bfloat16 doubled the clean
492
494
  logit difference and reordered the heads, and float16 overflows. Check results that matter
493
- against float32.
495
+ against float32: **Check robustness** reruns a 16-bit run in float32 and compares the two.
496
+
497
+ On Apple Silicon, **Automatic** runs models on the CPU. TransformerLens reports that Apple's MPS
498
+ can give silently wrong results, and Logogram hasn't been checked on it yet, so MPS is used only
499
+ when you choose it under **Device**. It is faster; check results that matter on the CPU.
494
500
 
495
501
  Rather than trusting a list, Logogram checks every model when it loads, on a short fixed input:
496
502
 
@@ -561,10 +567,18 @@ push a tag with the same version, such as `v0.1.0`. The Release workflow runs ev
561
567
  that commit, builds the wheel and source distribution, and uploads them to PyPI through Trusted
562
568
  Publishing, so no token is stored anywhere. PyPI never accepts the same version twice.
563
569
 
564
- Before tagging, manually check a first GPT-2 download, cancel and retry it, then run the example
565
- on each supported compute backend. CPU and CUDA are covered by local development checks; MPS
566
- still needs a check on Apple Silicon. Pythia and Qwen 2.5 have also been run end to end with
567
- their real weights; other families are checked when they load and in the tiny-model tests.
570
+ Before tagging, manually check a first GPT-2 download, cancel and retry it, then run
571
+ `scripts/validate_real_weights.py` on each supported compute backend. It runs every method on a
572
+ real model in a temporary project and checks exact identities, agreement between methods and
573
+ bit-identical reruns; `--compare-device` reruns the head sweep on a second device, and `--sae`
574
+ adds SAE features. CPU and CUDA pass it. Apple's MPS still needs it run on a Mac:
575
+
576
+ ```bash
577
+ uv run python scripts/validate_real_weights.py --device mps --compare-device cpu
578
+ ```
579
+
580
+ Pythia and Qwen 2.5 have also been run end to end with their real weights; other families are
581
+ checked when they load and in the tiny-model tests.
568
582
 
569
583
  Model access goes through `logogram.backends.base.ModelBackend`. TransformerLens is the only
570
584
  backend today; remote execution and other libraries can be added behind the same interface.
@@ -36,7 +36,7 @@ PyTorch is chosen per machine:
36
36
  |---|---|---|
37
37
  | Linux with an NVIDIA GPU | `uv tool install logogram` | The default Linux wheels include CUDA. |
38
38
  | Windows with an NVIDIA GPU | `uv tool install --torch-backend=auto logogram` | Picks the CUDA build that matches your driver. |
39
- | Apple Silicon | `uv tool install logogram` | Uses Metal (MPS). Use a native arm64 Python, not one under Rosetta. |
39
+ | Apple Silicon | `uv tool install logogram` | Runs on the CPU unless you choose MPS (see Models). Use a native arm64 Python, not one under Rosetta. |
40
40
  | CPU only | `uv tool install --torch-backend=cpu logogram` | Smaller download. GPT-2 small runs well on a CPU. |
41
41
 
42
42
  Then check the setup:
@@ -77,8 +77,9 @@ This starts a local server on 127.0.0.1, prints its address and opens your brows
77
77
  4. Press **A** on a head to see its attention pattern, or right-click any cell for **Patch here**,
78
78
  **Ablate here** and **Compare across runs**.
79
79
  5. Press **Check robustness** to rerun the sweep with a different baseline, direction or donor
80
- count. Logogram reports the rank correlation, the overlap of the top components and the
81
- components whose conclusion changed, and flags them on the map.
80
+ count, or, for a run in float16 or bfloat16, in float32. Logogram reports the rank
81
+ correlation, the overlap of the top components and the components whose conclusion changed,
82
+ and flags them on the map.
82
83
 
83
84
  Keyboard: arrow keys move across the map, **Ctrl/⌘ K** opens the command palette (type `L9H9` to
84
85
  jump to a head), **1–7** switch views, **[** and **]** step through prompts, **P**, **B**, **A**
@@ -274,7 +275,7 @@ its effect at every strength with intervals. **Check robustness** offers another
274
275
  or the other direction. A mean difference steers only when the pairs differ the same way. In IOI
275
276
  prompts that mix the ABBA and BABA orders, the difference at the last token flips with the order,
276
277
  so the mean cancels and steering does no more than the control; generate prompts of one order to
277
- steer.
278
+ steer. When no site and strength does more than the control, the results say so.
278
279
 
279
280
  **SAE features.** A sparse autoencoder (SAE) rewrites one of the model's activations as a few active
280
281
  features out of thousands, each a direction in the model, plus an error it misses. In
@@ -453,7 +454,11 @@ times faster on a GPU, but they round. In the IOI example on GPT-2 small and Qwe
453
454
  moved effects by at most 0.003 and bfloat16 by up to 0.03: the strongest sites stayed in place,
454
455
  but effects smaller than about 0.01 changed order. On Pythia-70m, bfloat16 doubled the clean
455
456
  logit difference and reordered the heads, and float16 overflows. Check results that matter
456
- against float32.
457
+ against float32: **Check robustness** reruns a 16-bit run in float32 and compares the two.
458
+
459
+ On Apple Silicon, **Automatic** runs models on the CPU. TransformerLens reports that Apple's MPS
460
+ can give silently wrong results, and Logogram hasn't been checked on it yet, so MPS is used only
461
+ when you choose it under **Device**. It is faster; check results that matter on the CPU.
457
462
 
458
463
  Rather than trusting a list, Logogram checks every model when it loads, on a short fixed input:
459
464
 
@@ -524,10 +529,18 @@ push a tag with the same version, such as `v0.1.0`. The Release workflow runs ev
524
529
  that commit, builds the wheel and source distribution, and uploads them to PyPI through Trusted
525
530
  Publishing, so no token is stored anywhere. PyPI never accepts the same version twice.
526
531
 
527
- Before tagging, manually check a first GPT-2 download, cancel and retry it, then run the example
528
- on each supported compute backend. CPU and CUDA are covered by local development checks; MPS
529
- still needs a check on Apple Silicon. Pythia and Qwen 2.5 have also been run end to end with
530
- their real weights; other families are checked when they load and in the tiny-model tests.
532
+ Before tagging, manually check a first GPT-2 download, cancel and retry it, then run
533
+ `scripts/validate_real_weights.py` on each supported compute backend. It runs every method on a
534
+ real model in a temporary project and checks exact identities, agreement between methods and
535
+ bit-identical reruns; `--compare-device` reruns the head sweep on a second device, and `--sae`
536
+ adds SAE features. CPU and CUDA pass it. Apple's MPS still needs it run on a Mac:
537
+
538
+ ```bash
539
+ uv run python scripts/validate_real_weights.py --device mps --compare-device cpu
540
+ ```
541
+
542
+ Pythia and Qwen 2.5 have also been run end to end with their real weights; other families are
543
+ checked when they load and in the tiny-model tests.
531
544
 
532
545
  Model access goes through `logogram.backends.base.ModelBackend`. TransformerLens is the only
533
546
  backend today; remote execution and other libraries can be added behind the same interface.
@@ -20,6 +20,7 @@ classifiers = [
20
20
  "Programming Language :: Python :: 3.11",
21
21
  "Programming Language :: Python :: 3.12",
22
22
  "Programming Language :: Python :: 3.13",
23
+ "Programming Language :: Python :: 3.14",
23
24
  "Topic :: Scientific/Engineering :: Artificial Intelligence",
24
25
  ]
25
26
  dependencies = [
@@ -51,7 +52,7 @@ logogram = "logogram.cli:main"
51
52
  [dependency-groups]
52
53
  dev = [
53
54
  "pytest>=8",
54
- "httpx>=0.27",
55
+ "httpx2>=2.13",
55
56
  "ruff==0.16.10",
56
57
  ]
57
58
 
@@ -0,0 +1,459 @@
1
+ #!/usr/bin/env python3
2
+ """Run every Logogram method on a real model and check what must hold.
3
+
4
+ A release check, and the check for a compute backend Logogram hasn't been verified on (Apple's
5
+ MPS). It loads a model from Hugging Face (downloading it once into the usual cache), runs each
6
+ method on the bundled IOI example in a temporary project, and checks:
7
+
8
+ * exact identities: patching the residual stream at the only token that differs restores the
9
+ whole effect, and so does patching the final residual stream; path patching from the last layer
10
+ to the logits equals patching; direct effects add up to the logit difference; the same spec
11
+ twice gives bit-identical results;
12
+ * agreement expected of a trained model: attribution patching ranks heads like patching, and
13
+ verifying its strongest estimates reverses no sign;
14
+ * that every other method runs (ablations, steering, path patching, layer predictions, and SAE
15
+ features with --sae).
16
+
17
+ With --compare-device it runs the head sweep again on a second device and reports how far the two
18
+ disagree. Nothing is written outside a temporary folder, apart from the Hugging Face cache.
19
+
20
+ uv run python scripts/validate_real_weights.py
21
+ uv run python scripts/validate_real_weights.py --device mps --compare-device cpu
22
+ uv run python scripts/validate_real_weights.py --sae jbloom/GPT2-Small-SAEs-Reformatted \\
23
+ blocks.8.hook_resid_pre
24
+
25
+ Exits with status 1 if a check fails.
26
+ """
27
+
28
+ from __future__ import annotations
29
+
30
+ import argparse
31
+ import json
32
+ import shutil
33
+ import sys
34
+ import tempfile
35
+ import time
36
+ from collections.abc import Callable
37
+ from dataclasses import dataclass
38
+ from pathlib import Path
39
+ from typing import Any
40
+
41
+ # The tolerance of an identity in each dtype: float32 rounds in the last digits, 16-bit types in
42
+ # the second or third.
43
+ TOLERANCE = {"float32": 1e-3, "float16": 2e-2, "bfloat16": 5e-2}
44
+
45
+
46
+ @dataclass
47
+ class Ran:
48
+ """A finished run: its summary, its folder, and the spec it ran."""
49
+
50
+ summary: dict[str, Any]
51
+ folder: Path
52
+ spec: Any
53
+
54
+
55
+ @dataclass
56
+ class Check:
57
+ name: str
58
+ passed: bool | None # None: skipped
59
+ detail: str
60
+ exact: bool # holds for any model, not only a trained one
61
+
62
+
63
+ def validate(
64
+ backend: Any,
65
+ project: Any,
66
+ dataset: str,
67
+ *,
68
+ prepend_bos: bool,
69
+ sae_ref: Any = None,
70
+ compare: Any = None,
71
+ log: Callable[[str], None] = print,
72
+ ) -> list[Check]:
73
+ """Run the methods on ``backend`` with ``dataset`` (a path in ``project``) and check them.
74
+ ``compare`` is a second backend of the same model on another device, if any."""
75
+ import numpy as np
76
+ import pyarrow.parquet as pq
77
+
78
+ from logogram.compare import compare_summaries
79
+ from logogram.datasets import file_sha256
80
+ from logogram.results import largest_change
81
+ from logogram.runner import run_spec
82
+ from logogram.spec import Spec
83
+ from logogram.verify import verification_spec
84
+
85
+ info = backend.info
86
+ tol = TOLERANCE[info.dtype]
87
+ checks: list[Check] = []
88
+ sha = file_sha256(project.root / dataset)
89
+
90
+ def check(name: str, passed: bool | None, detail: str, exact: bool = True) -> None:
91
+ checks.append(Check(name, passed, detail, exact))
92
+ mark = {True: "PASS", False: "FAIL", None: "SKIP"}[passed]
93
+ log(f"{mark} {name}: {detail}")
94
+
95
+ def spec(experiment: dict, scope: dict, model: Any = None, **extra: Any) -> Spec:
96
+ m = (model or backend).info
97
+ return Spec.model_validate(
98
+ {
99
+ "name": "validation",
100
+ "model": {
101
+ "id": m.id,
102
+ "revision": m.revision,
103
+ "dtype": m.dtype,
104
+ "device": m.device,
105
+ "process_weights": m.process_weights,
106
+ },
107
+ "dataset": {"path": dataset, "sha256": sha},
108
+ "tokenization": {"prepend_bos": prepend_bos},
109
+ "experiment": experiment,
110
+ "scope": scope,
111
+ "statistics": {"bootstrap": 200, "ci": 0.95, "seed": 0},
112
+ "execution": {"batch_size": 32},
113
+ **extra,
114
+ }
115
+ )
116
+
117
+ def run(label: str, s: Spec, model: Any = None, sae: Any = None) -> Ran | None:
118
+ started = time.perf_counter()
119
+ outcome = run_spec(s, project, backend=model or backend, sae=sae)
120
+ seconds = time.perf_counter() - started
121
+ if outcome.status != "finished":
122
+ check(f"{label} runs", False, str(outcome.manifest.get("error")))
123
+ return None
124
+ log(f" {label}: {seconds:.1f} s")
125
+ return Ran(outcome.summary or {}, outcome.folder, s)
126
+
127
+ def effect(outcome: Any, label: str) -> float | None:
128
+ site = next((s for s in outcome.summary["sites"] if s["label"] == label), None)
129
+ return None if site is None else site["effect"]["mean"]
130
+
131
+ patch = {"kind": "activation_patching", "direction": "clean_to_corrupt"}
132
+ estimate = {"kind": "attribution_patching", "direction": "clean_to_corrupt"}
133
+ every_head = {"kind": "heads", "position": {"kind": "all"}}
134
+ last = {"kind": "last"}
135
+ n_layers, n_heads = info.n_layers, info.n_heads
136
+ checks_info = info.extra.get("checks") or {}
137
+ check(
138
+ "The model reproduces itself when it loads",
139
+ True,
140
+ f"error {checks_info.get('function', 0):.1e}, {info.extra.get('block_structure')} layers",
141
+ )
142
+
143
+ heads = run("Patching every head", spec(patch, every_head))
144
+ resid = run(
145
+ "Patching the residual stream at each named position",
146
+ spec(patch, {"kind": "layer_position", "site": "resid_pre", "positions": "labels"}),
147
+ )
148
+ if resid is not None:
149
+ at_s2, at_end = effect(resid, "L0 resid pre @ S2"), effect(resid, "L0 resid pre @ end")
150
+ if at_s2 is None:
151
+ check("Patching the token that differs restores everything", None, "no S2 label")
152
+ else:
153
+ check(
154
+ "Patching the token that differs restores everything",
155
+ abs(at_s2 - 1) <= tol and abs(at_end or 0) <= tol,
156
+ f"effect {at_s2:.4f} at S2 and {at_end or 0:.4f} at the end, in layer 0",
157
+ )
158
+ final = run(
159
+ "Patching the final residual stream",
160
+ spec(
161
+ patch,
162
+ {
163
+ "kind": "sites",
164
+ "sites": [{"kind": "resid_post", "layer": n_layers - 1, "position": last}],
165
+ },
166
+ ),
167
+ )
168
+ if final is not None:
169
+ value = final.summary["sites"][0]["effect"]["mean"]
170
+ check(
171
+ "Patching the final residual stream restores everything",
172
+ abs(value - 1) <= tol,
173
+ f"effect {value:.4f}",
174
+ )
175
+ if heads is not None:
176
+ again = run("The same head sweep again", heads.spec)
177
+ if again is not None:
178
+ change = largest_change(
179
+ pq.read_table(heads.folder / "results.parquet"),
180
+ pq.read_table(again.folder / "results.parquet"),
181
+ )
182
+ check(
183
+ "The same spec twice gives bit-identical results",
184
+ change == 0.0,
185
+ "identical" if change == 0.0 else f"largest change {change}",
186
+ )
187
+
188
+ components = {
189
+ "kind": "layer_components",
190
+ "components": ["attn_out", "mlp_out"],
191
+ "position": {"kind": "all"},
192
+ }
193
+ for name, baseline in (
194
+ ("Zero ablation", {"kind": "zero"}),
195
+ ("Mean ablation", {"kind": "mean", "reference": "corrupt"}),
196
+ ("Resample ablation", {"kind": "resample", "pool": "corrupt", "donors": 4, "seed": 0}),
197
+ ):
198
+ if run(name, spec({"kind": "ablation", "baseline": baseline}, components)) is not None:
199
+ check(f"{name} runs", True, "finished")
200
+
201
+ direct = run(
202
+ "Direct logit attribution",
203
+ spec(
204
+ {"kind": "direct_logit_attribution", "prompts": "clean"},
205
+ {"kind": "layer_components", "components": ["attn_out", "mlp_out"], "position": last},
206
+ ),
207
+ )
208
+ best_heads: list[tuple[int, int]] = []
209
+ if direct is not None:
210
+ split = direct.summary["direct"]
211
+ measured = direct.summary["baseline"]["clean"]["logit_diff"]["mean"]
212
+ parts = split["embeddings"] + split["attention"] + split["mlp"] + split["biases"]
213
+ check(
214
+ "Direct effects add up to the logit difference",
215
+ abs(split["logit_diff"] - measured) <= tol * max(1.0, abs(measured))
216
+ and abs(parts - split["logit_diff"]) <= 1e-6 * max(1.0, abs(measured)),
217
+ f"{split['logit_diff']:.4f} split, {measured:.4f} measured",
218
+ )
219
+ per_head = run(
220
+ "Direct effects of every head",
221
+ spec(
222
+ {"kind": "direct_logit_attribution", "prompts": "clean"},
223
+ {"kind": "heads", "position": last},
224
+ ),
225
+ )
226
+ if per_head is not None:
227
+ later = [s for s in per_head.summary["sites"] if s["layer"] >= n_layers // 2]
228
+ later.sort(key=lambda s: -(s["effect"]["mean"] or 0))
229
+ best_heads = [(s["layer"], s["head"]) for s in later[:3]]
230
+
231
+ estimated = run("Attribution patching of every head", spec(estimate, every_head))
232
+ if estimated is not None and heads is not None:
233
+ cmp = compare_summaries(heads.summary, estimated.summary, top_k=10)
234
+ check(
235
+ "Attribution patching ranks heads like patching",
236
+ (cmp["spearman"] or 0) >= 0.8,
237
+ f"rank correlation {cmp['spearman']:.3f}, top 10 overlap {cmp['top_overlap']}",
238
+ exact=False,
239
+ )
240
+ verified = run(
241
+ "Verifying the top 10 estimates by patching",
242
+ verification_spec(estimated.spec, estimated.summary, 10),
243
+ )
244
+ if verified is not None:
245
+ cmp = compare_summaries(estimated.summary, verified.summary, top_k=5)
246
+ check(
247
+ "Patching reverses none of the top estimates",
248
+ cmp["n_sign_changes"] == 0,
249
+ f"{cmp['n_sign_changes']} reversed, rank correlation {cmp['spearman']:.3f}",
250
+ exact=False,
251
+ )
252
+
253
+ steering = {
254
+ "kind": "steering",
255
+ "apply_to": "clean",
256
+ "coefficients": [-1, 1],
257
+ "train_fraction": 0.5,
258
+ "seed": 0,
259
+ "control": True,
260
+ }
261
+ if run(
262
+ "Steering",
263
+ spec(steering, {"kind": "layer_components", "components": ["resid_pre"], "position": last}),
264
+ ):
265
+ check("Steering runs", True, "finished")
266
+
267
+ last_heads = [
268
+ {"kind": "head", "layer": n_layers - 1, "head": h, "position": {"kind": "all"}}
269
+ for h in range(n_heads)
270
+ ]
271
+ to_logits = {
272
+ "kind": "path_patching",
273
+ "direction": "clean_to_corrupt",
274
+ "receivers": [{"kind": "logits"}],
275
+ "freeze_mlps": False,
276
+ }
277
+ path = run(
278
+ "Path patching from the last layer to the logits",
279
+ spec(to_logits, {"kind": "sites", "sites": last_heads}),
280
+ )
281
+ patched = run(
282
+ "Patching the last layer's heads", spec(patch, {"kind": "sites", "sites": last_heads})
283
+ )
284
+ if path is not None and patched is not None:
285
+ a = np.asarray(pq.read_table(path.folder / "results.parquet").column("patched_logit_diff"))
286
+ b = np.asarray(
287
+ pq.read_table(patched.folder / "results.parquet").column("patched_logit_diff")
288
+ )
289
+ check(
290
+ "From the last layer, a path to the logits is the whole effect",
291
+ float(np.abs(a - b).max()) <= tol,
292
+ f"largest difference {float(np.abs(a - b).max()):.2e}",
293
+ )
294
+ if best_heads:
295
+ receivers = [
296
+ {"kind": "head", "layer": layer, "head": h, "input": "q"} for layer, h in best_heads
297
+ ]
298
+ to_queries = {
299
+ "kind": "path_patching",
300
+ "direction": "clean_to_corrupt",
301
+ "receivers": receivers,
302
+ "freeze_mlps": False,
303
+ }
304
+ if run("Path patching into the strongest heads' queries", spec(to_queries, every_head)):
305
+ check("Path patching into heads runs", True, f"receivers {best_heads}")
306
+
307
+ if info.extra.get("prediction_method"):
308
+ settings = {
309
+ "method": "final_norm_logit_lens",
310
+ "prompt_index": 0,
311
+ "which": "clean",
312
+ "position": last,
313
+ "top_k": 3,
314
+ }
315
+ lens = run(
316
+ "Layer predictions",
317
+ spec(patch, {"kind": "sites", "sites": last_heads[:1]}, predictions=settings),
318
+ )
319
+ if lens is not None:
320
+ report = json.loads((lens.folder / "predictions.json").read_text(encoding="utf-8"))
321
+ mine = report["layers"][-1]["answer_prob"]
322
+ model = lens.summary["per_prompt"]["clean_answer_prob"][0]
323
+ check(
324
+ "The logit lens at the last layer is the model's own prediction",
325
+ abs(mine - model) <= tol * max(model, 1e-6) + 1e-7,
326
+ f"answer probability {mine:.6f} lens, {model:.6f} model",
327
+ )
328
+
329
+ if sae_ref is not None:
330
+ from logogram.analysis import sae_fit_report
331
+ from logogram.datasets import load_dataset
332
+ from logogram.sae import load_sae
333
+
334
+ sae = load_sae(sae_ref, info.device)
335
+ fit = sae_fit_report(
336
+ backend,
337
+ sae,
338
+ load_dataset(project.root / dataset),
339
+ prepend_bos=prepend_bos,
340
+ batch_size=32,
341
+ )
342
+ check(
343
+ "The SAE fits these activations",
344
+ fit["variance_explained"] > 0.5,
345
+ f"explains {fit['variance_explained']:.1%} of the variance; check weight processing if not",
346
+ exact=False,
347
+ )
348
+ ref = {"repo": sae.repo, "path": sae.path, "revision": sae.revision}
349
+ features = run(
350
+ "Attribution patching of every SAE feature",
351
+ spec(
352
+ estimate,
353
+ {"kind": "features", "position": {"kind": "label", "label": "end"}, "top": 20},
354
+ sae=ref,
355
+ ),
356
+ sae=sae,
357
+ )
358
+ if features is not None:
359
+ verified = run(
360
+ "Patching the top 10 features",
361
+ verification_spec(features.spec, features.summary, 10),
362
+ sae=sae,
363
+ )
364
+ if verified is not None:
365
+ cmp = compare_summaries(features.summary, verified.summary, top_k=5)
366
+ check(
367
+ "Feature estimates match patching",
368
+ cmp["n_sign_changes"] == 0 and (cmp["spearman"] or 0) >= 0.8,
369
+ f"rank correlation {cmp['spearman']:.3f}, {cmp['n_sign_changes']} reversed",
370
+ exact=False,
371
+ )
372
+
373
+ if compare is not None and heads is not None:
374
+ other = run(
375
+ f"Patching every head on {compare.info.device}", spec(patch, every_head, model=compare)
376
+ )
377
+ if other is not None:
378
+ x = np.asarray(pq.read_table(heads.folder / "results.parquet").column("effect"))
379
+ y = np.asarray(pq.read_table(other.folder / "results.parquet").column("effect"))
380
+ cmp = compare_summaries(heads.summary, other.summary, top_k=10)
381
+ agree = TOLERANCE[compare.info.dtype] if compare.info.dtype != "float32" else tol
382
+ check(
383
+ f"{info.device} and {compare.info.device} agree",
384
+ float(np.abs(x - y).max()) <= agree and cmp["top_overlap"] == 10,
385
+ f"largest difference in a per-prompt effect {float(np.abs(x - y).max()):.2e}, "
386
+ f"top 10 overlap {cmp['top_overlap']}",
387
+ )
388
+ return checks
389
+
390
+
391
+ def main() -> int:
392
+ parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
393
+ parser.add_argument("--model", default="openai-community/gpt2")
394
+ parser.add_argument("--revision", default=None)
395
+ parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda", "mps"])
396
+ parser.add_argument("--dtype", default="float32", choices=sorted(TOLERANCE))
397
+ parser.add_argument(
398
+ "--no-processing", action="store_true", help="Load without weight processing."
399
+ )
400
+ parser.add_argument("--compare-device", choices=["cpu", "cuda", "mps"], default=None)
401
+ parser.add_argument("--sae", nargs=2, metavar=("REPO", "PATH"), default=None)
402
+ args = parser.parse_args()
403
+
404
+ from logogram.backends.transformer_lens import load_model
405
+ from logogram.datasets import load_dataset, write_dataset
406
+ from logogram.ioi import generate_ioi
407
+ from logogram.project import Project, example_source
408
+ from logogram.prompts import prepare_prompts
409
+ from logogram.spec import SAERef
410
+
411
+ started = time.perf_counter()
412
+
413
+ def load(device: str) -> Any:
414
+ print(f"Loading {args.model} on {device} ({args.dtype})…", flush=True)
415
+ return load_model(
416
+ args.model,
417
+ revision=args.revision,
418
+ dtype=args.dtype,
419
+ device=device,
420
+ process_weights=not args.no_processing,
421
+ )
422
+
423
+ backend = load(args.device)
424
+ compare = load(args.compare_device) if args.compare_device else None
425
+ folder = Path(tempfile.mkdtemp(prefix="logogram-validation-"))
426
+ try:
427
+ shutil.copytree(example_source(), folder / "project")
428
+ project = Project.open(folder / "project")
429
+ bos = bool(backend.info.extra.get("bos"))
430
+ dataset = "datasets/ioi.jsonl"
431
+ try:
432
+ prepare_prompts(backend, load_dataset(project.root / dataset), bos)
433
+ except ValueError as exc:
434
+ # The example's names were chosen for GPT-2; make prompts this tokenizer can use.
435
+ print(f"The example's prompts don't fit this tokenizer ({exc}); generating new ones.")
436
+ records = generate_ioi(
437
+ 64, seed=0, single_token=lambda w: backend.single_token_id(w) is not None
438
+ )
439
+ dataset = "datasets/ioi-validation.jsonl"
440
+ write_dataset(project.root / dataset, records)
441
+ sae_ref = SAERef(repo=args.sae[0], path=args.sae[1]) if args.sae else None
442
+ checks = validate(
443
+ backend, project, dataset, prepend_bos=bos, sae_ref=sae_ref, compare=compare
444
+ )
445
+ finally:
446
+ shutil.rmtree(folder, ignore_errors=True)
447
+ failed = [c for c in checks if c.passed is False]
448
+ print(
449
+ f"\n{len(checks) - len(failed)} of {len(checks)} checks passed "
450
+ f"in {time.perf_counter() - started:.0f} s on {backend.info.device_name}."
451
+ )
452
+ for c in failed:
453
+ kind = "an identity" if c.exact else "agreement expected of a trained model"
454
+ print(f"FAILED ({kind}): {c.name}: {c.detail}")
455
+ return 1 if failed else 0
456
+
457
+
458
+ if __name__ == "__main__":
459
+ sys.exit(main())