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
@@ -0,0 +1,382 @@
1
+ """
2
+ Comprehensive logging utilities for CircuitKit.
3
+ """
4
+
5
+ import json
6
+ import logging
7
+ import sys
8
+ import traceback
9
+ import warnings
10
+ from contextlib import contextmanager
11
+ from datetime import datetime
12
+ from functools import wraps
13
+ from pathlib import Path
14
+ from typing import Any, Dict, Optional
15
+
16
+
17
+ # Suppress common warnings that clutter output
18
+ def configure_warning_filters():
19
+ """Configure warning filters to reduce noise."""
20
+ # Suppress TransformerLens precision warnings
21
+ warnings.filterwarnings("ignore", message=".*reduced precision.*")
22
+ warnings.filterwarnings("ignore", message=".*from_pretrained_no_processing.*")
23
+
24
+ # Suppress lm-eval model warnings
25
+ warnings.filterwarnings("ignore", message=".*pretrained.*model kwarg is not of type.*")
26
+ warnings.filterwarnings("ignore", message=".*Passed an already-initialized model.*")
27
+ warnings.filterwarnings("ignore", message=".*Overwriting default num_fewshot.*")
28
+
29
+ # Suppress IOI dataset warnings (these are expected)
30
+ warnings.filterwarnings("ignore", message=".*S2 index has been computed.*")
31
+ warnings.filterwarnings("ignore", message=".*Some groups have less than 5 prompts.*")
32
+
33
+ # Suppress common torch warnings
34
+ warnings.filterwarnings("ignore", category=UserWarning, module="torch")
35
+ warnings.filterwarnings("ignore", category=UserWarning, module="circuitkit.data")
36
+
37
+ # Suppress via logging
38
+ logging.getLogger("transformers").setLevel(logging.ERROR)
39
+ logging.getLogger("lm_eval").setLevel(logging.ERROR)
40
+ logging.getLogger("accelerate").setLevel(logging.ERROR)
41
+
42
+ # Suppress root logger warnings from TransformerLens
43
+ logging.getLogger().setLevel(logging.ERROR)
44
+
45
+
46
+ class CircuitKitFormatter(logging.Formatter):
47
+ """Custom formatter for cleaner CircuitKit logs."""
48
+
49
+ # Color codes for terminal
50
+ COLORS = {
51
+ "DEBUG": "\033[36m", # Cyan
52
+ "INFO": "\033[32m", # Green
53
+ "WARNING": "\033[33m", # Yellow
54
+ "ERROR": "\033[31m", # Red
55
+ "CRITICAL": "\033[35m", # Magenta
56
+ "RESET": "\033[0m", # Reset
57
+ "BOLD": "\033[1m", # Bold
58
+ }
59
+
60
+ # Icons for different message types
61
+ ICONS = {
62
+ "step": "→",
63
+ "complete": "✓",
64
+ "error": "✗",
65
+ "performance": "⏱",
66
+ "model": "🔧",
67
+ "config": "⚙",
68
+ "start": "▶",
69
+ "info": "•",
70
+ }
71
+
72
+ def __init__(self, use_colors: bool = True):
73
+ super().__init__()
74
+ self.use_colors = use_colors and sys.stdout.isatty()
75
+
76
+ def format(self, record):
77
+ # Extract message
78
+ msg = record.getMessage()
79
+
80
+ # Determine icon and formatting
81
+ icon = self.ICONS["info"]
82
+ color = self.COLORS["INFO"] if self.use_colors else ""
83
+ reset = self.COLORS["RESET"] if self.use_colors else ""
84
+ self.COLORS["BOLD"] if self.use_colors else ""
85
+
86
+ if "Starting operation" in msg:
87
+ icon = self.ICONS["start"]
88
+ color = self.COLORS["INFO"] if self.use_colors else ""
89
+ elif "Step" in msg:
90
+ icon = self.ICONS["step"]
91
+ elif "Completed" in msg or "complete" in msg.lower():
92
+ icon = self.ICONS["complete"]
93
+ elif "Performance" in msg:
94
+ icon = self.ICONS["performance"]
95
+ elif "Model" in msg:
96
+ icon = self.ICONS["model"]
97
+ elif "Config" in msg:
98
+ icon = self.ICONS["config"]
99
+ elif record.levelno >= logging.ERROR:
100
+ icon = self.ICONS["error"]
101
+ color = self.COLORS["ERROR"] if self.use_colors else ""
102
+ elif record.levelno >= logging.WARNING:
103
+ color = self.COLORS["WARNING"] if self.use_colors else ""
104
+
105
+ # Format timestamp
106
+ timestamp = datetime.fromtimestamp(record.created).strftime("%H:%M:%S")
107
+
108
+ # Clean up the message - remove verbose JSON context for console
109
+ if "| Context:" in msg:
110
+ msg = msg.split("| Context:")[0].strip()
111
+
112
+ # Format final message
113
+ formatted = f"{color}{timestamp} {icon} {msg}{reset}"
114
+
115
+ return formatted
116
+
117
+
118
+ class CircuitKitLogger:
119
+ """Enhanced logger for CircuitKit with structured logging capabilities."""
120
+
121
+ def __init__(self, name: str = "circuitkit", level: int = logging.INFO):
122
+ self.logger = logging.getLogger(name)
123
+ self._context = {}
124
+
125
+ # Prevent duplicate handlers
126
+ if not self.logger.handlers:
127
+ self.logger.setLevel(level)
128
+ self._setup_handlers()
129
+
130
+ else:
131
+ # MODIFY: If handlers exist (e.g., from the global logger), ensure we update their levels
132
+ self.setLevel(level)
133
+
134
+ def _setup_handlers(self):
135
+ """Setup console and file handlers."""
136
+ # Avoid propagating to parent loggers to prevent duplicate logs
137
+ self.logger.propagate = False
138
+
139
+ # Console handler with custom formatter
140
+ console_handler = logging.StreamHandler(sys.stdout)
141
+ console_handler.setLevel(self.logger.level)
142
+ console_handler.setFormatter(CircuitKitFormatter(use_colors=True))
143
+ self.logger.addHandler(console_handler)
144
+
145
+ # File handler for detailed logs (with JSON context)
146
+ log_dir = Path("logs")
147
+ log_dir.mkdir(exist_ok=True)
148
+ file_handler = logging.FileHandler(
149
+ log_dir / f"circuitkit_{datetime.now().strftime('%Y%m%d')}.log"
150
+ )
151
+ file_handler.setLevel(logging.DEBUG)
152
+ file_formatter = logging.Formatter(
153
+ "%(asctime)s - %(name)s - %(levelname)s - %(funcName)s:%(lineno)d - %(message)s"
154
+ )
155
+ file_handler.setFormatter(file_formatter)
156
+ self.logger.addHandler(file_handler)
157
+
158
+ def setLevel(self, level: int):
159
+ """Update the level for the logger and its console handler."""
160
+ self.logger.setLevel(level)
161
+ for handler in self.logger.handlers:
162
+ # Only update the console stream handler, keep the file handler at DEBUG
163
+ if (
164
+ isinstance(handler, logging.StreamHandler)
165
+ and getattr(handler, "stream", None) == sys.stdout
166
+ ):
167
+ handler.setLevel(level)
168
+
169
+ def debug(self, message: str, **kwargs):
170
+ """Log debug message with optional context."""
171
+ self._log_with_context(logging.DEBUG, message, **kwargs)
172
+
173
+ def info(self, message: str, **kwargs):
174
+ """Log info message with optional context."""
175
+ self._log_with_context(logging.INFO, message, **kwargs)
176
+
177
+ def warning(self, message: str, **kwargs):
178
+ """Log warning message with optional context."""
179
+ self._log_with_context(logging.WARNING, message, **kwargs)
180
+
181
+ def error(self, message: str, **kwargs):
182
+ """Log error message with optional context."""
183
+ self._log_with_context(logging.ERROR, message, **kwargs)
184
+
185
+ def critical(self, message: str, **kwargs):
186
+ """Log critical message with optional context."""
187
+ self._log_with_context(logging.CRITICAL, message, **kwargs)
188
+
189
+ def _log_with_context(self, level: int, message: str, **kwargs):
190
+ """Log message with additional context."""
191
+ if kwargs:
192
+ context = json.dumps(kwargs, default=str)
193
+ message = f"{message} | Context: {context}"
194
+ self.logger.log(level, message)
195
+
196
+ def log_function_call(self, func_name: str, args: tuple, kwargs: dict, result: Any = None):
197
+ """Log function call details."""
198
+ self.debug(
199
+ f"Function call: {func_name}",
200
+ args=str(args)[:200],
201
+ kwargs=str(kwargs)[:200],
202
+ result_type=type(result).__name__ if result is not None else None,
203
+ )
204
+
205
+ def log_performance(self, operation: str, duration: float, **metrics):
206
+ """Log performance metrics."""
207
+ # Simple format for console
208
+ self.info(f"{operation}: {duration:.2f}s")
209
+
210
+ def log_model_info(self, model_name: str, **model_details):
211
+ """Log model information."""
212
+ params = model_details.get("parameters", 0)
213
+ if params > 1e9:
214
+ params_str = f"{params/1e9:.1f}B"
215
+ elif params > 1e6:
216
+ params_str = f"{params/1e6:.1f}M"
217
+ else:
218
+ params_str = f"{params:,}"
219
+ self.info(
220
+ f"Model: {model_name} ({params_str} params, {model_details.get('device', 'unknown')} device)"
221
+ )
222
+
223
+ def log_config(self, config: Dict[str, Any]):
224
+ """Log configuration details - simplified for console."""
225
+ algo = config.get("discovery", {}).get("algorithm", "unknown")
226
+ task = config.get("discovery", {}).get("task", "unknown")
227
+ level = config.get("discovery", {}).get("level", "node")
228
+ sparsity = config.get("pruning", {}).get("target_sparsity", 0)
229
+ self.info(f"Config: {algo.upper()} on {task} task, {level} level, {sparsity:.0%} sparsity")
230
+
231
+ def log_error_with_traceback(self, message: str, exception: Exception):
232
+ """Log error with full traceback."""
233
+ self.error(f"{message}: {str(exception)}")
234
+ self.debug("Full traceback:", traceback=traceback.format_exc())
235
+
236
+
237
+ # Global logger instance
238
+ logger = CircuitKitLogger()
239
+
240
+ # Configure warning filters on module import
241
+ configure_warning_filters()
242
+
243
+
244
+ def get_logger(name: Optional[str] = None) -> CircuitKitLogger:
245
+ """Get logger instance."""
246
+ if name:
247
+ return CircuitKitLogger(name)
248
+ return logger
249
+
250
+
251
+ def setup_logging(verbose: bool = False, log_file: Optional[str] = None):
252
+ """Setup logging configuration."""
253
+ level = logging.DEBUG if verbose else logging.INFO
254
+ logger = CircuitKitLogger(level=level)
255
+
256
+ # Reconfigure warning filters
257
+ configure_warning_filters()
258
+
259
+ if log_file:
260
+ # Add custom file handler
261
+ file_handler = logging.FileHandler(log_file)
262
+ file_handler.setLevel(logging.DEBUG)
263
+ formatter = logging.Formatter(
264
+ "%(asctime)s - %(name)s - %(levelname)s - %(funcName)s:%(lineno)d - %(message)s"
265
+ )
266
+ file_handler.setFormatter(formatter)
267
+ logger.logger.addHandler(file_handler)
268
+
269
+ return logger
270
+
271
+
272
+ @contextmanager
273
+ def log_execution_time(operation: str, logger: Optional[CircuitKitLogger] = None):
274
+ """Context manager to log execution time."""
275
+ if logger is None:
276
+ logger = get_logger()
277
+
278
+ start_time = datetime.now()
279
+
280
+ try:
281
+ yield
282
+ duration = (datetime.now() - start_time).total_seconds()
283
+ logger.log_performance(operation, duration)
284
+ except Exception as e:
285
+ duration = (datetime.now() - start_time).total_seconds()
286
+ logger.error(f"Failed: {operation} (took {duration:.3f}s)", error=str(e))
287
+ raise
288
+
289
+
290
+ def log_function_calls(logger: Optional[CircuitKitLogger] = None):
291
+ """Decorator to log function calls."""
292
+ if logger is None:
293
+ logger = get_logger()
294
+
295
+ def decorator(func):
296
+ @wraps(func)
297
+ def wrapper(*args, **kwargs):
298
+ logger.log_function_call(func.__name__, args, kwargs)
299
+ try:
300
+ result = func(*args, **kwargs)
301
+ logger.debug(f"Function {func.__name__} completed successfully")
302
+ return result
303
+ except Exception as e:
304
+ logger.log_error_with_traceback(f"Function {func.__name__} failed", e)
305
+ raise
306
+
307
+ return wrapper
308
+
309
+ return decorator
310
+
311
+
312
+ class ProgressLogger:
313
+ """Logger for progress tracking with structured output."""
314
+
315
+ def __init__(self, logger: Optional[CircuitKitLogger] = None):
316
+ self.logger = logger or get_logger()
317
+ self.steps = []
318
+ self.current_step = 0
319
+ self.start_time = None
320
+
321
+ def start_operation(self, operation: str, total_steps: int = 1):
322
+ """Start a new operation."""
323
+ self.operation = operation
324
+ self.total_steps = total_steps
325
+ self.current_step = 0
326
+ self.steps = []
327
+ self.start_time = datetime.now()
328
+ self.logger.info(f"{'='*50}")
329
+ self.logger.info(f"Starting: {operation}")
330
+ self.logger.info(f"{'='*50}")
331
+
332
+ def step(self, step_name: str, **context):
333
+ """Log a step in the operation."""
334
+ self.current_step += 1
335
+ self.steps.append(step_name)
336
+ # Format context nicely if present
337
+ if context:
338
+ ctx_str = ", ".join(f"{k}={v}" for k, v in context.items())
339
+ self.logger.info(f"[{self.current_step}/{self.total_steps}] {step_name} ({ctx_str})")
340
+ else:
341
+ self.logger.info(f"[{self.current_step}/{self.total_steps}] {step_name}")
342
+
343
+ def complete(self, **summary):
344
+ """Complete the operation."""
345
+ duration = (datetime.now() - self.start_time).total_seconds() if self.start_time else 0
346
+ summary_str = ", ".join(f"{k}={v}" for k, v in summary.items()) if summary else ""
347
+ self.logger.info(f"{'='*50}")
348
+ self.logger.info(f"Completed: {self.operation} in {duration:.1f}s")
349
+ if summary_str:
350
+ self.logger.info(f"Summary: {summary_str}")
351
+ self.logger.info(f"{'='*50}")
352
+
353
+ def fail(self, error: str, **context):
354
+ """Log operation failure."""
355
+ self.logger.error(f"Operation failed: {self.operation}")
356
+ self.logger.error(f"Error: {error}")
357
+
358
+
359
+ # Convenience functions
360
+ def debug(message: str, **kwargs):
361
+ """Log debug message."""
362
+ logger.debug(message, **kwargs)
363
+
364
+
365
+ def info(message: str, **kwargs):
366
+ """Log info message."""
367
+ logger.info(message, **kwargs)
368
+
369
+
370
+ def warning(message: str, **kwargs):
371
+ """Log warning message."""
372
+ logger.warning(message, **kwargs)
373
+
374
+
375
+ def error(message: str, **kwargs):
376
+ """Log error message."""
377
+ logger.error(message, **kwargs)
378
+
379
+
380
+ def critical(message: str, **kwargs):
381
+ """Log critical message."""
382
+ logger.critical(message, **kwargs)
@@ -0,0 +1,191 @@
1
+ """
2
+ Memory optimization utilities for CircuitKit.
3
+ Provides memory-efficient configurations and helpers.
4
+ """
5
+
6
+ import gc
7
+ from typing import Any, Dict
8
+
9
+ import torch
10
+
11
+ from circuitkit.utils.logging import get_logger
12
+
13
+ logger = get_logger(__name__)
14
+
15
+
16
+ def _estimate_model_params(model_name: str) -> int:
17
+ """Estimate parameter count using HuggingFace AutoConfig (no weight loading)."""
18
+ try:
19
+ from transformers import AutoConfig
20
+
21
+ cfg = AutoConfig.from_pretrained(model_name)
22
+ n_layers = getattr(cfg, "num_hidden_layers", 12)
23
+ d_model = getattr(cfg, "hidden_size", 768)
24
+ d_ffn = getattr(cfg, "intermediate_size", d_model * 4)
25
+ getattr(cfg, "num_attention_heads", 12)
26
+ vocab_size = getattr(cfg, "vocab_size", 50257)
27
+ # rough estimate: embed + n_layers*(attn + mlp) + lm_head
28
+ n_params = (
29
+ vocab_size * d_model # embedding
30
+ + n_layers * (4 * d_model * d_model + 2 * d_model * d_ffn) # attn + mlp
31
+ + vocab_size * d_model # lm_head
32
+ )
33
+ return n_params
34
+ except Exception:
35
+ return 0
36
+
37
+
38
+ def get_memory_efficient_config(model_name: str, algorithm: str = "eap-ig") -> Dict[str, Any]:
39
+ """Get memory-efficient configuration for any TL-supported model.
40
+
41
+ Uses HuggingFace AutoConfig to estimate model size without loading weights,
42
+ so this works for any architecture -- not just GPT-2 / Llama.
43
+
44
+ Args:
45
+ model_name: HuggingFace model ID or TransformerLens model name.
46
+ algorithm: Discovery algorithm.
47
+
48
+ Returns:
49
+ Memory-optimized configuration dictionary.
50
+ """
51
+ n_params = _estimate_model_params(model_name)
52
+ n_params_b = n_params / 1e9 # billions
53
+
54
+ # Scale settings by estimated size: <1B small, 1-10B medium, >10B large
55
+ if n_params_b >= 10:
56
+ precision = "bfloat16"
57
+ batch_size = 1
58
+ ig_steps = 1
59
+ sparsity = 0.05
60
+ mem_opt = {"gradient_checkpointing": True, "low_memory_mode": True, "max_memory_usage": 0.8}
61
+ elif n_params_b >= 1:
62
+ precision = "bfloat16"
63
+ batch_size = 1
64
+ ig_steps = 2
65
+ sparsity = 0.08
66
+ mem_opt = {
67
+ "gradient_checkpointing": False,
68
+ "low_memory_mode": False,
69
+ "max_memory_usage": 0.9,
70
+ }
71
+ else:
72
+ precision = "float32"
73
+ batch_size = 2
74
+ ig_steps = 3
75
+ sparsity = 0.1
76
+ mem_opt = {}
77
+
78
+ config: Dict[str, Any] = {
79
+ "model": {"name": model_name, "precision": precision},
80
+ "discovery": {
81
+ "algorithm": algorithm,
82
+ "level": "node",
83
+ "task": "ioi",
84
+ "batch_size": batch_size,
85
+ "ig_steps": ig_steps,
86
+ },
87
+ "pruning": {"target_sparsity": sparsity, "scope": "heads"},
88
+ "batch_size": batch_size,
89
+ }
90
+ if mem_opt:
91
+ config["memory_optimization"] = mem_opt
92
+
93
+ logger.info(
94
+ f"Generated memory-efficient config for {model_name} " f"(~{n_params_b:.1f}B params)",
95
+ context={"algorithm": algorithm, "batch_size": batch_size, "precision": precision},
96
+ )
97
+ return config
98
+
99
+
100
+ def optimize_memory_usage():
101
+ """Apply memory optimization settings."""
102
+ # Clear CUDA cache
103
+ if torch.cuda.is_available():
104
+ torch.cuda.empty_cache()
105
+ torch.cuda.synchronize()
106
+
107
+ # Force garbage collection
108
+ gc.collect()
109
+
110
+ # Set memory fraction if needed
111
+ if torch.cuda.is_available():
112
+ torch.cuda.set_per_process_memory_fraction(0.8)
113
+
114
+ logger.info("Applied memory optimizations")
115
+
116
+
117
+ def get_available_memory() -> Dict[str, float]:
118
+ """Get current memory usage information."""
119
+ memory_info = {}
120
+
121
+ if torch.cuda.is_available():
122
+ memory_info["cuda_allocated"] = torch.cuda.memory_allocated() / 1024**3 # GB
123
+ memory_info["cuda_reserved"] = torch.cuda.memory_reserved() / 1024**3 # GB
124
+ memory_info["cuda_max_allocated"] = torch.cuda.max_memory_allocated() / 1024**3 # GB
125
+ memory_info["cuda_total"] = torch.cuda.get_device_properties(0).total_memory / 1024**3 # GB
126
+ memory_info["cuda_free"] = memory_info["cuda_total"] - memory_info["cuda_allocated"]
127
+
128
+ return memory_info
129
+
130
+
131
+ def check_memory_requirements(model_name: str) -> bool:
132
+ """
133
+ Check if there's enough memory for the model.
134
+
135
+ Args:
136
+ model_name: Name of the model
137
+
138
+ Returns:
139
+ True if sufficient memory, False otherwise
140
+ """
141
+ memory_info = get_available_memory()
142
+
143
+ if not torch.cuda.is_available():
144
+ logger.warning("CUDA not available, cannot check memory requirements")
145
+ return False
146
+
147
+ # Estimate memory requirements from actual parameter count via HF config.
148
+ # Assume float32 (4 bytes/param) + 2x overhead for gradients/activations.
149
+ n_params = _estimate_model_params(model_name)
150
+ if n_params > 0:
151
+ required_memory = (n_params * 4 / 1024**3) * 2 # float32 * overhead
152
+ else:
153
+ required_memory = 8.0 # conservative default if config unavailable
154
+
155
+ available_memory = memory_info.get("cuda_free", 0)
156
+
157
+ logger.info(
158
+ "Memory check",
159
+ context={
160
+ "model": model_name,
161
+ "required": f"{required_memory}GB",
162
+ "available": f"{available_memory:.1f}GB",
163
+ "sufficient": available_memory >= required_memory,
164
+ },
165
+ )
166
+
167
+ return available_memory >= required_memory
168
+
169
+
170
+ def suggest_alternatives(model_name: str) -> list:
171
+ """
172
+ Suggest alternative models if the current one is too large.
173
+
174
+ Args:
175
+ model_name: Name of the model
176
+
177
+ Returns:
178
+ List of alternative model suggestions
179
+ """
180
+ model_name_lower = model_name.lower()
181
+
182
+ if "llama-3-8b" in model_name_lower or "llama-2-7b" in model_name_lower:
183
+ return ["gpt2", "gpt2-medium", "opt-125m", "opt-350m"]
184
+ elif "llama-3-70b" in model_name_lower or "llama-2-13b" in model_name_lower:
185
+ return ["gpt2-large", "gpt2-xl", "llama-2-7b", "llama-3-8b"]
186
+ elif "gpt2-xl" in model_name_lower:
187
+ return ["gpt2-large", "gpt2-medium", "gpt2"]
188
+ elif "gpt2-large" in model_name_lower:
189
+ return ["gpt2-medium", "gpt2", "opt-350m"]
190
+ else:
191
+ return ["gpt2", "gpt2-medium", "opt-125m"]