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/__init__.py ADDED
@@ -0,0 +1,128 @@
1
+ """
2
+ CircuitKit: A comprehensive toolkit for circuit discovery in transformer models.
3
+ """
4
+
5
+ # Version information
6
+ __version__ = "0.1.0"
7
+ __author__ = "Pratinav Seth, Hem Gosalia, Aditya Kasliwal, Vinay Kumar Sankarapu"
8
+ __description__ = "Unified Discover, Evaluate, Intervene toolkit for mechanistic interpretability"
9
+ _API_EXPORTS = {"discover_circuit", "evaluate_circuit", "load_circuit"}
10
+
11
+ # Flat front-door API (circuitkit.quick) — lazily imported so `import circuitkit`
12
+ # stays fast and torch-free until one of these is actually accessed.
13
+ _QUICK_EXPORTS = {
14
+ "load_model",
15
+ "discover",
16
+ "faithfulness",
17
+ "prune",
18
+ "quantize",
19
+ "export_checkpoint",
20
+ "benchmark",
21
+ "load_scores",
22
+ "selective_finetune",
23
+ "visualize_circuit",
24
+ }
25
+ _CIRCUIT_EXPORTS = {"Circuit"}
26
+ _PIPELINE_EXPORTS = {"Pipeline"}
27
+
28
+ # Main exports
29
+ __all__ = [
30
+ # Core dict-config API
31
+ "discover_circuit",
32
+ "evaluate_circuit",
33
+ "load_circuit",
34
+ # Flat front-door API
35
+ "load_model",
36
+ "discover",
37
+ "faithfulness",
38
+ "prune",
39
+ "quantize",
40
+ "export_checkpoint",
41
+ "benchmark",
42
+ "Circuit",
43
+ "load_scores",
44
+ "selective_finetune",
45
+ "visualize_circuit",
46
+ "Pipeline",
47
+ # Task management
48
+ "get_task",
49
+ "list_tasks",
50
+ "register_task",
51
+ # Version info
52
+ "__version__",
53
+ "__author__",
54
+ "__description__",
55
+ ]
56
+
57
+
58
+ def __getattr__(name):
59
+ """Lazily import heavy API helpers when accessed from the package root."""
60
+ if name in _API_EXPORTS:
61
+ from . import api
62
+
63
+ value = getattr(api, name)
64
+ globals()[name] = value
65
+ return value
66
+ if name in _QUICK_EXPORTS:
67
+ from . import quick
68
+
69
+ value = getattr(quick, name)
70
+ globals()[name] = value
71
+ return value
72
+ if name in _CIRCUIT_EXPORTS:
73
+ from .circuit import Circuit
74
+
75
+ globals()["Circuit"] = Circuit
76
+ return Circuit
77
+ if name in _PIPELINE_EXPORTS:
78
+ from .pipeline import Pipeline
79
+
80
+ globals()["Pipeline"] = Pipeline
81
+ return Pipeline
82
+ if name == "visualize":
83
+ # Deprecated pre-1.0 alias (H1 in the 1.0.0 audit): ``visualize`` was
84
+ # renamed to ``visualize_circuit`` in 1.0.0 and is now a subpackage.
85
+ # Return the (callable) subpackage so ``ck.visualize(...)`` keeps working
86
+ # and warns on use — see circuitkit/visualize/__init__.py. The warning
87
+ # fires on the call, not on attribute access, so it survives the
88
+ # submodule being imported (which shadows any parent __getattr__ shim).
89
+ from . import visualize as _visualize
90
+
91
+ return _visualize
92
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
93
+
94
+
95
+ def __dir__():
96
+ """Expose lazily-imported names for autocomplete / dir()."""
97
+ return sorted(set(globals()) | set(__all__))
98
+
99
+
100
+ def _ensure_builtin_tasks():
101
+ """Register built-in tasks only when task helpers are used."""
102
+ from .tasks.bootstrap import _bootstrap_builtin_tasks
103
+
104
+ return _bootstrap_builtin_tasks()
105
+
106
+
107
+ def get_task(name):
108
+ """Get a built-in or registered task by name."""
109
+ _ensure_builtin_tasks()
110
+ from .tasks.registry import get_task as _get_task
111
+
112
+ return _get_task(name)
113
+
114
+
115
+ def list_tasks():
116
+ """List registered task names."""
117
+ _ensure_builtin_tasks()
118
+ from .tasks.registry import list_tasks as _list_tasks
119
+
120
+ return _list_tasks()
121
+
122
+
123
+ def register_task(spec):
124
+ """Register a custom task specification."""
125
+ _ensure_builtin_tasks()
126
+ from .tasks.registry import register_task as _register_task
127
+
128
+ return _register_task(spec)
circuitkit/__main__.py ADDED
@@ -0,0 +1,9 @@
1
+ """
2
+ CircuitKit module execution entry point.
3
+ Allows running: python -m circuitkit
4
+ """
5
+
6
+ from .cli.main import main
7
+
8
+ if __name__ == "__main__":
9
+ main()
@@ -0,0 +1,19 @@
1
+ """
2
+ CircuitKit Analysis Module
3
+
4
+ Provides analysis tools for circuits including metrics, scoring, and statistical analysis.
5
+ """
6
+
7
+ from .cross_method_jaccard import CrossMethodJaccardResult, cross_method_jaccard
8
+ from .metrics import * # noqa: F401,F403 - intentional API re-export
9
+ from .scores import * # noqa: F401,F403 - intentional API re-export
10
+
11
+ __all__ = [ # noqa: F405 - names provided via star imports above
12
+ # Metrics
13
+ "compute_metrics",
14
+ # Scores
15
+ "compute_scores",
16
+ # Cross-method comparison (EMNLP 2026, Section 5)
17
+ "cross_method_jaccard",
18
+ "CrossMethodJaccardResult",
19
+ ]
@@ -0,0 +1,116 @@
1
+ """Cross-method circuit comparison via top-K Jaccard.
2
+
3
+ Used by the EMNLP 2026 submission's Finding B (Section 5) to
4
+ quantify component-level disagreement across discovery methods
5
+ discovering circuits for the same task on the same model.
6
+
7
+ Example
8
+ -------
9
+ >>> from circuitkit.analysis.cross_method_jaccard import (
10
+ ... cross_method_jaccard,
11
+ ... )
12
+ >>> # circuits: dict mapping method name to a CircuitScores artifact
13
+ >>> # (or a list of node names, or a dict node_name -> score)
14
+ >>> result = cross_method_jaccard(circuits, top_k=10)
15
+ >>> # result.matrix: numpy-style nested list of pairwise Jaccards
16
+ >>> # result.top_nodes: dict method -> top_k node names
17
+ >>> # result.range: (min, max) pairwise Jaccard over off-diagonal cells
18
+
19
+ The Jaccard at top-K is the standard descriptive statistic for the
20
+ "do methods agree on the circuit?" question reported in the EMNLP
21
+ paper, derived from the workshop result on GPT-2 IOI where the same
22
+ analysis showed Jaccard 0.11 to 1.00 across 8 EAP-family variants.
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ from dataclasses import dataclass
28
+ from typing import Dict, List, Tuple, Union
29
+
30
+ CircuitLike = Union[Dict[str, float], List[str], object]
31
+
32
+
33
+ def _extract_top_k_nodes(circuit: CircuitLike, top_k: int) -> List[str]:
34
+ """Coerce a circuit-like input into a top-K node-name list.
35
+
36
+ Accepts:
37
+ - dict[str, float]: node_name -> score; top-K by absolute score
38
+ - list[str]: assumed to be a pre-ranked list; truncated to top-K
39
+ - object with `.node_scores` attribute (e.g. CircuitScores): dict-like
40
+ """
41
+ if hasattr(circuit, "node_scores"):
42
+ node_scores = circuit.node_scores
43
+ else:
44
+ node_scores = circuit
45
+ if isinstance(node_scores, dict):
46
+ ranked = sorted(node_scores.items(), key=lambda kv: abs(float(kv[1])), reverse=True)
47
+ return [n for n, _ in ranked[:top_k]]
48
+ if isinstance(node_scores, list):
49
+ return list(node_scores[:top_k])
50
+ raise TypeError(f"Unsupported circuit type for jaccard: {type(circuit)!r}")
51
+
52
+
53
+ def _jaccard(a: List[str], b: List[str]) -> float:
54
+ sa, sb = set(a), set(b)
55
+ if not sa and not sb:
56
+ return 1.0
57
+ return len(sa & sb) / len(sa | sb)
58
+
59
+
60
+ @dataclass
61
+ class CrossMethodJaccardResult:
62
+ methods: List[str]
63
+ top_k: int
64
+ matrix: List[List[float]]
65
+ top_nodes: Dict[str, List[str]]
66
+ range: Tuple[float, float]
67
+ n_pairs: int
68
+
69
+
70
+ def cross_method_jaccard(
71
+ circuits: Dict[str, CircuitLike],
72
+ top_k: int = 10,
73
+ ) -> CrossMethodJaccardResult:
74
+ """Compute the symmetric pairwise Jaccard matrix between method
75
+ circuits at top-K nodes.
76
+
77
+ Parameters
78
+ ----------
79
+ circuits : dict
80
+ Maps method name to either a CircuitScores artifact, a dict of
81
+ node_name to score, or a pre-ranked list of node names.
82
+ top_k : int
83
+ Truncation depth for each method's circuit. Default 10 to match
84
+ the EMNLP paper's protocol.
85
+
86
+ Returns
87
+ -------
88
+ CrossMethodJaccardResult
89
+ Includes the pairwise Jaccard matrix, the top-K node sets per
90
+ method, and the off-diagonal range (min, max).
91
+ """
92
+ methods = sorted(circuits.keys())
93
+ n = len(methods)
94
+ top_nodes = {m: _extract_top_k_nodes(circuits[m], top_k) for m in methods}
95
+ matrix = [[1.0] * n for _ in range(n)]
96
+ off_diag: List[float] = []
97
+ for i, mi in enumerate(methods):
98
+ for j, mj in enumerate(methods):
99
+ if i <= j:
100
+ continue
101
+ j_val = _jaccard(top_nodes[mi], top_nodes[mj])
102
+ matrix[i][j] = j_val
103
+ matrix[j][i] = j_val
104
+ off_diag.append(j_val)
105
+ rng = (min(off_diag), max(off_diag)) if off_diag else (0.0, 0.0)
106
+ return CrossMethodJaccardResult(
107
+ methods=methods,
108
+ top_k=top_k,
109
+ matrix=matrix,
110
+ top_nodes=top_nodes,
111
+ range=rng,
112
+ n_pairs=len(off_diag),
113
+ )
114
+
115
+
116
+ __all__ = ["cross_method_jaccard", "CrossMethodJaccardResult"]
@@ -0,0 +1,54 @@
1
+ # FILE: circuitkit/analysis/metrics.py
2
+
3
+
4
+ def calculate_faithfulness(original_output, pruned_output):
5
+ """
6
+ Calculate faithfulness metric between original and pruned outputs.
7
+ Uses multiple metrics for comprehensive evaluation.
8
+ """
9
+ import torch
10
+ import torch.nn.functional as F
11
+
12
+ # Ensure tensors are the same shape
13
+ if original_output.shape != pruned_output.shape:
14
+ min_size = min(original_output.numel(), pruned_output.numel())
15
+ original_output = original_output.flatten()[:min_size]
16
+ pruned_output = pruned_output.flatten()[:min_size]
17
+
18
+ # L2 norm difference
19
+ l2_diff = (original_output - pruned_output).pow(2).sum().item()
20
+
21
+ # KL divergence
22
+ try:
23
+ kl_div = F.kl_div(
24
+ F.log_softmax(pruned_output, dim=-1),
25
+ F.softmax(original_output, dim=-1),
26
+ reduction="sum",
27
+ ).item()
28
+ except (RuntimeError, ValueError):
29
+ kl_div = float("inf")
30
+
31
+ # Cosine similarity
32
+ cos_sim = F.cosine_similarity(
33
+ original_output.flatten().unsqueeze(0), pruned_output.flatten().unsqueeze(0)
34
+ ).item()
35
+
36
+ # Relative difference
37
+ rel_diff = torch.abs(original_output - pruned_output).sum().item() / (
38
+ torch.abs(original_output).sum().item() + 1e-8
39
+ )
40
+
41
+ return {
42
+ "l2_difference": l2_diff,
43
+ "kl_divergence": kl_div,
44
+ "cosine_similarity": cos_sim,
45
+ "relative_difference": rel_diff,
46
+ "faithfulness_score": 1.0 - min(rel_diff, 1.0), # Higher is better
47
+ }
48
+
49
+
50
+ def calculate_complexity(graph):
51
+ """
52
+ Calculates the complexity of a circuit, e.g., by node or edge count.
53
+ """
54
+ return {"node_count": graph.number_of_nodes(), "edge_count": graph.number_of_edges()}
@@ -0,0 +1,44 @@
1
+ # FILE: circuitkit/analysis/scores.py
2
+
3
+ from collections import defaultdict
4
+
5
+ from ..backends.acdc.types import PruneScores
6
+ from ..backends.acdc.utils.patchable_model import PatchableModel
7
+
8
+
9
+ def calculate_node_scores_from_edges(
10
+ p_model: PatchableModel, edge_prune_scores: PruneScores
11
+ ) -> dict[str, float]:
12
+ """
13
+ Calculates an importance score for each source node based on its outgoing edges.
14
+
15
+ The score for a node is the average of the absolute scores of all its
16
+ outgoing edges. Nodes with no outgoing edges will not be included.
17
+
18
+ Args:
19
+ p_model: The patchable model, used to access the graph of nodes and edges.
20
+ edge_prune_scores: A dictionary mapping destination modules to tensors of
21
+ edge importance scores.
22
+
23
+ Returns:
24
+ A dictionary mapping source node names to their calculated importance score.
25
+ """
26
+ # Use absolute values of scores as importance can be positive or negative in EAP
27
+ abs_edge_scores = {mod: scores.abs() for mod, scores in edge_prune_scores.items()}
28
+
29
+ # Group edge scores by their source node
30
+ outgoing_scores_by_node = defaultdict(list)
31
+ for edge in p_model.edges:
32
+ # edge.prune_score looks up the score for this specific edge
33
+ score = edge.prune_score(abs_edge_scores).item()
34
+ outgoing_scores_by_node[edge.src.name].append(score)
35
+
36
+ # Calculate the average score for each node. Use Python's built-in sum/len
37
+ # rather than np.mean() so values stay plain Python floats — numpy scalars
38
+ # (numpy.float64) are rejected by torch.load with weights_only=True.
39
+ node_scores = {}
40
+ for node_name, scores in outgoing_scores_by_node.items():
41
+ if scores:
42
+ node_scores[node_name] = sum(scores) / len(scores)
43
+
44
+ return node_scores