circuitkit 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 (368) hide show
  1. circuitkit/__init__.py +128 -0
  2. circuitkit/__main__.py +9 -0
  3. circuitkit/analysis/__init__.py +19 -0
  4. circuitkit/analysis/cross_method_jaccard.py +116 -0
  5. circuitkit/analysis/metrics.py +54 -0
  6. circuitkit/analysis/scores.py +44 -0
  7. circuitkit/api.py +2682 -0
  8. circuitkit/applications/__init__.py +70 -0
  9. circuitkit/applications/arch_registry.py +315 -0
  10. circuitkit/applications/arch_utils.py +302 -0
  11. circuitkit/applications/common_utils/__init__.py +15 -0
  12. circuitkit/applications/common_utils/_covariance.py +223 -0
  13. circuitkit/applications/common_utils/_device.py +33 -0
  14. circuitkit/applications/common_utils/_metrics.py +429 -0
  15. circuitkit/applications/common_utils/_tokenization.py +510 -0
  16. circuitkit/applications/common_utils/benchmark_analysis.py +400 -0
  17. circuitkit/applications/common_utils/cure_clue.py +338 -0
  18. circuitkit/applications/common_utils/hallucination_detection.py +497 -0
  19. circuitkit/applications/common_utils/linear_probe.py +294 -0
  20. circuitkit/applications/editing/__init__.py +30 -0
  21. circuitkit/applications/editing/cake.py +237 -0
  22. circuitkit/applications/editing/circuit_guided_editing.py +487 -0
  23. circuitkit/applications/editing/fine_tune_editing.py +327 -0
  24. circuitkit/applications/editing/knowledge_editing.py +598 -0
  25. circuitkit/applications/editing/knowledge_editing_enhanced.py +863 -0
  26. circuitkit/applications/editing/mcircke.py +263 -0
  27. circuitkit/applications/editing/memit_wrapper.py +947 -0
  28. circuitkit/applications/editing/rome_wrapper.py +770 -0
  29. circuitkit/applications/finetuning/__init__.py +17 -0
  30. circuitkit/applications/finetuning/benchmark_peft.py +451 -0
  31. circuitkit/applications/finetuning/circuit_tuning.py +377 -0
  32. circuitkit/applications/finetuning/healing_metrics.py +304 -0
  33. circuitkit/applications/finetuning/peft_methods.py +563 -0
  34. circuitkit/applications/finetuning/soft_healing.py +717 -0
  35. circuitkit/applications/pruning/__init__.py +13 -0
  36. circuitkit/applications/pruning/eval_utils.py +347 -0
  37. circuitkit/applications/pruning/examples/__init__.py +0 -0
  38. circuitkit/applications/pruning/examples/prune.py +787 -0
  39. circuitkit/applications/pruning/examples/prune_llama.py +661 -0
  40. circuitkit/applications/pruning/examples/prune_qwen.py +591 -0
  41. circuitkit/applications/pruning/finetune_utils.py +477 -0
  42. circuitkit/applications/pruning/importance.py +97 -0
  43. circuitkit/applications/pruning/neuron_pruner.py +36 -0
  44. circuitkit/applications/pruning/node_pruner.py +186 -0
  45. circuitkit/applications/pruning/pruner.py +529 -0
  46. circuitkit/applications/pruning/score_extractor.py +541 -0
  47. circuitkit/applications/pruning/selectors/__init__.py +0 -0
  48. circuitkit/applications/pruning/selectors/multi_granular_selector.py +119 -0
  49. circuitkit/applications/pruning/selectors/taylor_selector.py +112 -0
  50. circuitkit/applications/pruning/weight_pruner.py +652 -0
  51. circuitkit/applications/quantization/__init__.py +19 -0
  52. circuitkit/applications/quantization/examples/__init__.py +0 -0
  53. circuitkit/applications/quantization/examples/quantize_llama.py +745 -0
  54. circuitkit/applications/quantization/examples/quantize_qwen.py +726 -0
  55. circuitkit/applications/quantization/llmcompressor_quantize.py +388 -0
  56. circuitkit/applications/quantization/quant_utils.py +825 -0
  57. circuitkit/applications/quantization/score_extractor.py +465 -0
  58. circuitkit/applications/quantization/selectors/__init__.py +0 -0
  59. circuitkit/applications/quantization/selectors/awq_selector.py +109 -0
  60. circuitkit/applications/quantization/selectors/tacq_selector.py +145 -0
  61. circuitkit/applications/selective_finetuning/__init__.py +0 -0
  62. circuitkit/applications/selective_finetuning/examples/__init__.py +0 -0
  63. circuitkit/applications/selective_finetuning/examples/finetune_llama.py +668 -0
  64. circuitkit/applications/selective_finetuning/examples/finetune_qwen.py +654 -0
  65. circuitkit/applications/selective_finetuning/finetune_utils.py +643 -0
  66. circuitkit/applications/selective_finetuning/score_loader.py +585 -0
  67. circuitkit/applications/selective_finetuning/selector.py +616 -0
  68. circuitkit/applications/steering/__init__.py +31 -0
  69. circuitkit/applications/steering/steering.py +791 -0
  70. circuitkit/applications/steering/steering_enhanced.py +556 -0
  71. circuitkit/applications/steering/weight_steering.py +407 -0
  72. circuitkit/artifacts/__init__.py +25 -0
  73. circuitkit/artifacts/circuit_artifact.py +559 -0
  74. circuitkit/artifacts/converters.py +405 -0
  75. circuitkit/artifacts/scores.py +195 -0
  76. circuitkit/backends/__init__.py +96 -0
  77. circuitkit/backends/acdc/__init__.py +0 -0
  78. circuitkit/backends/acdc/artifact_export.py +140 -0
  79. circuitkit/backends/acdc/data.py +183 -0
  80. circuitkit/backends/acdc/model_utils/__init__.py +0 -0
  81. circuitkit/backends/acdc/model_utils/micro_model_utils.py +143 -0
  82. circuitkit/backends/acdc/model_utils/transformer_lens_utils.py +232 -0
  83. circuitkit/backends/acdc/prune.py +123 -0
  84. circuitkit/backends/acdc/prune_algos/ACDC.py +136 -0
  85. circuitkit/backends/acdc/prune_algos/__init__.py +0 -0
  86. circuitkit/backends/acdc/prune_algos/mask_gradient.py +129 -0
  87. circuitkit/backends/acdc/prune_algos/prune_algos.py +33 -0
  88. circuitkit/backends/acdc/tasks/__init__.py +3 -0
  89. circuitkit/backends/acdc/tasks/docstring_prompts.py +837 -0
  90. circuitkit/backends/acdc/tasks/docstring_utils.py +88 -0
  91. circuitkit/backends/acdc/tasks/induction_utils.py +122 -0
  92. circuitkit/backends/acdc/tasks/ioi_dataset.py +156 -0
  93. circuitkit/backends/acdc/tasks/ioi_utils.py +96 -0
  94. circuitkit/backends/acdc/types.py +248 -0
  95. circuitkit/backends/acdc/utils/__init__.py +0 -0
  96. circuitkit/backends/acdc/utils/ablation_activations.py +160 -0
  97. circuitkit/backends/acdc/utils/custom_tqdm.py +14 -0
  98. circuitkit/backends/acdc/utils/graph_utils.py +497 -0
  99. circuitkit/backends/acdc/utils/misc.py +68 -0
  100. circuitkit/backends/acdc/utils/patch_wrapper.py +106 -0
  101. circuitkit/backends/acdc/utils/patchable_model.py +143 -0
  102. circuitkit/backends/acdc/utils/task_utils.py +31 -0
  103. circuitkit/backends/acdc/utils/tensor_ops.py +118 -0
  104. circuitkit/backends/acdc/visualize.py +253 -0
  105. circuitkit/backends/cdt/__init__.py +30 -0
  106. circuitkit/backends/cdt/adapter.py +245 -0
  107. circuitkit/backends/cdt/propagation.py +357 -0
  108. circuitkit/backends/cdt/pyfunctions/__init__.py +14 -0
  109. circuitkit/backends/cdt/pyfunctions/cdt_ablations.py +215 -0
  110. circuitkit/backends/cdt/pyfunctions/cdt_basic.py +235 -0
  111. circuitkit/backends/cdt/pyfunctions/cdt_core.py +418 -0
  112. circuitkit/backends/cdt/pyfunctions/cdt_from_source_nodes.py +237 -0
  113. circuitkit/backends/cdt/pyfunctions/cdt_source_to_target.py +685 -0
  114. circuitkit/backends/cdt/pyfunctions/faithfulness_ablations.py +251 -0
  115. circuitkit/backends/cdt/pyfunctions/general.py +314 -0
  116. circuitkit/backends/cdt/pyfunctions/ioi_dataset.py +958 -0
  117. circuitkit/backends/cdt/pyfunctions/local_importance.py +809 -0
  118. circuitkit/backends/cdt/pyfunctions/pathology.py +460 -0
  119. circuitkit/backends/cdt/pyfunctions/toy_model.py +190 -0
  120. circuitkit/backends/cdt/pyfunctions/wrappers.py +159 -0
  121. circuitkit/backends/eap/__init__.py +2 -0
  122. circuitkit/backends/eap/artifact_export.py +137 -0
  123. circuitkit/backends/eap/attribute.py +784 -0
  124. circuitkit/backends/eap/attribute_node.py +1795 -0
  125. circuitkit/backends/eap/circuit_kit_adapter.py +121 -0
  126. circuitkit/backends/eap/eap_utils.py +582 -0
  127. circuitkit/backends/eap/evaluate.py +762 -0
  128. circuitkit/backends/eap/graph.py +1569 -0
  129. circuitkit/backends/eap/metrics.py +793 -0
  130. circuitkit/backends/eap/py.typed +0 -0
  131. circuitkit/backends/eap/visualization.py +101 -0
  132. circuitkit/backends/ibcircuit/__init__.py +0 -0
  133. circuitkit/backends/ibcircuit/artifact_export.py +127 -0
  134. circuitkit/backends/ibcircuit/ib_noise.py +216 -0
  135. circuitkit/backends/ibcircuit/ib_utils.py +194 -0
  136. circuitkit/backends/ibcircuit/model_wrapper.py +537 -0
  137. circuitkit/backends/ibcircuit/trainer.py +603 -0
  138. circuitkit/benchmarks/__init__.py +47 -0
  139. circuitkit/benchmarks/baselines/__init__.py +20 -0
  140. circuitkit/benchmarks/baselines/gptq.py +200 -0
  141. circuitkit/benchmarks/baselines/magnitude.py +204 -0
  142. circuitkit/benchmarks/baselines/random.py +142 -0
  143. circuitkit/benchmarks/baselines/sparsegpt.py +236 -0
  144. circuitkit/benchmarks/baselines/wanda.py +263 -0
  145. circuitkit/benchmarks/benchmark.py +764 -0
  146. circuitkit/benchmarks/reporting.py +639 -0
  147. circuitkit/circuit.py +390 -0
  148. circuitkit/cli/__init__.py +1 -0
  149. circuitkit/cli/config.py +74 -0
  150. circuitkit/cli/debug.py +279 -0
  151. circuitkit/cli/main.py +2208 -0
  152. circuitkit/cli/utils.py +351 -0
  153. circuitkit/corruption/__init__.py +61 -0
  154. circuitkit/corruption/base.py +132 -0
  155. circuitkit/corruption/color_swap.py +170 -0
  156. circuitkit/corruption/distractor.py +319 -0
  157. circuitkit/corruption/distractor_variation.py +313 -0
  158. circuitkit/corruption/effectiveness.py +297 -0
  159. circuitkit/corruption/entity_swap.py +288 -0
  160. circuitkit/corruption/negation.py +364 -0
  161. circuitkit/corruption/paraphrase.py +390 -0
  162. circuitkit/corruption/pipeline.py +333 -0
  163. circuitkit/corruption/position_shift.py +106 -0
  164. circuitkit/corruption/role_swap.py +381 -0
  165. circuitkit/corruption/token_swap.py +257 -0
  166. circuitkit/corruption/validators.py +570 -0
  167. circuitkit/corruption/voice_swap.py +393 -0
  168. circuitkit/data/__init__.py +8 -0
  169. circuitkit/data/adapters/__init__.py +20 -0
  170. circuitkit/data/adapters/base.py +123 -0
  171. circuitkit/data/adapters/code.py +106 -0
  172. circuitkit/data/adapters/conversational.py +167 -0
  173. circuitkit/data/adapters/forget_retain.py +152 -0
  174. circuitkit/data/adapters/instruction.py +125 -0
  175. circuitkit/data/adapters/math.py +133 -0
  176. circuitkit/data/adapters/mcq.py +194 -0
  177. circuitkit/data/adapters/pairwise.py +182 -0
  178. circuitkit/data/adapters/safety_prompt.py +258 -0
  179. circuitkit/data/auto_detect.py +187 -0
  180. circuitkit/data/clean_only.py +124 -0
  181. circuitkit/data/corruption/__init__.py +38 -0
  182. circuitkit/data/corruption/base.py +239 -0
  183. circuitkit/data/corruption/benign_rewrite.py +127 -0
  184. circuitkit/data/corruption/code_syntax_corrupt.py +98 -0
  185. circuitkit/data/corruption/entity_swap.py +113 -0
  186. circuitkit/data/corruption/final_answer_swap.py +231 -0
  187. circuitkit/data/corruption/instruction_swap.py +149 -0
  188. circuitkit/data/corruption/llm_counterfactual.py +161 -0
  189. circuitkit/data/corruption/logical_negation.py +96 -0
  190. circuitkit/data/corruption/math_step_corrupt.py +88 -0
  191. circuitkit/data/corruption/mcq_choice_swap.py +106 -0
  192. circuitkit/data/corruption/operand_swap.py +103 -0
  193. circuitkit/data/corruption/profession_swap.py +125 -0
  194. circuitkit/data/corruption/resample.py +75 -0
  195. circuitkit/data/corruption/template.py +195 -0
  196. circuitkit/data/corruption/template_utils.py +328 -0
  197. circuitkit/data/corruption/token_swap.py +90 -0
  198. circuitkit/data/dataset_schema.py +169 -0
  199. circuitkit/data/eap_dataset.py +98 -0
  200. circuitkit/data/invariance_groups/__init__.py +33 -0
  201. circuitkit/data/invariance_groups/builder.py +323 -0
  202. circuitkit/data/invariance_groups/schema.py +274 -0
  203. circuitkit/data/normalized.py +259 -0
  204. circuitkit/data/normalized_task.py +594 -0
  205. circuitkit/data/task_data/__init__.py +11 -0
  206. circuitkit/data/task_data/core/TLACDCCorrespondence.py +263 -0
  207. circuitkit/data/task_data/core/TLACDCEdge.py +113 -0
  208. circuitkit/data/task_data/core/TLACDCExperiment.py +1052 -0
  209. circuitkit/data/task_data/core/TLACDCInterpNode.py +96 -0
  210. circuitkit/data/task_data/core/__init__.py +12 -0
  211. circuitkit/data/task_data/core/acdc_utils.py +614 -0
  212. circuitkit/data/task_data/generation/__init__.py +11 -0
  213. circuitkit/data/task_data/generation/cache.py +267 -0
  214. circuitkit/data/task_data/generation/manager.py +562 -0
  215. circuitkit/data/task_data/generation/utils.py +323 -0
  216. circuitkit/data/task_data/storage/__init__.py +18 -0
  217. circuitkit/data/task_data/storage/greaterthan/greaterthan_32_ffd33106.json +23 -0
  218. circuitkit/data/task_data/storage/ioi/ioi_16_8c879ddb.json +43 -0
  219. circuitkit/data/task_data/storage/ioi/ioi_32_a432ca4a.json +43 -0
  220. circuitkit/data/task_data/storage/ioi/ioi_500_1f7e7324.json +43 -0
  221. circuitkit/data/task_data/storage/ioi/ioi_64_3bba747e.json +43 -0
  222. circuitkit/data/task_data/storage/ioi/ioi_64_f4a164db.json +43 -0
  223. circuitkit/data/task_data/storage/ioi/ioi_8_e87df42e.json +43 -0
  224. circuitkit/data/task_data/tasks/__init__.py +10 -0
  225. circuitkit/data/task_data/tasks/binary_align/generate_binary_align.py +1167 -0
  226. circuitkit/data/task_data/tasks/binary_align/jailbreak_binary.csv +335 -0
  227. circuitkit/data/task_data/tasks/binary_align/safe_binary.csv +335 -0
  228. circuitkit/data/task_data/tasks/capital_country/__init__.py +5 -0
  229. circuitkit/data/task_data/tasks/capital_country/utils.py +395 -0
  230. circuitkit/data/task_data/tasks/docstring/__init__.py +5 -0
  231. circuitkit/data/task_data/tasks/docstring/prompts.py +1175 -0
  232. circuitkit/data/task_data/tasks/docstring/utils.py +282 -0
  233. circuitkit/data/task_data/tasks/double_io/__init__.py +0 -0
  234. circuitkit/data/task_data/tasks/double_io/double_io_dataset.py +485 -0
  235. circuitkit/data/task_data/tasks/gender_bias/__init__.py +5 -0
  236. circuitkit/data/task_data/tasks/gender_bias/utils.py +396 -0
  237. circuitkit/data/task_data/tasks/gender_bias/utils2.py +143 -0
  238. circuitkit/data/task_data/tasks/greaterthan/__init__.py +5 -0
  239. circuitkit/data/task_data/tasks/greaterthan/utils.py +534 -0
  240. circuitkit/data/task_data/tasks/hypernymy/__init__.py +5 -0
  241. circuitkit/data/task_data/tasks/hypernymy/utils.py +326 -0
  242. circuitkit/data/task_data/tasks/induction/__init__.py +5 -0
  243. circuitkit/data/task_data/tasks/induction/utils.py +222 -0
  244. circuitkit/data/task_data/tasks/ioi/__init__.py +8 -0
  245. circuitkit/data/task_data/tasks/ioi/ioi_dataset.py +962 -0
  246. circuitkit/data/task_data/tasks/ioi/utils.py +656 -0
  247. circuitkit/data/task_data/tasks/sva/__init__.py +5 -0
  248. circuitkit/data/task_data/tasks/sva/utils.py +132 -0
  249. circuitkit/data/task_data/tasks/wmdp/wmdp_utils.py +296 -0
  250. circuitkit/data/template.py +392 -0
  251. circuitkit/data/wikitext_calibration.py +164 -0
  252. circuitkit/data/worthiness.py +746 -0
  253. circuitkit/evaluation/__init__.py +83 -0
  254. circuitkit/evaluation/checkpoint_benchmark.py +857 -0
  255. circuitkit/evaluation/evaluate.py +929 -0
  256. circuitkit/evaluation/full.py +556 -0
  257. circuitkit/evaluation/hf_checkpoint.py +1219 -0
  258. circuitkit/evaluation/intervention_faithfulness.py +198 -0
  259. circuitkit/evaluation/lm_eval_simple.py +223 -0
  260. circuitkit/evaluation/lm_harness.py +681 -0
  261. circuitkit/evaluation/master_grid.py +307 -0
  262. circuitkit/evaluation/mmlu_eval.py +208 -0
  263. circuitkit/evaluation/pillars/__init__.py +28 -0
  264. circuitkit/evaluation/pillars/ablation.py +404 -0
  265. circuitkit/evaluation/pillars/baselines.py +902 -0
  266. circuitkit/evaluation/pillars/causal_patching.py +371 -0
  267. circuitkit/evaluation/pillars/generalization.py +623 -0
  268. circuitkit/evaluation/pillars/intervention_reliability.py +325 -0
  269. circuitkit/evaluation/pillars/robustness.py +854 -0
  270. circuitkit/evaluation/pillars/stability.py +571 -0
  271. circuitkit/evaluation/report.py +318 -0
  272. circuitkit/evaluation/reports/__init__.py +20 -0
  273. circuitkit/evaluation/reports/aggregator.py +533 -0
  274. circuitkit/evaluation/reports/robustness_report.py +331 -0
  275. circuitkit/evaluation/reports/stability_report.py +298 -0
  276. circuitkit/evaluation/stability_discovery.py +439 -0
  277. circuitkit/evaluation/transfer.py +510 -0
  278. circuitkit/evaluation/transfer_analysis.py +315 -0
  279. circuitkit/evaluation/transfer_visualizer.py +327 -0
  280. circuitkit/evaluation/weight_based_eval.py +227 -0
  281. circuitkit/pipeline.py +1000 -0
  282. circuitkit/quick.py +1157 -0
  283. circuitkit/selection/__init__.py +54 -0
  284. circuitkit/selection/cdt_selector.py +64 -0
  285. circuitkit/selection/eap_gp_selector.py +67 -0
  286. circuitkit/selection/eap_selector.py +83 -0
  287. circuitkit/selection/gptq_selector.py +164 -0
  288. circuitkit/selection/ibcircuit_selector.py +120 -0
  289. circuitkit/selection/magnitude_selector.py +28 -0
  290. circuitkit/selection/random_selector.py +16 -0
  291. circuitkit/selection/relp_selector.py +66 -0
  292. circuitkit/selection/wanda_selector.py +174 -0
  293. circuitkit/tasks/__init__.py +40 -0
  294. circuitkit/tasks/_algorithm_families.py +106 -0
  295. circuitkit/tasks/_chat.py +262 -0
  296. circuitkit/tasks/auto_schema.py +526 -0
  297. circuitkit/tasks/bootstrap.py +93 -0
  298. circuitkit/tasks/builtins/__init__.py +44 -0
  299. circuitkit/tasks/builtins/boolq.py +501 -0
  300. circuitkit/tasks/builtins/capital_country.py +261 -0
  301. circuitkit/tasks/builtins/double_io.py +380 -0
  302. circuitkit/tasks/builtins/gender_bias.py +284 -0
  303. circuitkit/tasks/builtins/glue.py +683 -0
  304. circuitkit/tasks/builtins/greater_than.py +471 -0
  305. circuitkit/tasks/builtins/gsm8k.py +563 -0
  306. circuitkit/tasks/builtins/hypernymy.py +262 -0
  307. circuitkit/tasks/builtins/ifeval.py +116 -0
  308. circuitkit/tasks/builtins/ioi.py +447 -0
  309. circuitkit/tasks/builtins/ioi_acdc.py +323 -0
  310. circuitkit/tasks/builtins/ioi_legacy.py +473 -0
  311. circuitkit/tasks/builtins/mmlu.py +1578 -0
  312. circuitkit/tasks/builtins/sva.py +257 -0
  313. circuitkit/tasks/builtins/truthfulqa.py +520 -0
  314. circuitkit/tasks/builtins/winogrande.py +647 -0
  315. circuitkit/tasks/builtins/winogrande_mc.py +484 -0
  316. circuitkit/tasks/builtins/wmdp.py +1120 -0
  317. circuitkit/tasks/generic.py +1405 -0
  318. circuitkit/tasks/hf_factory.py +480 -0
  319. circuitkit/tasks/inspect.py +95 -0
  320. circuitkit/tasks/registry.py +72 -0
  321. circuitkit/tasks/safety_datasets.py +139 -0
  322. circuitkit/tasks/specs.py +246 -0
  323. circuitkit/tasks/type_specs/__init__.py +30 -0
  324. circuitkit/tasks/type_specs/classification_spec.py +49 -0
  325. circuitkit/tasks/type_specs/generation_spec.py +47 -0
  326. circuitkit/tasks/type_specs/mcq_spec.py +49 -0
  327. circuitkit/tasks/type_specs/qa_spec.py +129 -0
  328. circuitkit/tasks/type_specs/summarization_spec.py +45 -0
  329. circuitkit/tasks/type_specs/translation_spec.py +45 -0
  330. circuitkit/tasks/validator.py +481 -0
  331. circuitkit/tasks/yaml_loader.py +381 -0
  332. circuitkit/tooling/__init__.py +7 -0
  333. circuitkit/tooling/validate_environment.py +61 -0
  334. circuitkit/utils/__init__.py +0 -0
  335. circuitkit/utils/artifacts.py +39 -0
  336. circuitkit/utils/async_processing.py +371 -0
  337. circuitkit/utils/bootstrap.py +194 -0
  338. circuitkit/utils/config.py +330 -0
  339. circuitkit/utils/corruption_validation.py +369 -0
  340. circuitkit/utils/dataset_cache.py +303 -0
  341. circuitkit/utils/debug.py +345 -0
  342. circuitkit/utils/debugging.py +354 -0
  343. circuitkit/utils/device.py +51 -0
  344. circuitkit/utils/distributed.py +531 -0
  345. circuitkit/utils/exceptions.py +355 -0
  346. circuitkit/utils/logging.py +382 -0
  347. circuitkit/utils/memory.py +191 -0
  348. circuitkit/utils/optimization.py +316 -0
  349. circuitkit/utils/profiling.py +414 -0
  350. circuitkit/utils/token_utils.py +159 -0
  351. circuitkit/visualize/__init__.py +79 -0
  352. circuitkit/visualize/comparison.py +499 -0
  353. circuitkit/visualize/d3_template.py +1064 -0
  354. circuitkit/visualize/editor.py +385 -0
  355. circuitkit/visualize/feature_saliency.py +408 -0
  356. circuitkit/visualize/gallery.py +392 -0
  357. circuitkit/visualize/graph_viz.py +884 -0
  358. circuitkit/visualize/jupyter_suite.py +223 -0
  359. circuitkit/visualize/plotter.py +246 -0
  360. circuitkit/visualize/saliency.py +402 -0
  361. circuitkit/visualize/streamlit_app.py +471 -0
  362. circuitkit/visualize/theme.py +331 -0
  363. circuitkit-0.1.0.dist-info/METADATA +192 -0
  364. circuitkit-0.1.0.dist-info/RECORD +368 -0
  365. circuitkit-0.1.0.dist-info/WHEEL +5 -0
  366. circuitkit-0.1.0.dist-info/entry_points.txt +2 -0
  367. circuitkit-0.1.0.dist-info/licenses/LICENSE.md +73 -0
  368. circuitkit-0.1.0.dist-info/top_level.txt +1 -0
circuitkit/api.py ADDED
@@ -0,0 +1,2682 @@
1
+ import logging
2
+ import os
3
+ import re
4
+ import warnings
5
+ from datetime import datetime
6
+ from functools import partial
7
+ from pathlib import Path
8
+ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
9
+
10
+ import torch as t
11
+
12
+ if TYPE_CHECKING:
13
+ from .evaluation.report import FaithfulnessReport
14
+
15
+ # Suppress verbose warnings BEFORE importing heavy libraries
16
+ # This must be done early to catch warnings during imports
17
+ warnings.filterwarnings("ignore", message=".*reduced precision.*")
18
+ warnings.filterwarnings("ignore", message=".*from_pretrained_no_processing.*")
19
+ warnings.filterwarnings("ignore", message=".*pretrained.*model kwarg is not of type.*")
20
+ warnings.filterwarnings("ignore", message=".*Passed an already-initialized model.*")
21
+ warnings.filterwarnings("ignore", message=".*Overwriting default num_fewshot.*")
22
+ warnings.filterwarnings("ignore", message=".*S2 index has been computed.*")
23
+ warnings.filterwarnings("ignore", category=UserWarning)
24
+
25
+ # Suppress verbose loggers
26
+ logging.getLogger("transformers").setLevel(logging.ERROR)
27
+ logging.getLogger("lm_eval").setLevel(logging.ERROR)
28
+ logging.getLogger("accelerate").setLevel(logging.ERROR)
29
+
30
+ from transformer_lens import ( # noqa: E402 - import after intentional pre-import setup
31
+ HookedTransformer,
32
+ )
33
+
34
+ # CircuitScores artifact (Workstream G)
35
+ from .artifacts.scores import ( # noqa: E402 - import after intentional pre-import setup
36
+ CircuitScores,
37
+ )
38
+ from .utils.debug import ( # noqa: E402 - import after intentional pre-import setup
39
+ debug_context,
40
+ debug_function,
41
+ )
42
+ from .utils.exceptions import ( # noqa: E402 - import after intentional pre-import setup
43
+ AlgorithmError,
44
+ handle_errors,
45
+ validate_discovery_algorithm,
46
+ validate_file_exists,
47
+ validate_model_name,
48
+ )
49
+
50
+ # CircuitKit imports
51
+ from circuitkit.utils.device import get_device, empty_cache
52
+ from .utils.logging import ( # noqa: E402 - import after intentional pre-import setup
53
+ ProgressLogger,
54
+ get_logger,
55
+ log_execution_time,
56
+ )
57
+
58
+ logger = get_logger(__name__)
59
+
60
+
61
+ def _fmt_opt_score(x, spec=".4f"):
62
+ """Format an optional score for logging. Pillar scores (patching/ablation)
63
+ are None when the underlying metric is invalid (e.g. inverted denominator);
64
+ formatting None with ``:.4f`` raises ``NoneType.__format__``."""
65
+ return format(x, spec) if x is not None else "invalid"
66
+
67
+
68
+ # Task management imports
69
+ # NOTE: imported lazily inside functions to avoid a circular import between
70
+ # `circuitkit.api` and `circuitkit.tasks.registry` (task builtins pull in
71
+ # helpers re-exported from this module), which breaks when `circuitkit.api`
72
+ # is the first module imported.
73
+ def _get_task(*args, **kwargs):
74
+ from .tasks.registry import get_task as _gt
75
+
76
+ return _gt(*args, **kwargs)
77
+
78
+
79
+ def _register_task(*args, **kwargs):
80
+ from .tasks.registry import register_task as _rt
81
+
82
+ return _rt(*args, **kwargs)
83
+
84
+
85
+ import warnings as _warnings # noqa: E402 - import after intentional pre-import setup
86
+
87
+ from .backends import ( # noqa: E402 - import after intentional pre-import setup
88
+ DEFAULT_ALGORITHM as _DEFAULT_ALGO,
89
+ )
90
+ from .backends import ( # noqa: E402 - import after intentional pre-import setup
91
+ EXPERIMENTAL_ALGORITHMS,
92
+ RESEARCH_ALGORITHMS,
93
+ )
94
+
95
+ # ACDC Backend Imports
96
+ from .backends.acdc.data import ( # noqa: E402 - import after intentional pre-import setup
97
+ load_task_data,
98
+ )
99
+ from .backends.acdc.prune_algos.ACDC import ( # noqa: E402 - import after intentional pre-import setup
100
+ acdc_prune_scores,
101
+ )
102
+ from .backends.acdc.utils.graph_utils import ( # noqa: E402 - import after intentional pre-import setup
103
+ patchable_model,
104
+ )
105
+ from .backends.eap.attribute_node import ( # noqa: E402 - import after intentional pre-import setup
106
+ attribute_node,
107
+ )
108
+ # EAP Backend Imports
109
+ from .backends.eap.graph import ( # noqa: E402 - import after intentional pre-import setup
110
+ AttentionNode,
111
+ Graph,
112
+ MLPNode,
113
+ )
114
+
115
+
116
+ def _log_gpu_mem(label: str, logger):
117
+ """Log GPU memory stats at DEBUG level. No-op if CUDA unavailable."""
118
+ import torch
119
+
120
+ if not torch.cuda.is_available():
121
+ return
122
+ allocated = torch.cuda.memory_allocated() / (1024**3)
123
+ reserved = torch.cuda.memory_reserved() / (1024**3)
124
+ free_reserved = reserved - allocated
125
+ total = torch.cuda.get_device_properties(0).total_memory / (1024**3)
126
+ free_total = total - reserved
127
+ logger.debug(
128
+ f"[GPU-MEM] {label}: "
129
+ f"alloc={allocated:.2f}GB, reserved={reserved:.2f}GB, "
130
+ f"free_in_reserved={free_reserved:.2f}GB, free_total={free_total:.2f}GB"
131
+ )
132
+
133
+
134
+ # EAPDiscoveryDataset is a torch Dataset — it lives in circuitkit.data, not in
135
+ # this front-door facade. Re-exported here for backward compatibility; new code
136
+ # should import it from circuitkit.data.eap_dataset.
137
+ from .data.eap_dataset import EAPDiscoveryDataset # noqa: F401,E402
138
+
139
+
140
+ from collections import defaultdict # noqa: E402 - import after intentional pre-import setup
141
+
142
+ from tqdm import tqdm # noqa: E402 - import after intentional pre-import setup
143
+
144
+ # CircuitKit Core Imports
145
+ from .analysis.scores import ( # noqa: E402 - import after intentional pre-import setup
146
+ calculate_node_scores_from_edges,
147
+ )
148
+ from .applications.pruning.node_pruner import ( # noqa: E402 - import after intentional pre-import setup
149
+ get_nodes_to_prune,
150
+ )
151
+ from .utils.config import ( # noqa: E402 - import after intentional pre-import setup
152
+ DEFAULT_CONFIG,
153
+ load_and_validate_config,
154
+ )
155
+
156
+
157
+ def _correct_token_prob(logits, clean_logits, input_lengths, labels, loss=False, mean=False):
158
+ """Correct-answer token probability at the answer position.
159
+
160
+ A bounded [0, 1] metric suitable for clean-only evaluation (e.g.
161
+ IBCircuit neuron-level on clean-only custom data) where no incorrect token
162
+ is available and logit_diff would collapse to zero.
163
+
164
+ ``labels`` shape: [batch, 2] where ``labels[:, 0]`` = correct token ID.
165
+ The second column is ignored (may be a duplicate or a dummy).
166
+ """
167
+ batch = logits.size(0)
168
+ idx = t.arange(batch, device=logits.device)
169
+ last = (input_lengths.long() - 1).clamp_min(0)
170
+ probs = t.softmax(logits[idx, last], dim=-1)
171
+ correct_ids = labels[:, 0].to(logits.device)
172
+ result = probs[idx, correct_ids]
173
+ if loss:
174
+ result = -result
175
+ if mean:
176
+ result = result.mean()
177
+ return result
178
+
179
+
180
+ def _build_clean_only_ib_eval_dataloader(task_spec, model, num_examples: int, batch_size: int):
181
+ """Build an EAP-format eval DataLoader for clean-only NormalizedTaskSpec.
182
+
183
+ Duplicates the clean prompt as the corrupt side so tokenize_batch_pair
184
+ in evaluate_baseline / evaluate_ibcircuit_neuron_circuit can run without
185
+ real paired data. Mean- and zero-ablation paths never consume corrupt
186
+ activations, so the duplicate is harmless.
187
+
188
+ Labels are shaped [batch, 2] with both columns = correct-answer token ID,
189
+ suitable for _correct_token_prob (which only reads column 0).
190
+
191
+ Returns a torch DataLoader yielding (clean_list, corrupt_list, label_tensor)
192
+ batches — the same EAP-format the evaluators expect.
193
+ """
194
+ import torch
195
+ from torch.utils.data import DataLoader, TensorDataset
196
+
197
+ tokenizer = model.tokenizer
198
+ try:
199
+ ws_probe = tokenizer.encode(" ", add_special_tokens=False)
200
+ ws_token_id = ws_probe[0] if len(ws_probe) == 1 else None
201
+ except Exception:
202
+ ws_token_id = None
203
+
204
+ clean_texts = []
205
+ label_ids = []
206
+
207
+ for r in task_spec.ds.records[:num_examples]:
208
+ # Derive correct-answer token ID using joint encoding (same logic as
209
+ # NormalizedTaskSpec._build_ibcircuit_dataloader).
210
+ precomputed = r.meta.get("_precomputed_labels")
211
+ if precomputed:
212
+ ans_token = precomputed["clean_label_id"]
213
+ else:
214
+ prompt_ids_solo = tokenizer.encode(r.clean_prompt, add_special_tokens=False)
215
+ full_ids = tokenizer.encode(r.clean_prompt + r.clean_answer, add_special_tokens=False)
216
+ boundary_clean = (
217
+ len(full_ids) > len(prompt_ids_solo)
218
+ and full_ids[: len(prompt_ids_solo)] == prompt_ids_solo
219
+ )
220
+ if boundary_clean:
221
+ first_cont = int(full_ids[len(prompt_ids_solo)])
222
+ if (
223
+ ws_token_id is not None
224
+ and first_cont == ws_token_id
225
+ and len(full_ids) > len(prompt_ids_solo) + 1
226
+ ):
227
+ ans_token = int(full_ids[len(prompt_ids_solo) + 1])
228
+ else:
229
+ ans_token = first_cont
230
+ else:
231
+ ans_ids = tokenizer.encode(r.clean_answer, add_special_tokens=False)
232
+ if not ans_ids:
233
+ continue
234
+ ans_token = ans_ids[0]
235
+ if ws_token_id is not None and ans_token == ws_token_id and len(ans_ids) > 1:
236
+ ans_token = ans_ids[1]
237
+
238
+ clean_texts.append(r.clean_prompt)
239
+ label_ids.append(ans_token)
240
+
241
+ if not clean_texts:
242
+ raise RuntimeError(
243
+ "clean-only eval dataloader: no records could be built from the task spec."
244
+ )
245
+
246
+ label_tensor = torch.tensor([[lid, lid] for lid in label_ids], dtype=torch.long)
247
+
248
+ # Batch into (clean_list, corrupt_list, label_chunk) tuples.
249
+ batches = []
250
+ for start in range(0, len(clean_texts), batch_size):
251
+ end = start + batch_size
252
+ chunk_clean = clean_texts[start:end]
253
+ chunk_labels = label_tensor[start:end]
254
+ batches.append((chunk_clean, chunk_clean, chunk_labels))
255
+
256
+ class _TextBatchLoader:
257
+ def __init__(self, batches, padding_side="left"):
258
+ self._batches = batches
259
+ self.pair_padding_side = padding_side
260
+
261
+ def __iter__(self):
262
+ return iter(self._batches)
263
+
264
+ def __len__(self):
265
+ return len(self._batches)
266
+
267
+ side = getattr(task_spec, "pair_padding_side", "left")
268
+ return _TextBatchLoader(batches, padding_side=side)
269
+
270
+
271
+ # Backwards-compatibility re-export: legacy code (e.g. tasks/builtins/ioi_acdc.py)
272
+ # imports `_eap_logit_diff` from this module. The canonical implementation now
273
+ # lives on each TaskSpec. We forward to IOITaskSpec._ioi_logit_diff because the
274
+ # legacy importer was IOI-specific.
275
+ def _eap_logit_diff(*args, **kwargs):
276
+ """Legacy IOI-style logit-difference metric.
277
+
278
+ Deprecated: use ``IOITaskSpec._ioi_logit_diff`` (or the equivalent on
279
+ your TaskSpec) instead.
280
+ """
281
+ from .tasks.builtins.ioi import IOITaskSpec
282
+
283
+ return IOITaskSpec._ioi_logit_diff(*args, **kwargs)
284
+
285
+
286
+ def _eap_kl_divergence(logits, clean_logits, input_length, labels, mean=True):
287
+ """Multi-token KL-divergence metric.
288
+
289
+ For tasks where the answer is multi-token (long names, full
290
+ sentences, free-form generation), the single-token logit-diff
291
+ metric truncates to the first BPE subword and loses semantics.
292
+ KL-divergence between the model's full distribution at the answer
293
+ position(s) and the reference distribution captures the whole
294
+ answer profile.
295
+
296
+ Used as a substitute for ``_eap_logit_diff`` on tasks where
297
+ ``clean_answer`` and ``corrupt_answer`` differ in tokens beyond
298
+ the first subword.
299
+
300
+ Args:
301
+ logits (Tensor): Model logits [batch, seq_len, vocab_size].
302
+ clean_logits (Tensor): Reference logits at the same shape.
303
+ KL is computed as KL(softmax(logits) || softmax(clean_logits))
304
+ at the last real token position.
305
+ input_length (Tensor): Number of real tokens per example [batch].
306
+ labels (Tensor): Unused for KL; kept for signature compatibility.
307
+ mean (bool): If True, return scalar mean KL. If False, per-example.
308
+
309
+ Returns:
310
+ Tensor: Scalar mean KL or per-sample [batch].
311
+ """
312
+ if clean_logits is None:
313
+ # KL needs the reference distribution; degrade gracefully to logit-diff.
314
+ from .tasks.builtins.ioi import IOITaskSpec
315
+
316
+ return IOITaskSpec._ioi_logit_diff(logits, clean_logits, input_length, labels, mean=mean)
317
+
318
+ batch = logits.size(0)
319
+ last = (input_length.long() - 1).clamp_min(0)
320
+ arange = t.arange(batch, device=logits.device)
321
+ last_logits = logits[arange, last] # [batch, vocab]
322
+ last_clean = clean_logits[arange, last] # [batch, vocab]
323
+ log_p = t.nn.functional.log_softmax(last_logits, dim=-1)
324
+ log_q = t.nn.functional.log_softmax(last_clean, dim=-1)
325
+ p = log_p.exp()
326
+ kl = (p * (log_p - log_q)).sum(dim=-1) # [batch]
327
+ return kl.mean() if mean else kl
328
+
329
+
330
+ def _eap_accuracy(logits, clean_logits, input_length, labels, mean=True):
331
+ """
332
+ Token-prediction accuracy metric for EAP attribution.
333
+
334
+ Selects the logit at each example's last real token position and checks
335
+ whether the argmax matches the correct token. Handles both single-token
336
+ labels (IOI) and multi-column label tensors (MMLU); in the latter case
337
+ column 0 is treated as the correct answer.
338
+
339
+ Args:
340
+ logits (Tensor): Model logits [batch, seq_len, vocab_size].
341
+ clean_logits (Tensor): Unused; kept for metric signature compatibility.
342
+ input_length (Tensor): Number of real tokens per example [batch].
343
+ labels (Tensor): Correct token indices. Shape [batch] or [batch, n_options];
344
+ if 2-D, labels[:, 0] is used as the correct token.
345
+ mean (bool): If True, return the batch mean accuracy scalar.
346
+ If False, return per-sample accuracy [batch]. Defaults to True.
347
+
348
+ Returns:
349
+ Tensor: Scalar mean accuracy (mean=True) or per-sample float tensor [batch].
350
+ """
351
+ # Added Debugging here because this is a frequent crash point
352
+ # debug_metric_shapes(logits, labels, input_length)
353
+
354
+ batch_size = logits.size(0)
355
+ idx = t.arange(batch_size, device=logits.device)
356
+ logits = logits[idx, input_length - 1]
357
+
358
+ if labels.ndim > 1:
359
+ correct_token = labels[:, 0]
360
+ else:
361
+ correct_token = labels
362
+
363
+ correct_token = correct_token.to(logits.device)
364
+ predictions = logits.argmax(dim=-1)
365
+
366
+ results = (predictions == correct_token).float()
367
+
368
+ if mean:
369
+ results = results.mean()
370
+ return results
371
+
372
+
373
+ # def _convert_eap_scores_to_ck_format(graph: Graph) -> dict[str, float]:
374
+ # """
375
+ # Convert EAP node scores to CircuitKit's name-keyed score dict.
376
+
377
+ # Maps graph node names to CircuitKit naming convention:
378
+ # AttentionNode 'a{L}.h{H}' → 'A{L}.{H}', MLPNode 'm{L}' → 'MLP {L}'.
379
+ # Scores are absolute values of the raw node scores.
380
+
381
+ # Args:
382
+ # graph (Graph): Graph with populated node scores after attribution.
383
+
384
+ # Returns:
385
+ # Dict[str, float]: {'A0.0': score, 'MLP 0': score, ...} for all
386
+ # AttentionNode and MLPNode instances in the graph.
387
+ # """
388
+ # node_scores_dict = {}
389
+ # for node in graph.nodes.values():
390
+ # if isinstance(node, (AttentionNode, MLPNode)):
391
+ # score = abs(node.score.item())
392
+ # if isinstance(node, AttentionNode):
393
+ # circuit_kit_name = f"A{node.layer}.{node.head}"
394
+ # else: # MLPNode
395
+ # circuit_kit_name = f"MLP {node.layer}"
396
+ # node_scores_dict[circuit_kit_name] = score
397
+ # return node_scores_dict
398
+
399
+ from .backends.eap.circuit_kit_adapter import ( # noqa: E402 - import after intentional pre-import setup
400
+ convert_eap_graph_to_circuitkit_scores as _convert_eap_scores_to_ck_format,
401
+ )
402
+
403
+
404
+ def _ib_name_to_graph_name(ib_name: str) -> Optional[str]:
405
+ """Convert IBCircuit score key ('A0.0', 'MLP 0') to EAP Graph node key ('a0.h0', 'm0')."""
406
+ attn_match = re.match(r"A(\d+)\.(\d+)$", ib_name)
407
+ if attn_match:
408
+ return f"a{attn_match.group(1)}.h{attn_match.group(2)}"
409
+ mlp_match = re.match(r"MLP (\d+)$", ib_name)
410
+ if mlp_match:
411
+ return f"m{mlp_match.group(1)}"
412
+ return None
413
+
414
+
415
+ def _populate_graph_from_ib_scores(graph: Graph, ib_node_scores: dict) -> Graph:
416
+ """
417
+ Write IBCircuit node scores into a Graph's nodes_scores tensor.
418
+
419
+ Converts IBCircuit naming ('A0.0', 'MLP 0') to EAP graph naming ('a0.h0',
420
+ 'm0') and records absolute scores. Any node absent from ib_node_scores
421
+ (i.e. out-of-scope) is pinned to inf so that graph.apply_topn() always
422
+ retains it — this is the mechanism that enforces scope constraints.
423
+
424
+ Args:
425
+ graph (Graph): Graph initialised with node_scores=True. Its
426
+ nodes_scores tensor is overwritten in-place.
427
+ ib_node_scores (Dict[str, float]): Scores keyed by IBCircuit node
428
+ names, e.g. {'A0.0': 0.42, 'MLP 3': 0.07}.
429
+
430
+ Returns:
431
+ Graph: The same graph object, mutated in-place.
432
+ """
433
+ graph.nodes_scores = t.full((graph.n_forward,), float("nan"))
434
+ for ib_name, score in ib_node_scores.items():
435
+ graph_name = _ib_name_to_graph_name(ib_name)
436
+ if graph_name and graph_name in graph.nodes:
437
+ node = graph.nodes[graph_name]
438
+ node.score = t.tensor(abs(float(score)))
439
+ fwd_idx = graph.forward_index(node, attn_slice=False)
440
+ graph.nodes_scores[fwd_idx] = abs(float(score))
441
+
442
+ # Any node not scored by IB (nan) is pinned to inf so apply_topn
443
+ # always keeps it. For scope='heads', this catches all MLPs.
444
+ # For scope='mlp', this catches all attention heads.
445
+ # For scope='both', all nodes are scored and nothing is pinned.
446
+ for node in graph.nodes.values():
447
+ if isinstance(node, (AttentionNode, MLPNode)):
448
+ fwd_idx = graph.forward_index(node, attn_slice=False)
449
+ if t.isnan(graph.nodes_scores[fwd_idx]).any():
450
+ node.score = t.tensor(float("inf"))
451
+ graph.nodes_scores[fwd_idx] = float("inf")
452
+
453
+ return graph
454
+
455
+
456
+ def _validate_ibcircuit_dataloader(dataloader) -> None:
457
+ """
458
+ Validate that dataloader provides IBCircuit-compatible batches.
459
+
460
+ IBCircuit requires batches with specific keys:
461
+ - 'tokens': Input token IDs [batch_size, seq_len]
462
+ - 'labels': Answer token IDs [batch_size]
463
+ - 'answer_positions': Positions where answers appear [batch_size]
464
+
465
+ Args:
466
+ dataloader: DataLoader to validate
467
+
468
+ Raises:
469
+ ValueError: If dataloader format is incompatible
470
+ StopIteration: If dataloader is empty
471
+ """
472
+ try:
473
+ # Extract one batch for validation
474
+ batch = next(iter(dataloader))
475
+ except StopIteration:
476
+ raise ValueError(
477
+ "IBCircuit dataloader is empty. Ensure your task's "
478
+ "build_dataloader() method returns a non-empty DataLoader."
479
+ )
480
+
481
+ # Check required keys
482
+ required_keys = {"tokens", "labels", "answer_positions"}
483
+ actual_keys = set(batch.keys())
484
+ missing_keys = required_keys - actual_keys
485
+
486
+ if missing_keys:
487
+ raise ValueError(
488
+ f"IBCircuit dataloader missing required keys: {missing_keys}.\n"
489
+ f"Got keys: {list(actual_keys)}\n"
490
+ f"Required keys: {list(required_keys)}\n\n"
491
+ f"Your task's build_dataloader() method must return a DataLoader "
492
+ f"that yields batches with these exact keys. See the IBCircuit "
493
+ f"documentation for the expected batch format."
494
+ )
495
+
496
+ # Validate types and shapes
497
+ if not isinstance(batch["tokens"], t.Tensor):
498
+ raise ValueError(f"batch['tokens'] must be a torch.Tensor, got {type(batch['tokens'])}")
499
+
500
+ if not isinstance(batch["labels"], t.Tensor):
501
+ raise ValueError(f"batch['labels'] must be a torch.Tensor, got {type(batch['labels'])}")
502
+
503
+ if not isinstance(batch["answer_positions"], t.Tensor):
504
+ raise ValueError(
505
+ f"batch['answer_positions'] must be a torch.Tensor, "
506
+ f"got {type(batch['answer_positions'])}"
507
+ )
508
+
509
+ # Validate shapes are consistent
510
+ batch_size = batch["tokens"].shape[0]
511
+
512
+ if batch["labels"].shape[0] != batch_size:
513
+ raise ValueError(
514
+ f"Batch size mismatch: tokens has {batch_size} examples but "
515
+ f"labels has {batch['labels'].shape[0]} examples"
516
+ )
517
+
518
+ if batch["answer_positions"].shape[0] != batch_size:
519
+ raise ValueError(
520
+ f"Batch size mismatch: tokens has {batch_size} examples but "
521
+ f"answer_positions has {batch['answer_positions'].shape[0]} examples"
522
+ )
523
+
524
+ # Validate answer_positions are within sequence bounds
525
+ seq_len = batch["tokens"].shape[1]
526
+ max_pos = batch["answer_positions"].max().item()
527
+
528
+ if max_pos >= seq_len:
529
+ raise ValueError(
530
+ f"Invalid answer_positions: max position {max_pos} is >= "
531
+ f"sequence length {seq_len}. All answer positions must be "
532
+ f"valid indices into the sequence."
533
+ )
534
+
535
+
536
+ # ── Shared helpers used by discover_circuit and evaluate_circuit ──────────────────
537
+
538
+
539
+ def _avg_scores(scores) -> float:
540
+ """
541
+ Reduce per-sample metric scores to a single Python float.
542
+
543
+ Args:
544
+ scores (Tensor | List[Tensor]): Per-sample scores. If a list, each
545
+ element is averaged first, then those averages are averaged.
546
+
547
+ Returns:
548
+ float: Mean score across all samples.
549
+ """
550
+ if isinstance(scores, list):
551
+ return t.mean(t.stack([t.mean(s.float()) for s in scores])).item()
552
+ return t.mean(scores.float()).item() if scores.numel() > 1 else scores.item()
553
+
554
+
555
+ def _make_eval_metric(task_spec):
556
+ """
557
+ Build a per-sample, non-loss metric callable from a TaskSpec.
558
+
559
+ For partial-based metrics, overrides 'loss=False' and 'mean=False' so the
560
+ metric returns raw per-sample scores suitable for faithfulness evaluation.
561
+ Non-partial callables are returned unchanged.
562
+
563
+ Args:
564
+ task_spec: A registered TaskSpec with a metric_fn() method.
565
+
566
+ Returns:
567
+ Callable: Metric with signature
568
+ (logits, clean_logits, input_lengths, labels) -> Tensor [batch].
569
+ """
570
+ base = task_spec.metric_fn()
571
+ if isinstance(base, partial):
572
+ kw = base.keywords.copy()
573
+ kw["loss"] = False
574
+ kw["mean"] = False
575
+ return partial(base.func, **kw)
576
+ return base
577
+
578
+
579
+ def _compute_n_topn(graph: Graph, scope: str, sparsity: float):
580
+ """
581
+ Compute the apply_topn budget for node-level pruning under a given scope.
582
+
583
+ Out-of-scope nodes are always kept, so n_topn includes their count on top
584
+ of the in-scope budget. n_to_keep reflects only in-scope nodes and is used
585
+ when building an equivalently-sized random baseline.
586
+
587
+ Args:
588
+ graph (Graph): Graph containing n_layers and n_heads in its cfg.
589
+ scope (str): Which components are prunable — 'heads', 'mlp', or 'both'.
590
+ sparsity (float): Fraction of in-scope nodes to remove (0.0-1.0).
591
+
592
+ Returns:
593
+ Tuple[int, int]: (n_topn, n_to_keep) where n_topn is passed to
594
+ graph.apply_topn() and n_to_keep is the in-scope keep count.
595
+ """
596
+ n_layers = graph.cfg["n_layers"]
597
+ n_heads = n_layers * graph.cfg["n_heads"]
598
+ n_mlps = n_layers
599
+
600
+ if scope == "heads":
601
+ n_to_keep, n_always = int(n_heads * (1 - sparsity)), n_mlps
602
+ elif scope == "mlp":
603
+ n_to_keep, n_always = int(n_mlps * (1 - sparsity)), n_heads
604
+ else: # both
605
+ n_to_keep, n_always = int((n_heads + n_mlps) * (1 - sparsity)), 0
606
+
607
+ return n_to_keep + n_always, n_to_keep
608
+
609
+
610
+ def _build_random_node_graph(model, scope: str, n_to_keep: int, seed=None) -> Graph:
611
+ """
612
+ Build a randomly-pruned node-level Graph for use as a faithfulness baseline.
613
+
614
+ Only in-scope nodes are assigned a score (1.0) so they participate in
615
+ random selection; out-of-scope nodes keep NaN scores and are always
616
+ retained by apply_random. This ensures a fair comparison across all
617
+ algorithms regardless of scope.
618
+
619
+ Args:
620
+ model (HookedTransformer): Model whose config defines the graph structure.
621
+ scope (str): Which components are prunable — 'heads', 'mlp', or 'both'.
622
+ n_to_keep (int): Number of in-scope nodes to keep (from _compute_n_topn).
623
+ seed (Optional[int]): Random seed for reproducibility. Defaults to None.
624
+
625
+ Returns:
626
+ Graph: A pruned Graph with nodes_in_graph and in_graph set randomly.
627
+ """
628
+ rand = Graph.from_model(model, node_scores=True, neuron_level=False)
629
+ for node in rand.nodes.values():
630
+ if isinstance(node, (AttentionNode, MLPNode)):
631
+ fwd_idx = rand.forward_index(node)
632
+ in_scope = (
633
+ (scope == "heads" and isinstance(node, AttentionNode))
634
+ or (scope == "mlp" and isinstance(node, MLPNode))
635
+ or scope == "both"
636
+ )
637
+ if in_scope:
638
+ rand.nodes_scores[fwd_idx] = 1.0
639
+ # out-of-scope stays NaN → apply_random always keeps it
640
+ rand.apply_random(n_to_keep, level="node", prune=True, seed=seed)
641
+ return rand
642
+
643
+
644
+ def _build_random_ibcircuit_neuron_pruning_dict(
645
+ model: HookedTransformer,
646
+ reference_pruning_dict: dict,
647
+ scope: str,
648
+ seed: int = None,
649
+ ) -> dict:
650
+ """
651
+ Build a random neuron pruning dict matching the IBCircuit discovery budget.
652
+
653
+ Samples the same total number of neurons as the reference dict, drawn
654
+ uniformly from the same neuron space IBCircuit searches:
655
+ - Attention: d_head neurons per head (hook_z space).
656
+ - MLP: d_mlp or d_model neurons per layer depending on mlp_hook.
657
+
658
+ Used as a random baseline to contextualise faithfulness scores.
659
+
660
+ Args:
661
+ model (HookedTransformer): Model whose architecture defines the neuron space.
662
+ reference_pruning_dict (dict): Pruning dict produced by IBCircuit discovery,
663
+ used solely to count the total neurons pruned. Expected keys: 'heads',
664
+ 'mlp', '_meta'.
665
+ scope (str): Which components to sample from — 'heads', 'mlp', or 'both'.
666
+ seed (Optional[int]): Random seed for reproducibility. Defaults to None.
667
+
668
+ Returns:
669
+ Dict: Pruning dict with keys 'mlp', 'heads', '_meta', in the same
670
+ format as the IBCircuit discovery output.
671
+ """
672
+ n_to_prune = sum(len(v) for v in reference_pruning_dict.get("heads", {}).values()) + sum(
673
+ len(v) for v in reference_pruning_dict.get("mlp", {}).values()
674
+ )
675
+
676
+ n_layers = model.cfg.n_layers
677
+ n_heads = model.cfg.n_heads
678
+ d_head = model.cfg.d_head
679
+ model.cfg.d_model
680
+
681
+ mlp_hook = reference_pruning_dict.get("_meta", {}).get("mlp_hook", "mlp_out")
682
+ mlp_dim = model.cfg.d_mlp if mlp_hook == "post_act" else model.cfg.d_model
683
+
684
+ all_neurons = []
685
+ if scope in ("heads", "both"):
686
+ for layer in range(n_layers):
687
+ for head in range(n_heads):
688
+ for ni in range(d_head):
689
+ all_neurons.append(("attn", layer, head, ni))
690
+ if scope in ("mlp", "both"):
691
+ for layer in range(n_layers):
692
+ for ni in range(mlp_dim):
693
+ all_neurons.append(("mlp", layer, None, ni))
694
+
695
+ if seed is not None:
696
+ t.manual_seed(seed)
697
+ perm = t.randperm(len(all_neurons)).tolist()
698
+ selected = [all_neurons[i] for i in perm[:n_to_prune]]
699
+
700
+ rand_mlp = defaultdict(list)
701
+ rand_heads = defaultdict(list)
702
+ for kind, layer, head, ni in selected:
703
+ if kind == "mlp":
704
+ rand_mlp[layer].append(ni)
705
+ else:
706
+ rand_heads[(layer, head)].append(ni)
707
+
708
+ return {"mlp": dict(rand_mlp), "heads": dict(rand_heads), "_meta": {"mlp_hook": mlp_hook}}
709
+
710
+
711
+ def _build_artifact_stem(config: Dict[str, Any]) -> str:
712
+ """
713
+ Build a descriptive filename stem from a discovery config.
714
+
715
+ Format: '{algo}_{task}_{model}_{scope_or_level}_sp{sparsity}[_{extras}]'
716
+ Examples:
717
+ eap-ig_ioi_gpt2_neuron_sp0.3
718
+ ibcircuit_ioi_gpt2-small_heads_sp0.2_e1000
719
+ acdc_greater-than_pythia-70m_node_sp0.5
720
+
721
+ Args:
722
+ config (Dict[str, Any]): Validated discovery config containing
723
+ 'model', 'discovery', and 'pruning' sub-dicts.
724
+
725
+ Returns:
726
+ str: Underscore-joined filename stem, safe for use in file paths.
727
+ """
728
+ disc = config["discovery"]
729
+ prune = config["pruning"]
730
+ algo = disc["algorithm"].lower()
731
+ task = disc["task"]
732
+ model = config["model"]["name"].split("/")[-1] # strip org prefix
733
+ # Use defaults from DEFAULT_CONFIG (single source of truth)
734
+ default_discovery = DEFAULT_CONFIG["discovery"]
735
+ default_pruning = DEFAULT_CONFIG["pruning"]
736
+ level = disc.get("level", default_discovery.get("level"))
737
+ scope = (
738
+ disc.get("scope", default_discovery.get("scope"))
739
+ if algo == "ibcircuit"
740
+ else prune.get("scope", default_pruning.get("scope"))
741
+ )
742
+ sp = prune.get("target_sparsity", default_pruning.get("target_sparsity"))
743
+
744
+ parts = [algo, task, model, scope if algo == "ibcircuit" else level, f"sp{sp}"]
745
+
746
+ # Algo-specific differentiators
747
+ if algo == "ibcircuit":
748
+ parts.append(f"e{disc.get('num_epochs', default_discovery.get('num_epochs'))}")
749
+ if disc.get("mlp_hook", default_discovery.get("mlp_hook")) != default_discovery.get(
750
+ "mlp_hook"
751
+ ):
752
+ parts.append(disc["mlp_hook"])
753
+ elif algo in ("eap", "eap-ig"):
754
+ if disc.get("method"):
755
+ parts.append(disc["method"].lower().replace("-", ""))
756
+
757
+ return "_".join(str(p) for p in parts)
758
+
759
+
760
+ def _save_artifact(
761
+ data: Any, output_path: str, suffix: str, logger, config: Dict = None
762
+ ) -> Optional[str]:
763
+ """
764
+ Save a discovery artifact to disk as a .pt file.
765
+
766
+ If output_path is not provided, a path is auto-generated from the config
767
+ stem under '{cwd}/outputs/'. A suffix (e.g. '_scores') is appended to
768
+ the stem before the extension to differentiate artifact types saved at
769
+ the same base path.
770
+
771
+ Args:
772
+ data (Any): Serialisable object to save (passed to torch.save).
773
+ output_path (str): Base output path or directory. If a directory or
774
+ no extension, a filename is generated from the config stem.
775
+ suffix (str): String appended to the filename stem, e.g. '_scores'
776
+ or '_ib_weights'.
777
+ logger: Logger instance for info messages.
778
+ config (Optional[Dict]): Discovery config used to build the filename
779
+ stem when output_path is absent or a directory.
780
+
781
+ Returns:
782
+ Optional[str]: Absolute path of the saved file, or None if no path
783
+ could be determined.
784
+ """
785
+ if not output_path and config:
786
+ stem = _build_artifact_stem(config)
787
+ output_path = os.path.join(os.getcwd(), "outputs", stem + ".pt")
788
+ if not output_path:
789
+ return None
790
+ from pathlib import Path
791
+
792
+ p = Path(output_path)
793
+ # If output_path is a directory, generate filename inside it
794
+ if p.is_dir() or not p.suffix:
795
+ stem = _build_artifact_stem(config) if config else "circuit"
796
+ p = p / (stem + ".pt")
797
+ dest = p.parent / (p.stem + suffix + p.suffix)
798
+ os.makedirs(str(p.parent), exist_ok=True)
799
+ t.save(data, str(dest))
800
+ logger.info(f"Saved '{suffix.lstrip('_')}' → {dest}")
801
+ return str(dest)
802
+
803
+
804
+ def _build_circuit_scores(
805
+ task: str,
806
+ model_name: str,
807
+ algorithm: str,
808
+ node_scores: Dict[str, float],
809
+ discovery_cfg: Optional[Dict] = None,
810
+ ) -> CircuitScores:
811
+ """
812
+ Build a CircuitScores artifact from discovered scores.
813
+
814
+ Helper to standardize the creation of CircuitScores across all backends.
815
+
816
+ Args:
817
+ task: Task name (e.g., 'ioi', 'mmlu').
818
+ model_name: Model identifier (e.g., 'gpt2').
819
+ algorithm: Algorithm name ('eap', 'eap-ig', 'acdc', 'ibcircuit').
820
+ node_scores: Dict mapping node names to scores.
821
+ discovery_cfg: Optional discovery configuration.
822
+
823
+ Returns:
824
+ CircuitScores artifact with timestamp and metadata.
825
+ """
826
+ return CircuitScores(
827
+ task=task,
828
+ model=model_name,
829
+ algorithm=algorithm,
830
+ level="node",
831
+ node_scores=node_scores,
832
+ timestamp=CircuitScores.create_timestamp(),
833
+ version="1.0",
834
+ discovery_cfg=discovery_cfg or {},
835
+ )
836
+
837
+
838
+ # ────────────────────────────────
839
+
840
+ def prepare_custom_task(
841
+ config: Dict[str, Any],
842
+ model: HookedTransformer,
843
+ task_name: Optional[str] = None,
844
+ ) -> str:
845
+ """
846
+ Normalise a config["data"] block into a registered CircuitKit task.
847
+
848
+ Must be called once before discover_circuit() and evaluate_circuit()
849
+ when config contains a "data" block. Mutates config in-place: sets
850
+ config["discovery"]["task"] to the registered name and removes
851
+ config["data"] so neither downstream function re-processes it.
852
+
853
+ Args:
854
+ config: Full CircuitKit config dict with a "data" block.
855
+ model: Loaded HookedTransformer (tokenizer used for alignment).
856
+ task_name: Explicit registry name. Defaults to "custom:{csv_stem}".
857
+
858
+ Returns:
859
+ The registered task name string.
860
+ """
861
+ data_cfg = config.get("data")
862
+ if not data_cfg:
863
+ return config["discovery"]["task"]
864
+
865
+ data_type = data_cfg.get("type")
866
+ if not data_type:
867
+ raise KeyError(
868
+ "config['data']['type'] is required ('template', 'auto', or 'clean_only')"
869
+ )
870
+
871
+ if task_name is None:
872
+ task_name = f"custom:{Path(data_cfg['path']).stem}"
873
+
874
+ logger = get_logger("circuitkit.custom_data")
875
+
876
+ if data_type == "template":
877
+ from .data.template import clean_only_from_template, template_normalize
878
+ from .tasks._algorithm_families import CDT_FAMILY, IB_FAMILY
879
+
880
+ template = data_cfg.get("template", {})
881
+ algo = config["discovery"].get("algorithm", "").lower()
882
+ is_clean_only_algo = algo in (IB_FAMILY | CDT_FAMILY)
883
+ has_corrupt_keys = bool(template.get("corrupt_prompt") and template.get("corrupt_answer"))
884
+
885
+ if is_clean_only_algo and not has_corrupt_keys:
886
+ # Algorithm only needs the clean side; skip full pairing pipeline.
887
+ if not template.get("clean_prompt"):
888
+ raise ValueError(
889
+ "config['data']['template'] must contain 'clean_prompt'."
890
+ )
891
+ ds = clean_only_from_template(
892
+ data_cfg["path"],
893
+ template_spec=template,
894
+ max_records=data_cfg.get("max_records"),
895
+ name=Path(data_cfg["path"]).stem,
896
+ source=data_cfg["path"],
897
+ )
898
+ logger.info(
899
+ f"template (clean-only extraction): {len(ds)} records loaded "
900
+ f"for algorithm '{algo}' (corrupt keys omitted, no alignment pass)"
901
+ )
902
+ else:
903
+ required = ["clean_prompt", "corrupt_prompt", "clean_answer", "corrupt_answer"]
904
+ missing = [k for k in required if not template.get(k)]
905
+ if missing:
906
+ raise ValueError(
907
+ f"config['data']['template'] is missing required fields: {missing}"
908
+ )
909
+ align_strategy = data_cfg.get("align_strategy", "filter")
910
+ ds = template_normalize(
911
+ data_cfg["path"],
912
+ template_spec=template,
913
+ pairing_mode=data_cfg.get("pairing_mode", "explicit"),
914
+ align_strategy=align_strategy,
915
+ tokenizer=model.tokenizer,
916
+ pad_region_end=data_cfg.get("pad_region_end"),
917
+ max_records=data_cfg.get("max_records"),
918
+ name=Path(data_cfg["path"]).stem,
919
+ source=data_cfg["path"],
920
+ )
921
+ align_meta = ds.meta.get("_alignment", {})
922
+ logger.info(
923
+ f"template dataset: {align_meta.get('kept')}/{align_meta.get('total_input')} "
924
+ f"records kept after alignment (strategy={align_strategy!r}, "
925
+ f"dropped_nondiscriminative={align_meta.get('dropped_nondiscriminative')}, "
926
+ f"dropped_misaligned={align_meta.get('dropped_misaligned')}, "
927
+ f"dropped_pad_failed={align_meta.get('dropped_pad_failed')}, "
928
+ f"recommended_metric={align_meta.get('recommended_metric')!r})"
929
+ )
930
+
931
+ elif data_type == "auto":
932
+ from .data.auto_detect import auto_normalize
933
+
934
+ ds = auto_normalize(
935
+ data_cfg["path"],
936
+ apply_default_strategy=True,
937
+ max_records=data_cfg.get("max_records"),
938
+ name=data_cfg.get("name", Path(data_cfg["path"]).stem),
939
+ source=data_cfg["path"],
940
+ )
941
+
942
+ elif data_type == "clean_only":
943
+ from .data.clean_only import clean_only_normalize
944
+
945
+ ds = clean_only_normalize(
946
+ data_cfg["path"],
947
+ prompt_column=data_cfg.get("prompt_column", "prompt"),
948
+ answer_column=data_cfg.get("answer_column", "answer"),
949
+ max_records=data_cfg.get("max_records"),
950
+ name=data_cfg.get("name", Path(data_cfg["path"]).stem),
951
+ source=data_cfg["path"],
952
+ )
953
+ logger.info(
954
+ f"clean_only dataset: {len(ds)} records loaded "
955
+ f"(no corrupt partner; compatible with ibcircuit, cdt)"
956
+ )
957
+
958
+ else:
959
+ raise ValueError(
960
+ f"Unknown data.type {data_type!r}. Use 'template', 'auto', or 'clean_only'."
961
+ )
962
+
963
+ from .data.normalized_task import NormalizedTaskSpec
964
+
965
+ task_spec = NormalizedTaskSpec(ds, name=task_name)
966
+ padding = data_cfg.get("pair_padding_side")
967
+ if padding in ("left", "right"):
968
+ task_spec.pair_padding_side = padding
969
+
970
+ try:
971
+ _register_task(task_spec)
972
+ except ValueError as e:
973
+ if "already registered" not in str(e):
974
+ raise
975
+ logger.info(f"Task '{task_name}' already registered, reusing.")
976
+
977
+ config["discovery"]["task"] = task_name
978
+ config.pop("data", None)
979
+ logger.info(f"Custom task '{task_name}' registered ({len(ds)} records, {ds.n_paired} paired)")
980
+ return task_name
981
+
982
+ @debug_function
983
+ @handle_errors(context={"operation": "discover_circuit"})
984
+ def discover_circuit( # noqa: C901 - complex function, refactor out of scope for lint pass
985
+ config: Union[str, Dict[str, Any]],
986
+ _model: Optional[HookedTransformer] = None,
987
+ ) -> Union[List[str], Dict]:
988
+ """
989
+ Run circuit discovery and return a pruning artifact.
990
+
991
+ Loads the model and task, runs the specified attribution algorithm,
992
+ applies sparsity-based pruning, and optionally evaluates faithfulness.
993
+
994
+ Args:
995
+ config: Path to a YAML file or a config dict with keys:
996
+ ``model.name``, ``model.precision``, ``discovery.algorithm``,
997
+ ``discovery.task``, ``discovery.level``, ``pruning.target_sparsity``,
998
+ ``pruning.scope``, ``output_path``.
999
+ _model: Internal. An already-loaded HookedTransformer to reuse.
1000
+ Leave as ``None`` for external callers.
1001
+
1002
+ Returns:
1003
+ Node-level: list of node name strings. Neuron-level: dict with
1004
+ ``mlp``, ``heads``, and ``_meta`` keys.
1005
+
1006
+ Raises:
1007
+ ValueError: If required config keys are missing.
1008
+ AlgorithmError: If the algorithm is not recognised.
1009
+ """
1010
+ # Bootstrap built-in tasks
1011
+ from .tasks.bootstrap import _bootstrap_builtin_tasks
1012
+
1013
+ _bootstrap_builtin_tasks()
1014
+
1015
+ logger = get_logger("circuitkit.discovery")
1016
+ progress = ProgressLogger(logger)
1017
+
1018
+ # Snapshot of the caller's global RNG state, captured iff we seed below so
1019
+ # the ``finally`` can restore it (see the seed block).
1020
+ _rng_snapshot = None
1021
+
1022
+ try:
1023
+ # Load, merge defaults, and validate the config in one step
1024
+ progress.start_operation("Circuit Discovery", 4)
1025
+ progress.step("Loading and validating configuration")
1026
+
1027
+ config = load_and_validate_config(config)
1028
+ logger.log_config(config)
1029
+
1030
+ model_cfg = config["model"]
1031
+ discovery_cfg = config["discovery"]
1032
+ pruning_cfg = config["pruning"]
1033
+
1034
+ # Seed all global RNGs when the config supplies a seed, so stochastic
1035
+ # algorithms (IBCircuit, CD-T/ACDC data generation via numpy/`random`)
1036
+ # are reproducible. We snapshot the caller's global RNG state first and
1037
+ # restore it in the ``finally`` — otherwise discovery would permanently
1038
+ # reseed the whole process's numpy/random/torch RNGs as a side effect.
1039
+ _seed = discovery_cfg.get("seed", discovery_cfg.get("data_params", {}).get("seed"))
1040
+ if _seed is not None:
1041
+ import random as _random_std
1042
+ import numpy as _np_std
1043
+ _rng_snapshot = (
1044
+ t.get_rng_state(),
1045
+ _np_std.random.get_state(),
1046
+ _random_std.getstate(),
1047
+ t.cuda.get_rng_state_all() if t.cuda.is_available() else None,
1048
+ )
1049
+ t.manual_seed(_seed)
1050
+ _np_std.random.seed(_seed)
1051
+ _random_std.seed(_seed)
1052
+ if t.cuda.is_available():
1053
+ t.cuda.manual_seed_all(_seed)
1054
+
1055
+ is_verbose = discovery_cfg.get("verbose", False)
1056
+ if is_verbose:
1057
+ import logging
1058
+
1059
+ get_logger("circuitkit").setLevel(logging.DEBUG)
1060
+ get_logger("data").setLevel(logging.DEBUG)
1061
+ logger.debug(f"Discovery Config: {config['discovery']}")
1062
+
1063
+ # Resolve and Sanitize Discovery Intervention
1064
+ # Use default from DEFAULT_CONFIG (single source of truth)
1065
+ default_discovery = DEFAULT_CONFIG["discovery"]
1066
+ discovery_intervention = discovery_cfg.get(
1067
+ "intervention", default_discovery.get("intervention")
1068
+ )
1069
+
1070
+ # IG methods strictly require patching to function
1071
+ if discovery_cfg["algorithm"].lower() in [
1072
+ "eap-ig",
1073
+ "eap-ig-activations",
1074
+ "clean-corrupted",
1075
+ ]:
1076
+ if discovery_intervention != "patching":
1077
+ logger.warning(
1078
+ f"Safety Override: {discovery_cfg['algorithm']} requires 'patching'. "
1079
+ f"Changing discovery intervention from '{discovery_intervention}' to 'patching'."
1080
+ )
1081
+ discovery_intervention = "patching"
1082
+
1083
+ if is_verbose:
1084
+ if discovery_cfg["algorithm"].lower() == "ibcircuit":
1085
+ # IBCircuit uses stochastic mean-ablation via IB Noise
1086
+ logger.debug("Discovery Phase Intervention: IB Noise (Stochastic Mean-Ablation)")
1087
+ else:
1088
+ logger.debug(f"Discovery Phase Intervention: {discovery_intervention}")
1089
+
1090
+ # Validate model name
1091
+ validate_model_name(model_cfg["name"])
1092
+ if "algorithm" not in discovery_cfg:
1093
+ from .backends import DISCOVERY_ALGORITHMS
1094
+
1095
+ raise ValueError(
1096
+ "Discovery config is missing the required key 'algorithm'. "
1097
+ "Add an 'algorithm' key under the discovery config. "
1098
+ "Supported discovery algorithms: "
1099
+ f"{', '.join(sorted(DISCOVERY_ALGORITHMS))}."
1100
+ )
1101
+ validate_discovery_algorithm(discovery_cfg["algorithm"])
1102
+
1103
+ progress.step("Setting up model", model=model_cfg["name"])
1104
+ device = get_device()
1105
+ # Use default from DEFAULT_CONFIG (single source of truth)
1106
+ default_model = DEFAULT_CONFIG["model"]
1107
+ dtype = getattr(t, model_cfg.get("precision", default_model.get("precision")))
1108
+
1109
+ if _model is not None:
1110
+ # Reuse the caller's already-loaded model instead of loading a
1111
+ # second full copy (e.g. quick.discover()/Pipeline.discover()
1112
+ # already built one via load_model()/_ensure_model()).
1113
+ model = _model
1114
+ logger.debug("discover_circuit: reusing pre-loaded model, skipping reload")
1115
+ else:
1116
+ with log_execution_time("Model loading", logger):
1117
+ model = HookedTransformer.from_pretrained(
1118
+ model_cfg["name"], device=device, dtype=dtype
1119
+ )
1120
+
1121
+ algo = discovery_cfg["algorithm"].lower()
1122
+
1123
+ if hasattr(model.cfg, "ungroup_grouped_query_attention"):
1124
+ model.cfg.ungroup_grouped_query_attention = True
1125
+
1126
+ # ── Inline data path: delegate to prepare_custom_task ───────────
1127
+ if config.get("data") and config["data"].get("type"):
1128
+ prepare_custom_task(config, model=model)
1129
+ discovery_cfg["task"] = config["discovery"]["task"]
1130
+ # ── End inline data path ─────────────────────────────────────────
1131
+
1132
+ # Resolve and validate task spec (explicit, no defaults)
1133
+ if "task" not in discovery_cfg:
1134
+ from .tasks.registry import list_tasks
1135
+
1136
+ raise ValueError(
1137
+ "Discovery config is missing the required key 'task'. "
1138
+ "Add a 'task' key under the discovery config naming the task to "
1139
+ f"discover. Registered tasks: {list_tasks()}."
1140
+ )
1141
+ task_spec = _get_task(discovery_cfg["task"])
1142
+ task_spec.validate_discovery_config(discovery_cfg)
1143
+
1144
+ if algo in (
1145
+ "acdc",
1146
+ "eap",
1147
+ "eap-ig",
1148
+ "eap-ig-activations",
1149
+ "eap-clean-corrupted",
1150
+ "eap-exact",
1151
+ "atp-gd",
1152
+ "eap-gp",
1153
+ "relp",
1154
+ "peap",
1155
+ "eap-ifr",
1156
+ ):
1157
+ model.cfg.use_attn_result = True
1158
+ model.cfg.use_split_qkv_input = True
1159
+ model.cfg.use_hook_mlp_in = True
1160
+
1161
+ # Warn about experimental / research algorithms
1162
+ if algo in RESEARCH_ALGORITHMS:
1163
+ _warnings.warn(
1164
+ f"Algorithm '{algo}' is research-quality (only validated on GPT-2 IOI). "
1165
+ f"Use '{_DEFAULT_ALGO}' for production.",
1166
+ UserWarning,
1167
+ stacklevel=2,
1168
+ )
1169
+ elif algo in EXPERIMENTAL_ALGORITHMS:
1170
+ _warnings.warn(
1171
+ f"Algorithm '{algo}' is experimental. May fail on larger models or non-IOI tasks. "
1172
+ f"Use '{_DEFAULT_ALGO}' for production.",
1173
+ UserWarning,
1174
+ stacklevel=2,
1175
+ )
1176
+
1177
+ logger.log_model_info(
1178
+ model_cfg["name"],
1179
+ device=device,
1180
+ dtype=str(dtype),
1181
+ parameters=sum(p.numel() for p in model.parameters()),
1182
+ )
1183
+
1184
+ progress.step("Running discovery algorithm", algorithm=algo)
1185
+ logger.info(f"Starting {algo.upper()} discovery algorithm")
1186
+
1187
+ # Use defaults from DEFAULT_CONFIG (single source of truth)
1188
+ default_discovery = DEFAULT_CONFIG["discovery"]
1189
+ _ib_scope = discovery_cfg.get("scope", default_discovery.get("scope"))
1190
+
1191
+ if algo == "acdc":
1192
+ with debug_context("ACDC Discovery"):
1193
+ p_model = patchable_model(
1194
+ model,
1195
+ factorized=True,
1196
+ slice_output="last_seq",
1197
+ separate_qkv=True,
1198
+ device=device,
1199
+ )
1200
+ train_loader, _ = load_task_data(
1201
+ task_name=discovery_cfg["task"],
1202
+ model=model,
1203
+ device=device,
1204
+ **discovery_cfg.get("data_params", {}),
1205
+ )
1206
+ # ACDC sweeps one full edge pass per (base, exp) tao value.
1207
+ # With the library defaults (5 bases x 4 exps = 20 sweeps of
1208
+ # ~32k edges) a single GPT-2 run takes hours. Expose the tao
1209
+ # grid via discovery_cfg so callers can scope the search;
1210
+ # fall back to the backend defaults when unspecified.
1211
+ _acdc_kwargs = {}
1212
+ if "tao_exps" in discovery_cfg:
1213
+ _acdc_kwargs["tao_exps"] = list(discovery_cfg["tao_exps"])
1214
+ if "tao_bases" in discovery_cfg:
1215
+ _acdc_kwargs["tao_bases"] = list(discovery_cfg["tao_bases"])
1216
+ if "faithfulness_target" in discovery_cfg:
1217
+ _acdc_kwargs["faithfulness_target"] = discovery_cfg["faithfulness_target"]
1218
+ # verbose=True shows tqdm bars; False (default) emits progress as
1219
+ # DEBUG log messages on circuitkit.backends.acdc.prune_algos.ACDC.
1220
+ _acdc_kwargs["verbose"] = discovery_cfg.get("verbose", False)
1221
+ edge_scores = acdc_prune_scores(
1222
+ p_model, train_loader, official_edges=None, **_acdc_kwargs
1223
+ )
1224
+ node_scores = calculate_node_scores_from_edges(p_model, edge_scores)
1225
+
1226
+ # Build unified CircuitScores artifact (Workstream G)
1227
+ circuit_scores = _build_circuit_scores(
1228
+ task=discovery_cfg["task"],
1229
+ model_name=model_cfg["name"],
1230
+ algorithm=algo,
1231
+ node_scores=node_scores,
1232
+ discovery_cfg=discovery_cfg,
1233
+ )
1234
+
1235
+ # Save CircuitScores as JSON
1236
+ if config.get("output_path"):
1237
+ scores_path = Path(config["output_path"]).parent / (
1238
+ Path(config["output_path"]).stem + "_scores.json"
1239
+ )
1240
+ circuit_scores.to_json(scores_path)
1241
+ logger.info(f"Saved unified CircuitScores → {scores_path}")
1242
+
1243
+ # Also save legacy format for compatibility
1244
+ _save_artifact(
1245
+ {"algo": algo, "level": "node", "node_scores": node_scores},
1246
+ config.get("output_path"),
1247
+ "_scores",
1248
+ logger,
1249
+ )
1250
+
1251
+ elif algo in [
1252
+ "eap",
1253
+ "eap-ig",
1254
+ # Tier-0 promotions: top-level keys for the
1255
+ # 4 EAP-internal methods that previously could
1256
+ # only be selected via discovery_cfg['method'].
1257
+ "eap-ig-activations",
1258
+ "eap-clean-corrupted",
1259
+ "eap-exact",
1260
+ # AtP+GradDrop (Kramár et al. 2024) — same EAP backbone
1261
+ # with L gradient passes, one residual gradient zeroed each.
1262
+ "atp-gd",
1263
+ # EAP-GP (Zhang et al. 2025) — adaptive integration path
1264
+ # in input embedding space; replaces EAP-IG's straight line.
1265
+ "eap-gp",
1266
+ # RelP (Mohebbi et al. 2025) — LRP-style relevance
1267
+ # propagation via forward detach hooks; same EAP cost.
1268
+ "relp",
1269
+ # PEAP (Haklay et al. 2025) — per-position retention
1270
+ # of EAP scores; node-level summary preserved.
1271
+ "peap",
1272
+ # IFR / Information Flow Routes (Ferrando et al. 2024)
1273
+ # — proximity-based attribution, no metric needed.
1274
+ "eap-ifr",
1275
+ ]:
1276
+ with debug_context("EAP Discovery"):
1277
+ # Use TaskSpec for dataloader and metric
1278
+ dataloader = task_spec.build_dataloader(model, discovery_cfg, device)
1279
+
1280
+ if "level" not in discovery_cfg:
1281
+ raise ValueError(
1282
+ "Discovery config is missing the required key 'level'. "
1283
+ "Add a 'level' key under the discovery config set to "
1284
+ "'node' or 'neuron'."
1285
+ )
1286
+ is_neuron_level = discovery_cfg["level"] == "neuron"
1287
+ # Use defaults from DEFAULT_CONFIG (single source of truth)
1288
+ default_discovery = DEFAULT_CONFIG["discovery"]
1289
+ mlp_hook = discovery_cfg.get("mlp_hook", default_discovery.get("mlp_hook"))
1290
+ graph = Graph.from_model(
1291
+ model, node_scores=True, neuron_level=is_neuron_level, mlp_hook=mlp_hook
1292
+ )
1293
+
1294
+ metric = task_spec.metric_fn()
1295
+
1296
+ logger.debug(
1297
+ f"Graph initialized. Nodes: {len(graph.nodes)}. Neuron Level: {is_neuron_level}"
1298
+ )
1299
+
1300
+ # Map top-level algorithm keys to internal `method` arg.
1301
+ _ALGO_METHOD_MAP = {
1302
+ "eap": "EAP",
1303
+ "eap-ig": "EAP-IG-inputs",
1304
+ "eap-ig-activations": "EAP-IG-activations",
1305
+ "eap-clean-corrupted": "clean-corrupted",
1306
+ "eap-exact": "exact",
1307
+ "atp-gd": "atp-gd",
1308
+ "eap-gp": "eap-gp",
1309
+ "relp": "relp",
1310
+ "peap": "peap",
1311
+ "eap-ifr": "ifr",
1312
+ }
1313
+ if algo == "eap-ig":
1314
+ # eap-ig still supports an explicit method override
1315
+ # (legacy behaviour) for users who want to dispatch
1316
+ # via discovery_cfg['method'].
1317
+ _valid_node_methods = (
1318
+ "EAP",
1319
+ "EAP-IG-inputs",
1320
+ "EAP-IG-activations",
1321
+ "exact",
1322
+ "clean-corrupted",
1323
+ )
1324
+ _method = discovery_cfg.get("method", default_discovery.get("method"))
1325
+ if _method not in _valid_node_methods:
1326
+ raise ValueError(
1327
+ f"discovery config key 'method' has invalid value "
1328
+ f"{_method!r} for algorithm 'eap-ig'. "
1329
+ f"Set 'method' to one of: {list(_valid_node_methods)}. "
1330
+ f"Note: the algorithm name 'eap-ig' is not itself a "
1331
+ f"valid 'method' string."
1332
+ )
1333
+ else:
1334
+ _method = _ALGO_METHOD_MAP[algo]
1335
+ attribute_node(
1336
+ model,
1337
+ graph,
1338
+ dataloader,
1339
+ metric,
1340
+ method=_method,
1341
+ ig_steps=discovery_cfg.get("ig_steps", default_discovery.get("ig_steps")),
1342
+ neuron=is_neuron_level,
1343
+ intervention=discovery_intervention,
1344
+ )
1345
+
1346
+ if not is_neuron_level:
1347
+ node_scores = _convert_eap_scores_to_ck_format(graph)
1348
+
1349
+ # Build unified CircuitScores artifact (Workstream G)
1350
+ circuit_scores = _build_circuit_scores(
1351
+ task=discovery_cfg["task"],
1352
+ model_name=model_cfg["name"],
1353
+ algorithm=algo,
1354
+ node_scores=node_scores,
1355
+ discovery_cfg=discovery_cfg,
1356
+ )
1357
+
1358
+ # Save CircuitScores as JSON
1359
+ if config.get("output_path"):
1360
+ scores_path = Path(config["output_path"]).parent / (
1361
+ Path(config["output_path"]).stem + "_scores.json"
1362
+ )
1363
+ circuit_scores.to_json(scores_path)
1364
+ logger.info(f"Saved unified CircuitScores → {scores_path}")
1365
+
1366
+ # Also save legacy format for compatibility
1367
+ _save_artifact(
1368
+ {"algo": algo, "level": "node", "node_scores": node_scores},
1369
+ config.get("output_path"),
1370
+ "_scores",
1371
+ logger,
1372
+ )
1373
+
1374
+ else:
1375
+ # Handle neuron-level results
1376
+ default_pruning = DEFAULT_CONFIG["pruning"]
1377
+ effective_scope = pruning_cfg.get("scope", default_pruning.get("scope"))
1378
+ logger.info(
1379
+ f"Processing neuron-level scores (scope: {effective_scope}, strategy: per-layer)"
1380
+ )
1381
+ pruned_mlp_neurons = defaultdict(list)
1382
+ pruned_attn_neurons = defaultdict(list)
1383
+
1384
+ all_neuron_scores = []
1385
+ for node in tqdm(graph.nodes.values(), desc="Extracting neuron scores"):
1386
+ if isinstance(node, (MLPNode, AttentionNode)):
1387
+ # Filter by scope
1388
+ if effective_scope == "mlp" and not isinstance(node, MLPNode):
1389
+ continue
1390
+ if effective_scope == "heads" and not isinstance(node, AttentionNode):
1391
+ continue
1392
+
1393
+ fwd_index = graph.forward_index(node, attn_slice=False)
1394
+ scores_tensor = graph.neurons_scores[fwd_index].clone().detach().cpu()
1395
+ # Truncate to actual activation dimension to avoid counting padding zeros
1396
+ valid_scores = scores_tensor[: node.d_neuron]
1397
+ for neuron_idx, score in enumerate(valid_scores):
1398
+ all_neuron_scores.append(
1399
+ (abs(score.item()), (node.name, neuron_idx))
1400
+ )
1401
+
1402
+ all_neuron_scores.sort(key=lambda x: x[0]) # Sort by absolute score, ascending
1403
+ num_to_prune = int(len(all_neuron_scores) * pruning_cfg["target_sparsity"])
1404
+
1405
+ logger.debug(
1406
+ f"Total Neurons: {len(all_neuron_scores)}, Pruning: {num_to_prune}"
1407
+ )
1408
+ neurons_to_prune_info = all_neuron_scores[:num_to_prune]
1409
+
1410
+ for score, (node_name, neuron_idx) in neurons_to_prune_info:
1411
+ mlp_match = re.match(r"m(\d+)", node_name)
1412
+ attn_match = re.match(r"a(\d+)\.h(\d+)", node_name)
1413
+ if mlp_match:
1414
+ pruned_mlp_neurons[int(mlp_match.group(1))].append(neuron_idx)
1415
+ elif attn_match:
1416
+ pruned_attn_neurons[
1417
+ (int(attn_match.group(1)), int(attn_match.group(2)))
1418
+ ].append(neuron_idx)
1419
+
1420
+ result = {
1421
+ "mlp": dict(pruned_mlp_neurons),
1422
+ "heads": dict(pruned_attn_neurons),
1423
+ "_meta": {
1424
+ "mlp_hook": discovery_cfg.get("mlp_hook", "mlp_out"),
1425
+ "heads_hook": "attn.hook_result", # EAP uses attn.hook_result for heads
1426
+ },
1427
+ }
1428
+ if config.get("output_path"):
1429
+ os.makedirs(os.path.dirname(config["output_path"]), exist_ok=True)
1430
+ t.save(result, config["output_path"])
1431
+ logger.info(f"Neuron pruning dictionary saved to {config['output_path']}")
1432
+ _save_artifact(
1433
+ {
1434
+ "algo": algo,
1435
+ "level": "neuron",
1436
+ "neurons_scores": graph.neurons_scores.cpu(),
1437
+ "total_neurons": len(all_neuron_scores),
1438
+ },
1439
+ config.get("output_path"),
1440
+ "_scores",
1441
+ logger,
1442
+ )
1443
+
1444
+ # `graph` (and its GPU-resident neurons_scores tensor) is
1445
+ # fully consumed at this point - the pruning dict and the
1446
+ # CPU-side scores side-car are already built/saved above,
1447
+ # and nothing below this line reads `graph` again. Free it
1448
+ # before the optional inline evaluation, which loads/uses
1449
+ # its own evaluation dataloaders and graph reconstruction
1450
+ # and does not need this one.
1451
+ del graph
1452
+ if t.cuda.is_available():
1453
+ empty_cache()
1454
+
1455
+ if discovery_cfg.get("evaluate", False):
1456
+ config["_eval_result"] = evaluate_circuit(
1457
+ config,
1458
+ pruned_artifact_path=config.get("output_path"),
1459
+ _model=model,
1460
+ )
1461
+
1462
+ progress.complete(neurons_pruned=num_to_prune)
1463
+ return result
1464
+
1465
+ elif algo == "ibcircuit":
1466
+ from .backends.ibcircuit.trainer import run_ib_discovery as _run_ib
1467
+
1468
+ # Build IBCircuit-format dataloader via TaskSpec.
1469
+ # TaskSpec.build_dataloader is the abstraction boundary:
1470
+ # it knows the task format, we don't need to.
1471
+ dataloader = task_spec.build_dataloader(model, discovery_cfg, device)
1472
+
1473
+ # Forward only the training hyperparameters - no path/save keys.
1474
+ # Saving is api.py's responsibility, not the trainer's.
1475
+ # All defaults come from DEFAULT_CONFIG in utils/config.py (single source of truth)
1476
+ default_discovery = DEFAULT_CONFIG["discovery"]
1477
+ ib_config = {
1478
+ "num_epochs": discovery_cfg.get("num_epochs", default_discovery.get("num_epochs")),
1479
+ "learning_rate": discovery_cfg.get(
1480
+ "learning_rate", default_discovery.get("learning_rate")
1481
+ ),
1482
+ "alpha": discovery_cfg.get("alpha", default_discovery.get("alpha")),
1483
+ "beta": discovery_cfg.get("beta", default_discovery.get("beta")),
1484
+ "alpha_loss": discovery_cfg.get("alpha_loss", default_discovery.get("alpha_loss")),
1485
+ "log_interval": discovery_cfg.get(
1486
+ "log_interval", default_discovery.get("log_interval")
1487
+ ),
1488
+ "scope": discovery_cfg.get("scope", default_discovery.get("scope")),
1489
+ "mask_type": discovery_cfg.get("mask_type", default_discovery.get("mask_type")),
1490
+ "level": discovery_cfg.get("level", default_discovery.get("level")),
1491
+ "mlp_hook": discovery_cfg.get("mlp_hook", default_discovery.get("mlp_hook")),
1492
+ "batch_size": discovery_cfg.get("batch_size", default_discovery.get("batch_size")),
1493
+ }
1494
+
1495
+ _validate_ibcircuit_dataloader(dataloader)
1496
+
1497
+ # Returns {"A{layer}.{head}": float, ...}
1498
+ # Higher score = more important head.
1499
+ node_scores, ib_model = _run_ib(
1500
+ model=model, dataloader=dataloader, config=ib_config, device=device
1501
+ )
1502
+
1503
+ # Save IB model weights and discovery scores (api.py owns persistence)
1504
+ _save_artifact(
1505
+ {
1506
+ "attn_ib_weights": ib_model.attn_ib_weights.state_dict(),
1507
+ "mlp_ib_weights": ib_model.mlp_ib_weights.state_dict(),
1508
+ "scope": ib_model.scope,
1509
+ "batch_size": ib_model.batch_size,
1510
+ "n_layers": ib_model.n_layers,
1511
+ "n_heads": ib_model.n_heads,
1512
+ "mask_type": ib_model.mask_type,
1513
+ "level": ib_model.level,
1514
+ "mlp_hook": ib_model.mlp_hook,
1515
+ },
1516
+ config.get("output_path"),
1517
+ "_ib_weights",
1518
+ logger,
1519
+ )
1520
+
1521
+ # Capture the one attribute this branch still needs below, then
1522
+ # free ib_model - its weights are already saved to disk above,
1523
+ # and nothing after this point references ib_model again.
1524
+ _ib_level = ib_model.level
1525
+ del ib_model
1526
+ if t.cuda.is_available():
1527
+ empty_cache()
1528
+
1529
+ if _ib_level == "neuron":
1530
+ # Neuron-level: convert {IBCircuit_name: tensor} → pruning dict,
1531
+ # save in the same format as EAP neuron so evaluate_circuit works.
1532
+ pruned_mlp_neurons = defaultdict(list)
1533
+ pruned_attn_neurons = defaultdict(list)
1534
+ all_neuron_scores = []
1535
+
1536
+ for ib_name, score_tensor in node_scores.items():
1537
+ attn_match = re.match(r"A(\d+)\.(\d+)$", ib_name)
1538
+ mlp_match = re.match(r"MLP (\d+)$", ib_name)
1539
+ for neuron_idx, score in enumerate(score_tensor):
1540
+ if attn_match:
1541
+ all_neuron_scores.append(
1542
+ (
1543
+ abs(score.item()),
1544
+ (
1545
+ "attn",
1546
+ int(attn_match.group(1)),
1547
+ int(attn_match.group(2)),
1548
+ neuron_idx,
1549
+ ),
1550
+ )
1551
+ )
1552
+ elif mlp_match:
1553
+ all_neuron_scores.append(
1554
+ (
1555
+ abs(score.item()),
1556
+ ("mlp", int(mlp_match.group(1)), None, neuron_idx),
1557
+ )
1558
+ )
1559
+
1560
+ all_neuron_scores.sort(key=lambda x: x[0]) # ascending: lowest = least important
1561
+ n_to_prune = int(len(all_neuron_scores) * pruning_cfg["target_sparsity"])
1562
+
1563
+ logger.info(
1564
+ f"IBCircuit neuron discovery: {len(all_neuron_scores)} total neurons, pruning {n_to_prune}"
1565
+ )
1566
+
1567
+ for _, (kind, layer, head, neuron_idx) in all_neuron_scores[:n_to_prune]:
1568
+ if kind == "mlp":
1569
+ pruned_mlp_neurons[layer].append(neuron_idx)
1570
+ else:
1571
+ pruned_attn_neurons[(layer, head)].append(neuron_idx)
1572
+
1573
+ result = {
1574
+ "mlp": dict(pruned_mlp_neurons),
1575
+ "heads": dict(pruned_attn_neurons),
1576
+ "_meta": {"mlp_hook": ib_config["mlp_hook"]},
1577
+ }
1578
+
1579
+ _save_artifact(
1580
+ {
1581
+ "algo": algo,
1582
+ "level": "neuron",
1583
+ "neurons_scores": node_scores,
1584
+ "total_neurons": len(all_neuron_scores),
1585
+ },
1586
+ config.get("output_path"),
1587
+ "_scores",
1588
+ logger,
1589
+ )
1590
+
1591
+ output_path = config.get("output_path")
1592
+ if output_path:
1593
+ parent = os.path.dirname(output_path)
1594
+ if parent:
1595
+ os.makedirs(parent, exist_ok=True)
1596
+ t.save(result, output_path)
1597
+ logger.info(f"Neuron pruning dict saved to {output_path}")
1598
+
1599
+ if discovery_cfg.get("evaluate", False):
1600
+ config["_eval_result"] = evaluate_circuit(
1601
+ config,
1602
+ pruned_artifact_path=output_path,
1603
+ _model=model,
1604
+ )
1605
+
1606
+ progress.complete(neurons_pruned=n_to_prune)
1607
+ return result
1608
+
1609
+ else:
1610
+ # Node-level: existing behaviour
1611
+ # Build unified CircuitScores artifact (Workstream G)
1612
+ circuit_scores = _build_circuit_scores(
1613
+ task=discovery_cfg["task"],
1614
+ model_name=model_cfg["name"],
1615
+ algorithm=algo,
1616
+ node_scores=node_scores,
1617
+ discovery_cfg=discovery_cfg,
1618
+ )
1619
+
1620
+ # Save CircuitScores as JSON
1621
+ if config.get("output_path"):
1622
+ scores_path = Path(config["output_path"]).parent / (
1623
+ Path(config["output_path"]).stem + "_scores.json"
1624
+ )
1625
+ circuit_scores.to_json(scores_path)
1626
+ logger.info(f"Saved unified CircuitScores → {scores_path}")
1627
+
1628
+ # Also save legacy format for compatibility
1629
+ _save_artifact(
1630
+ {"algo": algo, "level": "node", "node_scores": node_scores},
1631
+ config.get("output_path"),
1632
+ "_scores",
1633
+ logger,
1634
+ )
1635
+
1636
+ elif algo == "cdt":
1637
+ with debug_context("CD-T Discovery"):
1638
+
1639
+ if discovery_cfg.get("level") == "neuron":
1640
+ raise ValueError(
1641
+ "CD-T only supports node-level discovery in the current version. "
1642
+ "Set discovery config key 'level' to 'node', or choose an "
1643
+ "algorithm that supports neuron-level (e.g. eap, eap-ig, ibcircuit)."
1644
+ )
1645
+
1646
+ from .backends.cdt.adapter import run_cdt_discovery
1647
+
1648
+ dataloader = task_spec.build_dataloader(model, discovery_cfg, device)
1649
+ node_scores = run_cdt_discovery(
1650
+ tl_model=model,
1651
+ dataloader=dataloader,
1652
+ device=device,
1653
+ n_examples=discovery_cfg.get("data_params", {}).get("num_examples", 16),
1654
+ )
1655
+
1656
+ circuit_scores = _build_circuit_scores(
1657
+ task=discovery_cfg["task"],
1658
+ model_name=model_cfg["name"],
1659
+ algorithm=algo,
1660
+ node_scores=node_scores,
1661
+ discovery_cfg=discovery_cfg,
1662
+ )
1663
+ if config.get("output_path"):
1664
+ scores_path = Path(config["output_path"]).parent / (
1665
+ Path(config["output_path"]).stem + "_scores.json"
1666
+ )
1667
+ circuit_scores.to_json(scores_path)
1668
+ _save_artifact(
1669
+ {"algo": algo, "level": "node", "node_scores": node_scores},
1670
+ config.get("output_path"),
1671
+ "_scores",
1672
+ logger,
1673
+ )
1674
+ else:
1675
+ from .backends import DISCOVERY_ALGORITHMS
1676
+
1677
+ raise AlgorithmError(
1678
+ f"Unknown discovery algorithm '{algo}'. Set the discovery config "
1679
+ f"key 'algorithm' to one of: {sorted(DISCOVERY_ALGORITHMS)}."
1680
+ )
1681
+
1682
+ progress.step("Identifying nodes to prune")
1683
+
1684
+ effective_scope = _ib_scope if algo == "ibcircuit" else pruning_cfg.get("scope", "both")
1685
+ nodes_to_prune = get_nodes_to_prune(
1686
+ node_scores,
1687
+ target_sparsity=pruning_cfg["target_sparsity"],
1688
+ pruning_scope=effective_scope,
1689
+ )
1690
+
1691
+ logger.info(f" Pruned {len(nodes_to_prune)} nodes out of {len(node_scores)} candidates")
1692
+
1693
+ if config.get("output_path"):
1694
+ parent = os.path.dirname(config["output_path"])
1695
+ if parent:
1696
+ os.makedirs(parent, exist_ok=True)
1697
+ t.save(nodes_to_prune, config["output_path"])
1698
+ logger.info(f"Pruned node list saved to {config['output_path']}")
1699
+
1700
+ if discovery_cfg.get("evaluate", False):
1701
+ config["_eval_result"] = evaluate_circuit(
1702
+ config,
1703
+ pruned_artifact_path=config.get("output_path"),
1704
+ _model=model,
1705
+ )
1706
+
1707
+ progress.complete(nodes_pruned=len(nodes_to_prune))
1708
+ return nodes_to_prune
1709
+
1710
+ except Exception as e:
1711
+ progress.fail(str(e))
1712
+ raise
1713
+ finally:
1714
+ # Restore the caller's global RNG so a seeded discovery run doesn't
1715
+ # leak its deterministic RNG state into the surrounding process.
1716
+ if _rng_snapshot is not None:
1717
+ import random as _random_std
1718
+ import numpy as _np_std
1719
+ t.set_rng_state(_rng_snapshot[0])
1720
+ _np_std.random.set_state(_rng_snapshot[1])
1721
+ _random_std.setstate(_rng_snapshot[2])
1722
+ if _rng_snapshot[3] is not None:
1723
+ t.cuda.set_rng_state_all(_rng_snapshot[3])
1724
+
1725
+
1726
+ def _save_evaluation_results_to_txt(
1727
+ evaluation_results: List[Dict[str, Any]],
1728
+ model_name: str,
1729
+ pruned_artifact_path: str,
1730
+ evaluation_mode: str,
1731
+ logger,
1732
+ custom_path: str = None,
1733
+ ) -> str:
1734
+ """
1735
+ Write lm-eval benchmark results to a plain-text file.
1736
+
1737
+ Each entry in evaluation_results is rendered as a labelled block. On
1738
+ failure the error is logged and an empty string is returned rather than
1739
+ propagating the exception.
1740
+
1741
+ Args:
1742
+ evaluation_results (List[Dict]): List of result dicts, each with keys
1743
+ 'model_type' (str: 'original' | 'pruned') and 'results' (dict).
1744
+ model_name (str): HuggingFace model identifier, used in the filename
1745
+ when custom_path is not provided.
1746
+ pruned_artifact_path (str): Path to the pruning artifact, recorded in
1747
+ the file header for traceability.
1748
+ evaluation_mode (str): Evaluation mode label ('both', 'original', 'pruned').
1749
+ logger: Logger instance for info/error messages.
1750
+ custom_path (Optional[str]): Explicit output file path. If None, a
1751
+ timestamped file is created in the current working directory.
1752
+
1753
+ Returns:
1754
+ str: Absolute path to the written file, or '' on failure.
1755
+ """
1756
+ try:
1757
+ if custom_path:
1758
+ file_path = custom_path
1759
+ else:
1760
+ # Generate timestamp for unique filename
1761
+ timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
1762
+ filename = f"evaluation_results_{model_name.replace('/', '_')}_{timestamp}.txt"
1763
+ # Create the file path
1764
+ file_path = os.path.join(os.getcwd(), filename)
1765
+
1766
+ with open(file_path, "w", encoding="utf-8") as f:
1767
+ f.write("=" * 80 + "\n")
1768
+ f.write("CIRCUITKIT EVALUATION RESULTS\n")
1769
+ f.write("=" * 80 + "\n\n")
1770
+
1771
+ # Write metadata
1772
+ f.write(f"Model: {model_name}\n")
1773
+ f.write(f"Pruned Artifact: {pruned_artifact_path}\n")
1774
+ f.write(f"Evaluation Mode: {evaluation_mode}\n")
1775
+ f.write(f"Timestamp: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n")
1776
+ f.write("Generated by: CircuitKit\n\n")
1777
+
1778
+ # Write evaluation results
1779
+ for i, result in enumerate(evaluation_results, 1):
1780
+ f.write("-" * 60 + "\n")
1781
+ f.write(f"EVALUATION {i}: {result['model_type'].upper()} MODEL\n")
1782
+ f.write("-" * 60 + "\n\n")
1783
+
1784
+ # Format and write the results
1785
+ results_data = result["results"]
1786
+ if isinstance(results_data, dict):
1787
+ for task, score in results_data.items():
1788
+ if isinstance(score, dict):
1789
+ f.write(f"Task: {task}\n")
1790
+ for metric, value in score.items():
1791
+ f.write(f" {metric}: {value}\n")
1792
+ f.write("\n")
1793
+ else:
1794
+ f.write(f"{task}: {score}\n")
1795
+ else:
1796
+ f.write(f"Results: {results_data}\n")
1797
+
1798
+ f.write("\n")
1799
+
1800
+ f.write("=" * 80 + "\n")
1801
+ f.write("END OF EVALUATION RESULTS\n")
1802
+ f.write("=" * 80 + "\n")
1803
+
1804
+ logger.info(f"Evaluation results saved to: {file_path}")
1805
+ return file_path
1806
+
1807
+ except Exception as e:
1808
+ logger.error(f"Failed to save evaluation results to txt file: {e}")
1809
+ return ""
1810
+
1811
+
1812
+ def _reconstruct_circuit_graph(
1813
+ model: HookedTransformer,
1814
+ scores_data: Dict,
1815
+ discovery_cfg: Dict[str, Any],
1816
+ pruning_cfg: Dict[str, Any],
1817
+ device: str,
1818
+ ) -> Graph:
1819
+ """
1820
+ Helper: Reconstruct pruned graph from scores.
1821
+
1822
+ Handles all algorithms (ACDC, EAP, IBCircuit) and levels (node, neuron).
1823
+ Returns the reconstructed circuit graph with topn applied.
1824
+ """
1825
+ algo = discovery_cfg["algorithm"].lower()
1826
+ level = discovery_cfg.get("level", "node")
1827
+ scope = (
1828
+ discovery_cfg.get("scope", "heads")
1829
+ if algo == "ibcircuit"
1830
+ else pruning_cfg.get("scope", "both")
1831
+ )
1832
+ sparsity = pruning_cfg.get("target_sparsity", 0.0)
1833
+
1834
+ is_neuron = level == "neuron"
1835
+
1836
+ if is_neuron:
1837
+ # EAP neuron path
1838
+ mlp_hook = discovery_cfg.get("mlp_hook", "mlp_out")
1839
+ graph = Graph.from_model(model, node_scores=True, neuron_level=True, mlp_hook=mlp_hook)
1840
+ graph.neurons_scores = scores_data["neurons_scores"].to(device)
1841
+ if graph.neurons_scores.shape[1] > graph.neurons_in_graph.shape[1]:
1842
+ graph.neurons_scores = graph.neurons_scores[:, : graph.neurons_in_graph.shape[1]]
1843
+
1844
+ total_to_keep_global = 0
1845
+ for node in graph.nodes.values():
1846
+ if isinstance(node, (AttentionNode, MLPNode)):
1847
+ out_of_scope = (scope == "heads" and isinstance(node, MLPNode)) or (
1848
+ scope == "mlp" and isinstance(node, AttentionNode)
1849
+ )
1850
+ fwd_idx = graph.forward_index(node, attn_slice=False)
1851
+ if out_of_scope:
1852
+ graph.neurons_scores[fwd_idx] = float("inf")
1853
+ total_to_keep_global += node.d_neuron
1854
+ else:
1855
+ num_keep_local = int(node.d_neuron * (1 - sparsity))
1856
+ total_to_keep_global += num_keep_local
1857
+ abs_scores = t.abs(graph.neurons_scores[fwd_idx, : node.d_neuron])
1858
+ if num_keep_local < node.d_neuron and num_keep_local > 0:
1859
+ # Keep exactly num_keep_local neurons by index. A
1860
+ # threshold compare (abs_scores < kth-largest) over-keeps
1861
+ # every neuron tied at the boundary, drifting the
1862
+ # effective sparsity below the requested target; topk
1863
+ # indices break ties deterministically.
1864
+ keep_idx = t.topk(abs_scores, num_keep_local).indices
1865
+ prune_mask = t.ones_like(abs_scores, dtype=t.bool)
1866
+ prune_mask[keep_idx] = False
1867
+ graph.neurons_scores[fwd_idx, : node.d_neuron][prune_mask] = -float("inf")
1868
+ elif num_keep_local == 0:
1869
+ graph.neurons_scores[fwd_idx, : node.d_neuron] = -float("inf")
1870
+
1871
+ graph.apply_topn(total_to_keep_global, level="neuron", prune=True)
1872
+ else:
1873
+ # Node-level (ACDC, EAP)
1874
+ graph = Graph.from_model(
1875
+ model,
1876
+ node_scores=True,
1877
+ neuron_level=False,
1878
+ mlp_hook=discovery_cfg.get("mlp_hook", "mlp_out"),
1879
+ )
1880
+ _populate_graph_from_ib_scores(graph, scores_data["node_scores"])
1881
+ for node in graph.nodes.values():
1882
+ if isinstance(node, (AttentionNode, MLPNode)):
1883
+ fwd_idx = graph.forward_index(node, attn_slice=False)
1884
+ out_of_scope = (scope == "heads" and isinstance(node, MLPNode)) or (
1885
+ scope == "mlp" and isinstance(node, AttentionNode)
1886
+ )
1887
+ if out_of_scope:
1888
+ node.score = t.tensor(float("inf"))
1889
+ graph.nodes_scores[fwd_idx] = float("inf")
1890
+ n_topn, n_to_keep = _compute_n_topn(graph, scope, sparsity)
1891
+ graph.apply_topn(n_topn, level="node", prune=True)
1892
+
1893
+ return graph
1894
+
1895
+
1896
+ @debug_function
1897
+ @handle_errors(context={"operation": "evaluate_circuit"})
1898
+ def evaluate_circuit(
1899
+ config: Union[str, Dict[str, Any]],
1900
+ pruned_artifact_path: str = None,
1901
+ scores_path: str = None,
1902
+ _model: Optional[HookedTransformer] = None,
1903
+ ) -> "FaithfulnessReport":
1904
+ """
1905
+ Evaluate circuit faithfulness using the 6-pillar framework.
1906
+
1907
+ Thin wrapper around run_full_faithfulness(). Reconstructs the circuit
1908
+ graph from saved scores, loads the model, and runs comprehensive
1909
+ faithfulness evaluation via run_full_faithfulness().
1910
+
1911
+ Args:
1912
+ config: Path to YAML config or config dict.
1913
+ pruned_artifact_path: Path to .pt pruning artifact (defaults to config['output_path']).
1914
+ scores_path: Path to _scores.pt file (auto-derived if not provided).
1915
+ _model: Optional pre-loaded HookedTransformer. Internal parameter used
1916
+ by discover_circuit() to avoid loading the model a second time
1917
+ when discovery_cfg["evaluate"]=True triggers an inline evaluation.
1918
+ When provided, this function still unconditionally (re-)asserts
1919
+ the config flags it needs (use_split_qkv_input, use_attn_result,
1920
+ use_hook_mlp_in, ungroup_grouped_query_attention) on it before
1921
+ use, since discover_circuit only sets these for EAP-family
1922
+ algorithms — algorithms like 'ibcircuit'/'cdt' may hand over a
1923
+ model that doesn't have them yet. Setting an already-true flag
1924
+ is a no-op, so this is safe either way. External callers should
1925
+ leave this as None; behavior is identical to before this
1926
+ parameter existed.
1927
+
1928
+ Returns:
1929
+ FaithfulnessReport: Structured evaluation result. The two always-present
1930
+ fields are:
1931
+ - ``.patching_score``: Pillar 1 (causal patching) — original vs
1932
+ circuit performance under intervention.
1933
+ - ``.ablation_score``: Pillar 2 (ablation) — circuit sufficiency.
1934
+ The full-faithfulness path additionally populates ``.stability``,
1935
+ ``.robustness``, ``.baseline_comparison``, ``.generalization`` and
1936
+ ``.intervention_reliability``. A random-circuit baseline, when
1937
+ requested, is carried in ``.metadata["random_avg"]``.
1938
+
1939
+ Prior to 1.0 the fast path returned a dict with the misleadingly
1940
+ named keys ``baseline_avg`` (= patching), ``circuit_avg``
1941
+ (= ablation) and ``random_avg``; that dict has been removed. Use the
1942
+ attributes above.
1943
+ """
1944
+ from pathlib import Path
1945
+
1946
+ from .evaluation import run_full_faithfulness
1947
+ from .tasks.bootstrap import _bootstrap_builtin_tasks
1948
+ from .utils.config import load_and_validate_config
1949
+
1950
+ _bootstrap_builtin_tasks()
1951
+ logger = get_logger("circuitkit.evaluate_circuit")
1952
+ progress = ProgressLogger(logger)
1953
+
1954
+ try:
1955
+ progress.start_operation("Circuit Evaluation", 4)
1956
+ progress.step("Loading config and model")
1957
+
1958
+ config = load_and_validate_config(config)
1959
+ discovery_cfg = config["discovery"]
1960
+ pruning_cfg = config["pruning"]
1961
+
1962
+ # Resolve paths
1963
+ artifact_path = pruned_artifact_path or config.get("output_path")
1964
+ if not artifact_path:
1965
+ raise ValueError("Provide pruned_artifact_path or set config['output_path']")
1966
+ if scores_path is None:
1967
+ scores_path = str(
1968
+ Path(artifact_path).parent / (Path(artifact_path).stem + "_scores.pt")
1969
+ )
1970
+
1971
+ validate_file_exists(artifact_path, "pruned artifact")
1972
+ validate_file_exists(scores_path, "discovery scores")
1973
+
1974
+ # Load model and data
1975
+ device = get_device()
1976
+ dtype = getattr(t, config["model"].get("precision", "bfloat16"))
1977
+ if _model is not None:
1978
+ # Reuse the caller's already-loaded model (e.g. discover_circuit's
1979
+ # inline evaluate path) instead of loading a second full copy.
1980
+ model = _model
1981
+ logger.debug("evaluate_circuit: reusing pre-loaded model, skipping reload")
1982
+ else:
1983
+ with log_execution_time("Model loading", logger):
1984
+ model = HookedTransformer.from_pretrained(
1985
+ config["model"]["name"], device=device, dtype=dtype
1986
+ )
1987
+ # These flags are required by the graph reconstruction / faithfulness
1988
+ # evaluation below regardless of model provenance. discover_circuit()
1989
+ # only sets them for the EAP-family algorithms (see its algo-dispatch
1990
+ # block); algorithms like 'ibcircuit'/'cdt' reach this function with
1991
+ # a model that may not have them set yet. Setting an already-true
1992
+ # flag is a no-op, so it's always safe to (re-)assert these here,
1993
+ # whether `model` was just loaded or reused via `_model`.
1994
+ model.cfg.use_split_qkv_input = True
1995
+ model.cfg.use_attn_result = True
1996
+ model.cfg.use_hook_mlp_in = True
1997
+ if hasattr(model.cfg, "ungroup_grouped_query_attention"):
1998
+ model.cfg.ungroup_grouped_query_attention = True
1999
+
2000
+ scores_data = t.load(scores_path, map_location="cpu", weights_only=True)
2001
+ task_spec = _get_task(discovery_cfg["task"])
2002
+
2003
+ # Build evaluation dataloader
2004
+ eval_cfg = config.get("eval", {})
2005
+ eval_num_examples = eval_cfg.get(
2006
+ "num_examples", discovery_cfg.get("data_params", {}).get("num_examples", 256)
2007
+ )
2008
+ eval_seed = eval_cfg.get(
2009
+ "seed",
2010
+ discovery_cfg.get(
2011
+ "seed", # top-level seed (WMDP, MMLU style)
2012
+ discovery_cfg.get("data_params", {}).get("seed", 42), # nested seed (IOI style)
2013
+ ),
2014
+ )
2015
+ dl_cfg = {
2016
+ **discovery_cfg,
2017
+ "algorithm": "eap",
2018
+ "data_params": {
2019
+ **discovery_cfg.get("data_params", {}),
2020
+ "num_examples": eval_num_examples,
2021
+ "seed": eval_seed,
2022
+ },
2023
+ "batch_size": discovery_cfg.get("batch_size", 16),
2024
+ }
2025
+
2026
+ algo = discovery_cfg["algorithm"].lower()
2027
+ level = discovery_cfg.get("level", "node")
2028
+
2029
+ # Detect clean-only IBCircuit neuron-level (custom data with no corrupt
2030
+ # prompts). The EAP dataloader path would crash because it requires
2031
+ # fully-paired data. Instead, build a self-paired EAP-format loader and
2032
+ # switch to the correct-token probability metric (bounded [0, 1]) which
2033
+ # doesn't need an incorrect token.
2034
+ _is_clean_only_ib = (
2035
+ algo == "ibcircuit"
2036
+ and level == "neuron"
2037
+ and hasattr(task_spec, "ds")
2038
+ and not getattr(task_spec.ds, "fully_paired", True)
2039
+ )
2040
+ if _is_clean_only_ib:
2041
+ dataloader = _build_clean_only_ib_eval_dataloader(
2042
+ task_spec,
2043
+ model,
2044
+ num_examples=eval_num_examples,
2045
+ batch_size=int(discovery_cfg.get("batch_size", 8)),
2046
+ )
2047
+ else:
2048
+ dataloader = task_spec.build_dataloader(model, dl_cfg, device)
2049
+ # eval_cfg was already set above (line 1294); use_full_faithfulness_eval resolved here.
2050
+ # For clean-only IBCircuit neuron-level we always force the fast path: only
2051
+ # sufficiency (baseline avg vs circuit avg) is computable without paired data.
2052
+ use_full_faithfulness_eval = eval_cfg.get("full_faithfulness_eval", False)
2053
+ if _is_clean_only_ib:
2054
+ use_full_faithfulness_eval = False
2055
+
2056
+ # ── Shared setup for ALL algorithms ───────────────────────────────────
2057
+ # Must live here — before the IBCircuit branch — so every path has
2058
+ # access to eval_intervention, intervention_dataloader, corruption
2059
+ # dataloaders, and target task. EAP/EAP-IG behaviour is unchanged:
2060
+ # they fall through this block and hit _reconstruct_circuit_graph below.
2061
+
2062
+ eval_intervention = pruning_cfg.get("intervention", "zero")
2063
+ discovery_cfg["eval_intervention"] = (
2064
+ eval_intervention # read by run_full_faithfulness pillars 2/4/6
2065
+ )
2066
+
2067
+ intervention_dataloader = None
2068
+ if eval_intervention in ("mean", "mean-positional"):
2069
+ if _is_clean_only_ib:
2070
+ intervention_dataloader = dataloader
2071
+ else:
2072
+ intervention_dataloader = task_spec.build_dataloader(model, dl_cfg, device)
2073
+
2074
+ corruption_dataloaders = {}
2075
+ pillars_to_run = eval_cfg.get("pillars")
2076
+ if pillars_to_run is None or "robustness" in pillars_to_run:
2077
+ corruption_variants = eval_cfg.get("corruption_variants", ["paraphrase"])
2078
+ for variant in corruption_variants:
2079
+ var_cfg = dl_cfg.copy()
2080
+ var_cfg["data_params"] = var_cfg.get("data_params", {}).copy()
2081
+ var_cfg["data_params"]["corruption_variant"] = variant
2082
+ try:
2083
+ corruption_dataloaders[variant] = task_spec.build_dataloader(
2084
+ model, var_cfg, device
2085
+ )
2086
+ except Exception as e:
2087
+ logger.warning(f"Could not build corruption dataloader for '{variant}': {e}")
2088
+
2089
+ if not corruption_dataloaders and (
2090
+ pillars_to_run is None or "robustness" in pillars_to_run
2091
+ ):
2092
+ logger.error(
2093
+ f"Robustness pillar requested but no corruption dataloaders could be built "
2094
+ f"for variants {corruption_variants}. Skipping robustness evaluation."
2095
+ )
2096
+ # Remove 'robustness' from pillars to prevent meaningless zero-delta results
2097
+ if pillars_to_run:
2098
+ pillars_to_run = [p for p in pillars_to_run if p != "robustness"]
2099
+ target_task_name = eval_cfg.get("target_task", None)
2100
+ target_task_spec = None
2101
+ target_dataloader = None
2102
+ if target_task_name is not None:
2103
+ target_task_spec = _get_task(target_task_name)
2104
+ target_configs = eval_cfg.get("target_configs")
2105
+ if target_configs is not None:
2106
+ target_dl_cfg = {**dl_cfg, "configs": target_configs}
2107
+ else:
2108
+ target_dl_cfg = dl_cfg
2109
+ target_dataloader = target_task_spec.build_dataloader(model, target_dl_cfg, device)
2110
+ logger.info(
2111
+ f"Target task for generalization: {target_task_name}"
2112
+ + (f" (configs: {target_configs})" if target_configs else "")
2113
+ )
2114
+
2115
+ # For clean-only IBCircuit neuron-level, substitute a metric that
2116
+ # doesn't need an incorrect token. _correct_token_prob is bounded [0, 1]
2117
+ # and uses only labels[:, 0] (the correct-answer token).
2118
+ if _is_clean_only_ib:
2119
+ from functools import partial as _partial
2120
+
2121
+ metric = _partial(_correct_token_prob, loss=False, mean=False)
2122
+ else:
2123
+ metric = _make_eval_metric(task_spec)
2124
+
2125
+ # ── IBCircuit neuron-level: special graph construction and eval ────────
2126
+ # Cannot use _reconstruct_circuit_graph because IBCircuit neuron scores
2127
+ # are stored as Dict[str, Tensor] in _scores.pt, incompatible with the
2128
+ # 2-D Tensor that the neuron branch of _reconstruct_circuit_graph expects.
2129
+ # Instead the graph is built directly from the pruning dict artifact.
2130
+ if level == "neuron" and algo == "ibcircuit":
2131
+ from .evaluation.evaluate import evaluate_baseline, evaluate_ibcircuit_neuron_circuit
2132
+
2133
+ scope = discovery_cfg.get("scope", "heads")
2134
+ seed = eval_cfg.get("seed", discovery_cfg.get("data_params", {}).get("seed", 42))
2135
+
2136
+ pruning_dict = t.load(artifact_path, map_location=device, weights_only=True)
2137
+
2138
+ _log_gpu_mem("api.evaluate_circuit: after model+data load, before P1/P2", logger)
2139
+
2140
+ # IBCircuit training uses mean-ablation; honour pruning_cfg but default to 'mean'.
2141
+ # For clean-only data, patching is unavailable (no corrupt side); keep mean/zero.
2142
+ _ib_intervention = pruning_cfg.get("intervention", "mean")
2143
+ ib_eval_intervention = (
2144
+ _ib_intervention if _ib_intervention in ("zero", "patching") else "mean"
2145
+ )
2146
+ if _is_clean_only_ib and ib_eval_intervention == "patching":
2147
+ ib_eval_intervention = "mean" # patching needs a real corrupt side
2148
+
2149
+ progress.step("Running IBCircuit neuron evaluation")
2150
+
2151
+ # Pillars 1 & 2 — computed via IBCircuit-specific evaluators.
2152
+ # evaluate_graph (used inside run_full_faithfulness for P1/P2) ablates
2153
+ # via activation-difference hooks which are incompatible with the
2154
+ # per-neuron hook mechanism of evaluate_ibcircuit_neuron_circuit.
2155
+ baseline_avg = _avg_scores(evaluate_baseline(model, dataloader, metric))
2156
+ circuit_avg = _avg_scores(
2157
+ evaluate_ibcircuit_neuron_circuit(
2158
+ model,
2159
+ pruning_dict,
2160
+ dataloader,
2161
+ metric,
2162
+ intervention=ib_eval_intervention,
2163
+ )
2164
+ )
2165
+
2166
+ _log_gpu_mem("api.evaluate_circuit: after IBCircuit P1/P2 eval", logger)
2167
+
2168
+ random_avg = None
2169
+ if pruning_cfg.get("random", False):
2170
+ rand_pruning_dict = _build_random_ibcircuit_neuron_pruning_dict(
2171
+ model,
2172
+ pruning_dict,
2173
+ scope=scope,
2174
+ seed=seed,
2175
+ )
2176
+ random_avg = _avg_scores(
2177
+ evaluate_ibcircuit_neuron_circuit(
2178
+ model,
2179
+ rand_pruning_dict,
2180
+ dataloader,
2181
+ metric,
2182
+ intervention=ib_eval_intervention,
2183
+ )
2184
+ )
2185
+
2186
+ if not use_full_faithfulness_eval:
2187
+ # Fast path: P1/P2 only, returned as a FaithfulnessReport. The
2188
+ # random-circuit baseline (when computed) is carried in metadata.
2189
+ # For clean-only IBCircuit: patching_score = full-model correct-token
2190
+ # probability (baseline), ablation_score = circuit sufficiency score.
2191
+ from .evaluation.report import FaithfulnessReport
2192
+
2193
+ report = FaithfulnessReport(
2194
+ patching_score=baseline_avg,
2195
+ ablation_score=circuit_avg,
2196
+ )
2197
+ report.metadata = {"random_avg": random_avg} if random_avg is not None else {}
2198
+ if _is_clean_only_ib:
2199
+ report.metadata["eval_mode"] = "clean_only_sufficiency"
2200
+ logger.info(
2201
+ f"Clean-only sufficiency: full-model P(correct)={_fmt_opt_score(baseline_avg)} "
2202
+ f"| circuit P(correct)={_fmt_opt_score(circuit_avg)}"
2203
+ )
2204
+ else:
2205
+ logger.info(f"Original: {_fmt_opt_score(baseline_avg)} | Circuit: {_fmt_opt_score(circuit_avg)}")
2206
+ progress.complete(
2207
+ **{
2208
+ k: round(v, 4)
2209
+ for k, v in {"patching_score": baseline_avg, "ablation_score": circuit_avg}.items()
2210
+ if v is not None
2211
+ }
2212
+ )
2213
+ return report
2214
+
2215
+ # Full faithfulness path — build a proper neuron-level Graph from the
2216
+ # pruning dict so that graph-based pillars (baselines, robustness,
2217
+ # stability, generalization) receive a correctly populated graph.
2218
+ # neurons_in_graph defaults to all-ones (all in circuit); we zero
2219
+ # out the pruned neurons to match the IBCircuit discovery result.
2220
+ mlp_hook = discovery_cfg.get("mlp_hook", "mlp_out")
2221
+ graph = Graph.from_model(model, node_scores=True, neuron_level=True, mlp_hook=mlp_hook)
2222
+
2223
+ _log_gpu_mem("api.evaluate_circuit: after neuron-level Graph construction", logger)
2224
+
2225
+ ib_mlp_neurons = pruning_dict.get("mlp", {}) # {layer: [neuron_idx, ...]}
2226
+ ib_attn_neurons = pruning_dict.get("heads", {}) # {(layer, head): [neuron_idx, ...]}
2227
+
2228
+ for node in graph.nodes.values():
2229
+ if isinstance(node, MLPNode):
2230
+ pruned = ib_mlp_neurons.get(node.layer, [])
2231
+ if pruned:
2232
+ fwd_idx = graph.forward_index(node, attn_slice=False)
2233
+ graph.neurons_in_graph[fwd_idx, pruned] = 0
2234
+ elif isinstance(node, AttentionNode):
2235
+ pruned = ib_attn_neurons.get((node.layer, node.head), [])
2236
+ if pruned:
2237
+ fwd_idx = graph.forward_index(node, attn_slice=False)
2238
+ graph.neurons_in_graph[fwd_idx, pruned] = 0
2239
+
2240
+ # Run remaining pillars (baselines, robustness, stability,
2241
+ # generalization) via run_full_faithfulness. patching and ablation
2242
+ # are intentionally excluded — they were computed above with the
2243
+ # IBCircuit-correct evaluators.
2244
+ requested_pillars = eval_cfg.get("pillars") or [
2245
+ "patching",
2246
+ "ablation",
2247
+ "baselines",
2248
+ "robustness",
2249
+ "stability",
2250
+ "generalization",
2251
+ ]
2252
+ graph_pillars = [p for p in requested_pillars if p not in ("patching", "ablation")]
2253
+
2254
+ import gc
2255
+
2256
+ gc.collect()
2257
+ if t.cuda.is_available():
2258
+ empty_cache()
2259
+
2260
+ extra_report = None
2261
+
2262
+ _log_gpu_mem("api.evaluate_circuit: before run_full_faithfulness", logger)
2263
+
2264
+ if graph_pillars:
2265
+ extra_report = run_full_faithfulness(
2266
+ model=model,
2267
+ graph=graph,
2268
+ task_spec=task_spec,
2269
+ discovery_cfg=discovery_cfg,
2270
+ pruning_cfg=pruning_cfg,
2271
+ device=device,
2272
+ pillars=graph_pillars,
2273
+ n_stability_runs=eval_cfg.get("n_stability_runs", 5),
2274
+ metric_fn=metric,
2275
+ dataloader=dataloader,
2276
+ intervention_dataloader=intervention_dataloader,
2277
+ corruption_dataloaders=corruption_dataloaders,
2278
+ target_task_spec=target_task_spec,
2279
+ target_dataloader=target_dataloader,
2280
+ )
2281
+
2282
+ # Assemble FaithfulnessReport: P1/P2 from IBCircuit evaluators,
2283
+ # remaining pillars from extra_report (if computed).
2284
+ from .evaluation.report import FaithfulnessReport
2285
+
2286
+ report = FaithfulnessReport(
2287
+ patching_score=baseline_avg,
2288
+ ablation_score=circuit_avg,
2289
+ )
2290
+ if extra_report is not None:
2291
+ report.baseline_comparison = getattr(extra_report, "baseline_comparison", None)
2292
+ report.robustness = getattr(extra_report, "robustness", None)
2293
+ report.stability = getattr(extra_report, "stability", None)
2294
+ report.generalization = getattr(extra_report, "generalization", None)
2295
+ # Carry over metadata set by run_full_faithfulness; patch in our values
2296
+ report.metadata = getattr(extra_report, "metadata", {})
2297
+ else:
2298
+ report.metadata = {}
2299
+
2300
+ report.metadata.update(
2301
+ {
2302
+ "algorithm": algo,
2303
+ "model": config["model"]["name"],
2304
+ "task": discovery_cfg.get("task", "unknown"),
2305
+ "level": level,
2306
+ "scope": scope,
2307
+ "sparsity": pruning_cfg.get("target_sparsity", 0.0),
2308
+ "pillars_computed": requested_pillars,
2309
+ "random_avg": random_avg,
2310
+ }
2311
+ )
2312
+
2313
+ logger.info(f"Original: {_fmt_opt_score(baseline_avg)} | Circuit: {_fmt_opt_score(circuit_avg)}")
2314
+ logger.info("Full faithfulness report complete (IBCircuit neuron)")
2315
+ progress.complete()
2316
+ return report
2317
+
2318
+ # ── All other algorithms (EAP, EAP-IG, ACDC, IBCircuit node-level) ────
2319
+ # Reconstruct graph from scores and run faithfulness evaluation.
2320
+ # This path is identical to the original code — no changes.
2321
+ progress.step("Reconstructing circuit")
2322
+ graph = _reconstruct_circuit_graph(model, scores_data, discovery_cfg, pruning_cfg, device)
2323
+
2324
+ progress.step("Running faithfulness evaluation")
2325
+
2326
+ if use_full_faithfulness_eval:
2327
+ report = run_full_faithfulness(
2328
+ model=model,
2329
+ graph=graph,
2330
+ task_spec=task_spec,
2331
+ discovery_cfg=discovery_cfg,
2332
+ pruning_cfg=pruning_cfg,
2333
+ device=device,
2334
+ pillars=eval_cfg.get("pillars", None),
2335
+ n_stability_runs=eval_cfg.get("n_stability_runs", 5),
2336
+ metric_fn=metric,
2337
+ dataloader=dataloader,
2338
+ intervention_dataloader=intervention_dataloader,
2339
+ corruption_dataloaders=corruption_dataloaders,
2340
+ baseline_types=eval_cfg.get("baseline_types", None),
2341
+ target_task_spec=target_task_spec,
2342
+ target_dataloader=target_dataloader,
2343
+ )
2344
+ logger.info("Full faithfulness report complete")
2345
+ progress.complete()
2346
+ return report
2347
+ else:
2348
+ report = run_full_faithfulness(
2349
+ model=model,
2350
+ graph=graph,
2351
+ task_spec=task_spec,
2352
+ discovery_cfg=discovery_cfg,
2353
+ pruning_cfg=pruning_cfg,
2354
+ device=device,
2355
+ pillars=["patching", "ablation"],
2356
+ metric_fn=metric,
2357
+ dataloader=dataloader,
2358
+ intervention_dataloader=intervention_dataloader,
2359
+ target_task_spec=target_task_spec,
2360
+ target_dataloader=target_dataloader,
2361
+ )
2362
+ logger.info(
2363
+ f"Original: {_fmt_opt_score(report.patching_score)} | "
2364
+ f"Circuit: {_fmt_opt_score(report.ablation_score)}"
2365
+ )
2366
+ progress.complete(
2367
+ **{
2368
+ k: round(v, 4)
2369
+ for k, v in {
2370
+ "patching_score": report.patching_score,
2371
+ "ablation_score": report.ablation_score,
2372
+ }.items()
2373
+ if v is not None
2374
+ }
2375
+ )
2376
+ return report
2377
+
2378
+ except Exception as e:
2379
+ progress.fail(str(e))
2380
+ raise
2381
+
2382
+
2383
+ @debug_function
2384
+ @handle_errors(context={"operation": "benchmark_circuit"})
2385
+ def benchmark_circuit(
2386
+ model_name: str,
2387
+ pruned_artifact_path: str,
2388
+ eval_params: Dict[str, Any],
2389
+ config_for_report: Dict[str, Any],
2390
+ report_path: str = None,
2391
+ precision: str = "bfloat16",
2392
+ use_weight_based_pruning: bool = False,
2393
+ evaluation_mode: str = "both",
2394
+ save_to_txt: bool = False,
2395
+ ):
2396
+ """
2397
+ Evaluate a pruned circuit on lm-eval benchmarks.
2398
+
2399
+ Loads the model and pruning artifact, then runs the lm-evaluation-harness
2400
+ on the original and/or pruned model depending on evaluation_mode. Pruning
2401
+ is applied either via forward hooks (default) or by directly zeroing weights
2402
+ (use_weight_based_pruning=True, which avoids the use_attn_result overhead).
2403
+
2404
+ Args:
2405
+ model_name (str): HuggingFace model identifier (e.g. 'gpt2').
2406
+ pruned_artifact_path (str): Path to the .pt pruning artifact produced
2407
+ by discover_circuit — either a List[str] of node names (node-level)
2408
+ or a Dict with 'mlp'/'heads'/'_meta' keys (neuron-level).
2409
+ eval_params (Dict[str, Any]): Evaluation configuration. Recognised
2410
+ sub-key 'lm_eval' supports:
2411
+ enabled (bool): Skip lm-eval entirely if False. Default True.
2412
+ tasks (List[str]): lm-eval task names. Default: gsm8k, mmlu,
2413
+ truthfulqa, humaneval, hellaswag.
2414
+ fewshot (int): Number of few-shot examples. Default 0.
2415
+ limit (Optional[int]): Cap examples per task. Default None.
2416
+ max_gen_toks (int): Max generation tokens. Default 64.
2417
+ confirm_run_unsafe_code (bool): Required for code tasks. Default False.
2418
+ config_for_report (Dict[str, Any]): Original discovery config, currently
2419
+ used for logging context only.
2420
+ report_path (Optional[str]): If save_to_txt=True, write results to this
2421
+ path instead of an auto-generated timestamped file.
2422
+ precision (str): Torch dtype string for model loading
2423
+ ('bfloat16', 'float16', 'float32'). Defaults to 'bfloat16'.
2424
+ use_weight_based_pruning (bool): If True, prune by zeroing weights directly
2425
+ (more efficient, no hook overhead, does not require use_attn_result).
2426
+ If False, prune via forward hooks. Defaults to False.
2427
+ evaluation_mode (str): Which model variants to evaluate.
2428
+ 'both' runs original then pruned; 'original' skips pruned;
2429
+ 'pruned' skips original. Defaults to 'both'.
2430
+ save_to_txt (bool): If True, write results to a text file via
2431
+ _save_evaluation_results_to_txt. Defaults to False.
2432
+
2433
+ Returns:
2434
+ None: Results are printed to stdout and optionally written to a file.
2435
+
2436
+ Raises:
2437
+ ValueError: If evaluation_mode is not one of 'both', 'original', 'pruned'.
2438
+ TypeError: If the pruning artifact type is not a list or dict.
2439
+ FileNotFoundError: If pruned_artifact_path does not exist.
2440
+ """
2441
+ logger = get_logger("circuitkit.evaluation")
2442
+ progress = ProgressLogger(logger)
2443
+
2444
+ try:
2445
+ progress.start_operation("Circuit Evaluation", 3)
2446
+
2447
+ # Validate inputs
2448
+ validate_model_name(model_name)
2449
+ validate_file_exists(pruned_artifact_path, "pruned artifact")
2450
+
2451
+ # Validate evaluation_mode
2452
+ valid_modes = ["both", "original", "pruned"]
2453
+ if evaluation_mode not in valid_modes:
2454
+ raise ValueError(
2455
+ f"evaluation_mode must be one of {valid_modes}, got '{evaluation_mode}'"
2456
+ )
2457
+
2458
+ progress.step("Loading model and pruning artifacts", model=model_name)
2459
+ device = get_device()
2460
+ if not isinstance(getattr(t, precision, None), t.dtype):
2461
+ raise ValueError(
2462
+ f"Invalid precision '{precision}'. Pass 'precision' as a torch "
2463
+ f"dtype name such as 'float32', 'float16', or 'bfloat16'."
2464
+ )
2465
+ dtype = getattr(t, precision)
2466
+
2467
+ with log_execution_time("Model loading", logger):
2468
+ model = HookedTransformer.from_pretrained(model_name, device=device, dtype=dtype)
2469
+
2470
+ # Configure model for proper hook support (only needed for hook-based pruning)
2471
+ if not use_weight_based_pruning:
2472
+ model.cfg.use_attn_result = True
2473
+ model.cfg.use_hook_mlp_in = True
2474
+
2475
+ with log_execution_time("Artifact loading", logger):
2476
+ pruned_artifact = t.load(pruned_artifact_path, map_location="cpu", weights_only=True)
2477
+
2478
+ progress.step("Running evaluation")
2479
+
2480
+ if isinstance(pruned_artifact, list):
2481
+ logger.info(f"Detected node-level pruning artifact with {len(pruned_artifact)} nodes")
2482
+ elif isinstance(pruned_artifact, dict):
2483
+ mlp_count = sum(len(neurons) for neurons in pruned_artifact.get("mlp", {}).values())
2484
+ attn_count = sum(len(neurons) for neurons in pruned_artifact.get("heads", {}).values())
2485
+ logger.info(
2486
+ f"Detected neuron-level pruning artifact: {mlp_count} MLP neurons, {attn_count} attention neurons"
2487
+ )
2488
+ else:
2489
+ raise TypeError(f"Unknown artifact type for pruning: {type(pruned_artifact)}")
2490
+
2491
+ if report_path:
2492
+ logger.info(f"Report path specified: {report_path}")
2493
+ logger.warning("Report generation not yet implemented - results printed to console")
2494
+
2495
+ # Initialize results collection for txt file saving
2496
+ evaluation_results = []
2497
+
2498
+ # Run lm-evaluation-harness benchmarks
2499
+ lm_eval_cfg = eval_params.get("lm_eval", {}) if isinstance(eval_params, dict) else {}
2500
+ if lm_eval_cfg.get("enabled", True):
2501
+ tasks = lm_eval_cfg.get(
2502
+ "tasks",
2503
+ [
2504
+ "gsm8k",
2505
+ "mmlu",
2506
+ "truthfulqa",
2507
+ "humaneval",
2508
+ "hellaswag",
2509
+ ],
2510
+ )
2511
+ fewshot = int(lm_eval_cfg.get("fewshot", 0))
2512
+ limit = lm_eval_cfg.get("limit", None)
2513
+ int(lm_eval_cfg.get("max_gen_toks", 64))
2514
+ confirm_unsafe = bool(lm_eval_cfg.get("confirm_run_unsafe_code", False))
2515
+
2516
+ try:
2517
+ logger.info(f"Running lm-eval on tasks: {tasks}")
2518
+
2519
+ if use_weight_based_pruning:
2520
+ # Use weight-based pruning
2521
+ from .evaluation.weight_based_eval import (
2522
+ compare_original_vs_pruned_weight_based,
2523
+ evaluate_lm_eval_weight_based,
2524
+ )
2525
+
2526
+ if evaluation_mode == "both":
2527
+ results = compare_original_vs_pruned_weight_based(
2528
+ model,
2529
+ pruned_artifact,
2530
+ tasks=tasks,
2531
+ fewshot=fewshot,
2532
+ limit=limit,
2533
+ confirm_run_unsafe_code=confirm_unsafe,
2534
+ verbosity="WARNING",
2535
+ )
2536
+
2537
+ original_results = results["original"].get("results", results["original"])
2538
+ pruned_results = results["pruned"].get("results", results["pruned"])
2539
+
2540
+ logger.info("Original model results: %s", original_results)
2541
+ logger.info("Weight-pruned model results: %s", pruned_results)
2542
+
2543
+ # Collect results for txt file
2544
+ if save_to_txt:
2545
+ evaluation_results.append(
2546
+ {"model_type": "original", "results": original_results}
2547
+ )
2548
+ evaluation_results.append(
2549
+ {"model_type": "pruned", "results": pruned_results}
2550
+ )
2551
+ elif evaluation_mode == "original":
2552
+ # Only evaluate original model (empty artifact = no pruning)
2553
+ results = evaluate_lm_eval_weight_based(
2554
+ model,
2555
+ tasks=tasks,
2556
+ pruned_artifact=[],
2557
+ fewshot=fewshot,
2558
+ limit=limit,
2559
+ confirm_run_unsafe_code=confirm_unsafe,
2560
+ verbosity="WARNING",
2561
+ )
2562
+ original_results = results.get("results", results)
2563
+ logger.info("Original model results: %s", original_results)
2564
+
2565
+ # Collect results for txt file
2566
+ if save_to_txt:
2567
+ evaluation_results.append(
2568
+ {"model_type": "original", "results": original_results}
2569
+ )
2570
+ elif evaluation_mode == "pruned":
2571
+ # Only evaluate pruned model
2572
+ results = evaluate_lm_eval_weight_based(
2573
+ model,
2574
+ tasks=tasks,
2575
+ pruned_artifact=pruned_artifact,
2576
+ fewshot=fewshot,
2577
+ limit=limit,
2578
+ confirm_run_unsafe_code=confirm_unsafe,
2579
+ verbosity="WARNING",
2580
+ )
2581
+ pruned_results = results.get("results", results)
2582
+ logger.info("Weight-pruned model results: %s", pruned_results)
2583
+
2584
+ # Collect results for txt file
2585
+ if save_to_txt:
2586
+ evaluation_results.append(
2587
+ {"model_type": "pruned", "results": pruned_results}
2588
+ )
2589
+
2590
+ else:
2591
+ # Use hook-based pruning
2592
+ from .evaluation.lm_eval_simple import evaluate_lm_eval
2593
+
2594
+ if evaluation_mode in ["both", "original"]:
2595
+ # Original model
2596
+ orig = evaluate_lm_eval(
2597
+ model,
2598
+ tasks=tasks,
2599
+ fewshot=fewshot,
2600
+ limit=limit,
2601
+ confirm_run_unsafe_code=confirm_unsafe,
2602
+ verbosity="WARNING",
2603
+ )
2604
+ original_results = orig.get("results", orig)
2605
+ logger.info("Original model results: %s", original_results)
2606
+
2607
+ # Collect results for txt file
2608
+ if save_to_txt:
2609
+ evaluation_results.append(
2610
+ {"model_type": "original", "results": original_results}
2611
+ )
2612
+
2613
+ if evaluation_mode in ["both", "pruned"]:
2614
+ # Pruned model view via hooks
2615
+ pruned = evaluate_lm_eval(
2616
+ model,
2617
+ tasks=tasks,
2618
+ pruned_artifact=pruned_artifact,
2619
+ fewshot=fewshot,
2620
+ limit=limit,
2621
+ confirm_run_unsafe_code=confirm_unsafe,
2622
+ verbosity="WARNING",
2623
+ )
2624
+ pruned_results = pruned.get("results", pruned)
2625
+ logger.info("Pruned model results: %s", pruned_results)
2626
+
2627
+ # Collect results for txt file
2628
+ if save_to_txt:
2629
+ evaluation_results.append(
2630
+ {"model_type": "pruned", "results": pruned_results}
2631
+ )
2632
+
2633
+ except Exception as lm_err:
2634
+ logger.warning(f"lm-eval run skipped/failed: {lm_err}")
2635
+
2636
+ # Save results to txt file if requested
2637
+ if save_to_txt and evaluation_results:
2638
+ _save_evaluation_results_to_txt(
2639
+ evaluation_results,
2640
+ model_name,
2641
+ pruned_artifact_path,
2642
+ evaluation_mode,
2643
+ logger,
2644
+ report_path if report_path else None,
2645
+ )
2646
+
2647
+ progress.complete()
2648
+
2649
+ except Exception as e:
2650
+ progress.fail(str(e))
2651
+ raise
2652
+
2653
+
2654
+ def load_circuit(circuit_path: str) -> Union[List[str], Dict]:
2655
+ """
2656
+ Load a saved circuit pruning artifact from disk as its **raw** form.
2657
+
2658
+ This returns the low-level pruning artifact exactly as ``torch.save`` wrote
2659
+ it — a plain ``list[str]`` (node-level) or ``dict`` (neuron-level). It does
2660
+ **not** return a :class:`~circuitkit.Circuit` object and carries no scores
2661
+ or metadata. If you want a ready-to-use ``Circuit`` (with ``.scores``,
2662
+ ``.top_nodes()``, ``.task``, ...), use :func:`circuitkit.load_scores`
2663
+ instead — that is the loader most callers want. The two are not
2664
+ interchangeable: this one feeds ``prune``/``export``; ``load_scores`` feeds
2665
+ ``selective_finetune``/``Pipeline.from_scores``.
2666
+
2667
+ Args:
2668
+ circuit_path (str): Path to a .pt file produced by discover_circuit.
2669
+
2670
+ Returns:
2671
+ Union[List[str], Dict]: List of node name strings for node-level
2672
+ circuits, or a dict with keys 'mlp', 'heads', '_meta' for
2673
+ neuron-level circuits.
2674
+
2675
+ Raises:
2676
+ FileNotFoundError: If circuit_path does not exist.
2677
+
2678
+ See Also:
2679
+ circuitkit.load_scores: Load the same artifact as a rich ``Circuit``.
2680
+ """
2681
+ validate_file_exists(circuit_path, "circuit file")
2682
+ return t.load(circuit_path, map_location="cpu", weights_only=True)