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,962 @@
1
+ """
2
+ This is taken from mlab2 repo; arthur/induction branch
3
+
4
+ It is a very slightly edited version of https://github.com/redwoodresearch/Easy-Transformer/blob/main/easy_transformer/ioi_dataset.py
5
+ """
6
+
7
+ import copy
8
+ import random
9
+ import re
10
+ import warnings
11
+ from typing import List, Union
12
+
13
+ import numpy as np
14
+ import torch
15
+
16
+ from .....utils.logging import get_logger
17
+
18
+ logger = get_logger("data.task_data.ioi_dataset")
19
+
20
+ NAMES = [
21
+ "Michael",
22
+ "Christopher",
23
+ "Jessica",
24
+ "Matthew",
25
+ "Ashley",
26
+ "Jennifer",
27
+ "Joshua",
28
+ "Amanda",
29
+ "Daniel",
30
+ "David",
31
+ "James",
32
+ "Robert",
33
+ "John",
34
+ "Joseph",
35
+ "Andrew",
36
+ "Ryan",
37
+ "Brandon",
38
+ "Jason",
39
+ "Justin",
40
+ "Sarah",
41
+ "William",
42
+ "Jonathan",
43
+ "Stephanie",
44
+ "Brian",
45
+ "Nicole",
46
+ "Nicholas",
47
+ "Anthony",
48
+ "Heather",
49
+ "Eric",
50
+ "Elizabeth",
51
+ "Adam",
52
+ "Megan",
53
+ "Melissa",
54
+ "Kevin",
55
+ "Steven",
56
+ "Thomas",
57
+ "Timothy",
58
+ "Christina",
59
+ "Kyle",
60
+ "Rachel",
61
+ "Laura",
62
+ "Lauren",
63
+ "Amber",
64
+ "Brittany",
65
+ "Danielle",
66
+ "Richard",
67
+ "Kimberly",
68
+ "Jeffrey",
69
+ "Amy",
70
+ "Crystal",
71
+ "Michelle",
72
+ "Tiffany",
73
+ "Jeremy",
74
+ "Benjamin",
75
+ "Mark",
76
+ "Emily",
77
+ "Aaron",
78
+ "Charles",
79
+ "Rebecca",
80
+ "Jacob",
81
+ "Stephen",
82
+ "Patrick",
83
+ "Sean",
84
+ "Erin",
85
+ "Jamie",
86
+ "Kelly",
87
+ "Samantha",
88
+ "Nathan",
89
+ "Sara",
90
+ "Dustin",
91
+ "Paul",
92
+ "Angela",
93
+ "Tyler",
94
+ "Scott",
95
+ "Katherine",
96
+ "Andrea",
97
+ "Gregory",
98
+ "Erica",
99
+ "Mary",
100
+ "Travis",
101
+ "Lisa",
102
+ "Kenneth",
103
+ "Bryan",
104
+ "Lindsey",
105
+ "Kristen",
106
+ "Jose",
107
+ "Alexander",
108
+ "Jesse",
109
+ "Katie",
110
+ "Lindsay",
111
+ "Shannon",
112
+ "Vanessa",
113
+ "Courtney",
114
+ "Christine",
115
+ "Alicia",
116
+ "Cody",
117
+ "Allison",
118
+ "Bradley",
119
+ "Samuel",
120
+ ]
121
+
122
+ ABC_TEMPLATES = [
123
+ "Then, [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
124
+ "Afterwards [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
125
+ "When [A], [B] and [C] arrived at the [PLACE], [B] and [C] gave a [OBJECT] to [A]",
126
+ "Friends [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
127
+ ]
128
+
129
+ BAC_TEMPLATES = [
130
+ template.replace("[B]", "[A]", 1).replace("[A]", "[B]", 1) for template in ABC_TEMPLATES
131
+ ]
132
+
133
+ BABA_TEMPLATES = [
134
+ "Then, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
135
+ "Then, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
136
+ "Then, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
137
+ "Then, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
138
+ "Then, [B] and [A] had a long argument, and afterwards [B] said to [A]",
139
+ "After [B] and [A] went to the [PLACE], [B] gave a [OBJECT] to [A]",
140
+ "When [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give it to [A]",
141
+ "When [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give the [OBJECT] to [A]",
142
+ "While [B] and [A] were working at the [PLACE], [B] gave a [OBJECT] to [A]",
143
+ "While [B] and [A] were commuting to the [PLACE], [B] gave a [OBJECT] to [A]",
144
+ "After the lunch, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
145
+ "Afterwards, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
146
+ "Then, [B] and [A] had a long argument. Afterwards [B] said to [A]",
147
+ "The [PLACE] [B] and [A] went to had a [OBJECT]. [B] gave it to [A]",
148
+ "Friends [B] and [A] found a [OBJECT] at the [PLACE]. [B] gave it to [A]",
149
+ ]
150
+
151
+ BABA_LONG_TEMPLATES = [
152
+ "Then in the morning, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
153
+ "Then in the morning, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
154
+ "Then in the morning, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
155
+ "Then in the morning, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
156
+ "Then in the morning, [B] and [A] had a long argument, and afterwards [B] said to [A]",
157
+ "After taking a long break [B] and [A] went to the [PLACE], [B] gave a [OBJECT] to [A]",
158
+ "When soon afterwards [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give it to [A]",
159
+ "When soon afterwards [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give the [OBJECT] to [A]",
160
+ "While spending time together [B] and [A] were working at the [PLACE], [B] gave a [OBJECT] to [A]",
161
+ "While spending time together [B] and [A] were commuting to the [PLACE], [B] gave a [OBJECT] to [A]",
162
+ "After the lunch in the afternoon, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
163
+ "Afterwards, while spending time together [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
164
+ "Then in the morning afterwards, [B] and [A] had a long argument. Afterwards [B] said to [A]",
165
+ "The local big [PLACE] [B] and [A] went to had a [OBJECT]. [B] gave it to [A]",
166
+ "Friends separated at birth [B] and [A] found a [OBJECT] at the [PLACE]. [B] gave it to [A]",
167
+ ]
168
+
169
+ BABA_LATE_IOS = [
170
+ "Then, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
171
+ "Then, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
172
+ "Then, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
173
+ "Then, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
174
+ "Then, [B] and [A] had a long argument and after that [B] said to [A]",
175
+ "After the lunch, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
176
+ "Afterwards, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
177
+ "Then, [B] and [A] had a long argument. Afterwards [B] said to [A]",
178
+ ]
179
+
180
+ BABA_EARLY_IOS = [
181
+ "Then [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
182
+ "Then [B] and [A] had a lot of fun at the [PLACE], and [B] gave a [OBJECT] to [A]",
183
+ "Then [B] and [A] were working at the [PLACE], and [B] decided to give a [OBJECT] to [A]",
184
+ "Then [B] and [A] were thinking about going to the [PLACE], and [B] wanted to give a [OBJECT] to [A]",
185
+ "Then [B] and [A] had a long argument, and after that [B] said to [A]",
186
+ "After the lunch [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
187
+ "Afterwards [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
188
+ "Then [B] and [A] had a long argument, and afterwards [B] said to [A]",
189
+ ]
190
+
191
+ TEMPLATES_VARIED_MIDDLE = [
192
+ "",
193
+ ]
194
+
195
+ # no end of texts, GPT-2 small wasn't trained this way (ask Arthur)
196
+ # warnings.warn("Adding end of text prefixes!")
197
+ # for TEMPLATES in [BABA_TEMPLATES, BABA_EARLY_IOS, BABA_LATE_IOS]:
198
+ # for i in range(len(TEMPLATES)):
199
+ # TEMPLATES[i] = "<|endoftext|>" + TEMPLATES[i]
200
+
201
+ ABBA_TEMPLATES = BABA_TEMPLATES[:]
202
+ ABBA_LATE_IOS = BABA_LATE_IOS[:]
203
+ ABBA_EARLY_IOS = BABA_EARLY_IOS[:]
204
+
205
+ for TEMPLATES in [ABBA_TEMPLATES, ABBA_LATE_IOS, ABBA_EARLY_IOS]:
206
+ for i in range(len(TEMPLATES)):
207
+ first_clause = True
208
+ for j in range(1, len(TEMPLATES[i]) - 1):
209
+ if TEMPLATES[i][j - 1 : j + 2] == "[B]" and first_clause:
210
+ TEMPLATES[i] = TEMPLATES[i][:j] + "A" + TEMPLATES[i][j + 1 :]
211
+ elif TEMPLATES[i][j - 1 : j + 2] == "[A]" and first_clause:
212
+ first_clause = False
213
+ TEMPLATES[i] = TEMPLATES[i][:j] + "B" + TEMPLATES[i][j + 1 :]
214
+
215
+ VERBS = [" tried", " said", " decided", " wanted", " gave"]
216
+ PLACES = [
217
+ "store",
218
+ "garden",
219
+ "restaurant",
220
+ "school",
221
+ "hospital",
222
+ "office",
223
+ "house",
224
+ "station",
225
+ ]
226
+ OBJECTS = [
227
+ "ring",
228
+ "kiss",
229
+ "bone",
230
+ "basketball",
231
+ "computer",
232
+ "necklace",
233
+ "drink",
234
+ "snack",
235
+ ]
236
+
237
+ ANIMALS = [
238
+ "dog",
239
+ "cat",
240
+ "snake",
241
+ "elephant",
242
+ "beetle",
243
+ "hippo",
244
+ "giraffe",
245
+ "tiger",
246
+ "husky",
247
+ "lion",
248
+ "panther",
249
+ "whale",
250
+ "dolphin",
251
+ "beaver",
252
+ "rabbit",
253
+ "fox",
254
+ "lamb",
255
+ "ferret",
256
+ ]
257
+
258
+ def multiple_replace(dict, text):
259
+ # from: https://stackoverflow.com/questions/15175142/how-can-i-do-multiple-substitutions-using-regex
260
+ # Create a regular expression from the dictionary keys
261
+ regex = re.compile("(%s)" % "|".join(map(re.escape, dict.keys())))
262
+
263
+ # For each match, look-up corresponding value in dictionary
264
+ return regex.sub(lambda mo: dict[mo.string[mo.start() : mo.end()]], text)
265
+
266
+ def iter_sample_fast(iterable, samplesize, seed):
267
+ random.seed(seed)
268
+ results = []
269
+ # Fill in the first samplesize elements:
270
+ try:
271
+ for _ in range(samplesize):
272
+ results.append(next(iterable))
273
+ except StopIteration:
274
+ raise ValueError("Sample larger than population.")
275
+ random.shuffle(results) # Randomize their positions
276
+
277
+ return results
278
+
279
+ NOUNS_DICT = NOUNS_DICT = {"[PLACE]": PLACES, "[OBJECT]": OBJECTS}
280
+
281
+ def gen_prompt_uniform(
282
+ templates,
283
+ names,
284
+ nouns_dict,
285
+ N,
286
+ symmetric,
287
+ prefixes=None,
288
+ abc=False,
289
+ seed=None,
290
+ ):
291
+ assert seed is not None
292
+ random.seed(seed)
293
+
294
+ nb_gen = 0
295
+ ioi_prompts = []
296
+ while nb_gen < N:
297
+ temp = random.choice(templates)
298
+ temp_id = templates.index(temp)
299
+ name_1 = ""
300
+ name_2 = ""
301
+ name_3 = ""
302
+ while len(set([name_1, name_2, name_3])) < 3:
303
+ name_1 = random.choice(names)
304
+ name_2 = random.choice(names)
305
+ name_3 = random.choice(names)
306
+
307
+ nouns = {}
308
+ ioi_prompt = {}
309
+ for k in nouns_dict:
310
+ nouns[k] = random.choice(nouns_dict[k])
311
+ ioi_prompt[k] = nouns[k]
312
+ prompt = temp
313
+ for k in nouns_dict:
314
+ prompt = prompt.replace(k, nouns[k])
315
+
316
+ if prefixes is not None:
317
+ L = random.randint(30, 40)
318
+ pref = ".".join(random.choice(prefixes).split(".")[:L])
319
+ pref += "<|endoftext|>"
320
+ else:
321
+ pref = ""
322
+
323
+ prompt1 = prompt.replace("[A]", name_1)
324
+ prompt1 = prompt1.replace("[B]", name_2)
325
+ if abc:
326
+ prompt1 = prompt1.replace("[C]", name_3)
327
+ prompt1 = pref + prompt1
328
+ ioi_prompt["text"] = prompt1
329
+ ioi_prompt["IO"] = name_1
330
+ ioi_prompt["S"] = name_2
331
+ ioi_prompt["TEMPLATE_IDX"] = temp_id
332
+ ioi_prompts.append(ioi_prompt)
333
+ if abc:
334
+ ioi_prompts[-1]["C"] = name_3
335
+
336
+ nb_gen += 1
337
+
338
+ if symmetric and nb_gen < N:
339
+ prompt2 = prompt.replace("[A]", name_2)
340
+ prompt2 = prompt2.replace("[B]", name_1)
341
+ prompt2 = pref + prompt2
342
+ ioi_prompts.append(
343
+ {"text": prompt2, "IO": name_2, "S": name_1, "TEMPLATE_IDX": temp_id}
344
+ )
345
+ nb_gen += 1
346
+ return ioi_prompts
347
+
348
+ def gen_flipped_prompts( # noqa: C901 - complex function, refactor out of scope for lint pass
349
+ prompts, names, flip=("S2", "IO"), seed=None
350
+ ):
351
+ """_summary_
352
+
353
+ Args:
354
+ prompts (List[D]): _description_
355
+ flip (tuple, optional): First element is the string to be replaced, Second is what to replace with. Defaults to ("S2", "IO").
356
+
357
+ Returns:
358
+ _type_: _description_
359
+ """
360
+
361
+ assert seed is not None
362
+ np.random.seed(seed)
363
+
364
+ flipped_prompts = []
365
+
366
+ for prompt in prompts:
367
+ t = prompt["text"].split(" ")
368
+ prompt = prompt.copy()
369
+ if flip[0] == "S2":
370
+ if flip[1] == "IO":
371
+ t[len(t) - t[::-1].index(prompt["S"]) - 1] = prompt["IO"]
372
+ temp = prompt["IO"]
373
+ prompt["IO"] = prompt["S"]
374
+ prompt["S"] = temp
375
+ elif flip[1] == "RAND":
376
+ rand_name = names[np.random.randint(len(names))]
377
+ while rand_name == prompt["IO"] or rand_name == prompt["S"]:
378
+ rand_name = names[np.random.randint(len(names))]
379
+ t[len(t) - t[::-1].index(prompt["S"]) - 1] = rand_name
380
+ else:
381
+ raise ValueError("Invalid flip[1] value")
382
+
383
+ elif flip[0] == "IO":
384
+ if flip[1] == "RAND":
385
+ rand_name = names[np.random.randint(len(names))]
386
+ while rand_name == prompt["IO"] or rand_name == prompt["S"]:
387
+ rand_name = names[np.random.randint(len(names))]
388
+
389
+ t[t.index(prompt["IO"])] = rand_name
390
+ t[t.index(prompt["IO"])] = rand_name
391
+ prompt["IO"] = rand_name
392
+ elif flip[1] == "ANIMAL":
393
+ rand_animal = ANIMALS[np.random.randint(len(ANIMALS))]
394
+ t[t.index(prompt["IO"])] = rand_animal
395
+ prompt["IO"] = rand_animal
396
+ elif flip[1] == "S1":
397
+ io_index = t.index(prompt["IO"])
398
+ s1_index = t.index(prompt["S"])
399
+ io = t[io_index]
400
+ s1 = t[s1_index]
401
+ t[io_index] = s1
402
+ t[s1_index] = io
403
+ else:
404
+ raise ValueError("Invalid flip[1] value")
405
+
406
+ elif flip[0] in ["S", "S1"]:
407
+ if flip[1] == "ANIMAL":
408
+ new_s = ANIMALS[np.random.randint(len(ANIMALS))]
409
+ if flip[1] == "RAND":
410
+ new_s = names[np.random.randint(len(names))]
411
+ t[t.index(prompt["S"])] = new_s
412
+ if flip[0] == "S": # literally just change the first S if this is S1
413
+ t[len(t) - t[::-1].index(prompt["S"]) - 1] = new_s
414
+ prompt["S"] = new_s
415
+ elif flip[0] == "END":
416
+ if flip[1] == "S":
417
+ t[len(t) - t[::-1].index(prompt["IO"]) - 1] = prompt["S"]
418
+ elif flip[0] == "PUNC":
419
+ n = []
420
+
421
+ # separate the punctuation from the words
422
+ for i, word in enumerate(t):
423
+ if "." in word:
424
+ n.append(word[:-1])
425
+ n.append(".")
426
+ elif "," in word:
427
+ n.append(word[:-1])
428
+ n.append(",")
429
+ else:
430
+ n.append(word)
431
+
432
+ # remove punctuation, important that you check for period first
433
+ if flip[1] == "NONE":
434
+ if "." in n:
435
+ n[n.index(".")] = ""
436
+ elif "," in n:
437
+ n[len(n) - n[::-1].index(",") - 1] = ""
438
+
439
+ # remove empty strings
440
+ while "" in n:
441
+ n.remove("")
442
+
443
+ # add punctuation back to the word before it
444
+ while "," in n:
445
+ n[n.index(",") - 1] += ","
446
+ n.remove(",")
447
+
448
+ while "." in n:
449
+ n[n.index(".") - 1] += "."
450
+ n.remove(".")
451
+
452
+ t = n
453
+
454
+ elif flip[0] == "C2":
455
+ if flip[1] == "A":
456
+ t[len(t) - t[::-1].index(prompt["C"]) - 1] = prompt["A"]
457
+ elif flip[0] == "S+1":
458
+ if t[t.index(prompt["S"]) + 1] == "and":
459
+ t[t.index(prompt["S"]) + 1] = [
460
+ "with one friend named",
461
+ "accompanied by",
462
+ ][np.random.randint(2)]
463
+ else:
464
+ t[t.index(prompt["S"]) + 1] = (
465
+ t[t.index(prompt["S"])] + ", after a great day, " + t[t.index(prompt["S"]) + 1]
466
+ )
467
+ del t[t.index(prompt["S"])]
468
+ else:
469
+ raise ValueError(f"Invalid flipper {flip[0]}")
470
+
471
+ if "IO" in prompt:
472
+ prompt["text"] = " ".join(t)
473
+ flipped_prompts.append(prompt)
474
+ else:
475
+ flipped_prompts.append(
476
+ {
477
+ "A": prompt["A"],
478
+ "B": prompt["B"],
479
+ "C": prompt["C"],
480
+ "text": " ".join(t),
481
+ }
482
+ )
483
+
484
+ return flipped_prompts
485
+
486
+ # *Tok Idxs Methods
487
+
488
+ def get_name_idxs(prompts, tokenizer, idx_types=["IO", "S", "S2"], prepend_bos=False):
489
+ name_idx_dict = dict((idx_type, []) for idx_type in idx_types)
490
+ double_s2 = False
491
+ for prompt in prompts:
492
+ t = prompt["text"].split(" ")
493
+ toks = tokenizer.tokenize(" ".join(t[:-1]))
494
+ for idx_type in idx_types:
495
+ if "2" in idx_type:
496
+ idx = (
497
+ len(toks)
498
+ - toks[::-1].index(tokenizer.tokenize(" " + prompt[idx_type[:-1]])[0])
499
+ - 1
500
+ )
501
+ else:
502
+ idx = toks.index(tokenizer.tokenize(" " + prompt[idx_type])[0])
503
+ name_idx_dict[idx_type].append(idx)
504
+ if "S" in idx_types and "S2" in idx_types:
505
+ if name_idx_dict["S"][-1] == name_idx_dict["S2"][-1]:
506
+ double_s2 = True
507
+ if double_s2:
508
+ warnings.warn("S2 index has been computed as the same for S and S2")
509
+
510
+ return [int(prepend_bos) + torch.tensor(name_idx_dict[idx_type]) for idx_type in idx_types]
511
+
512
+ def get_word_idxs(prompts, word_list, tokenizer):
513
+ """Get the index of the words in word_list in the prompts. Exactly one of the word_list word has to be present in each prompt"""
514
+ idxs = []
515
+ tokenized_words = [tokenizer.decode(tokenizer(word)["input_ids"][0]) for word in word_list]
516
+ for pr_idx, prompt in enumerate(prompts):
517
+ toks = [
518
+ tokenizer.decode(t)
519
+ for t in tokenizer(prompt["text"], return_tensors="pt", padding=True)["input_ids"][0]
520
+ ]
521
+ idx = None
522
+ for i, w_tok in enumerate(tokenized_words):
523
+ if word_list[i] in prompt["text"]:
524
+ try:
525
+ idx = toks.index(w_tok)
526
+ if toks.count(w_tok) > 1:
527
+ idx = len(toks) - toks[::-1].index(w_tok) - 1
528
+ except Exception:
529
+ idx = toks.index(w_tok)
530
+ # raise ValueError(toks, w_tok, prompt["text"])
531
+ if idx is None:
532
+ raise ValueError(f"Word {word_list} and {i} not found {prompt}")
533
+ idxs.append(idx)
534
+ return torch.tensor(idxs)
535
+
536
+ def get_end_idxs(prompts, tokenizer, name_tok_len=1, prepend_bos=False, toks=None):
537
+
538
+ # toks = torch.Tensor(tokenizer([prompt["text"] for prompt in prompts], padding=True).input_ids).type(torch.int)
539
+ relevant_idx = int(prepend_bos)
540
+ # if the sentence begins with an end token
541
+ # AND the model pads at the end with the same end token,
542
+ # then we need make special arrangements
543
+
544
+ pad_token_id = tokenizer.pad_token_id
545
+
546
+ end_idxs_raw = []
547
+ for i in range(toks.shape[0]):
548
+ if pad_token_id not in toks[i][1:]:
549
+ end_idxs_raw.append(toks.shape[1])
550
+ continue
551
+ nonzers = (toks[i] == pad_token_id).nonzero()
552
+ try:
553
+ nonzers = nonzers[relevant_idx]
554
+ except Exception:
555
+ logger.error(toks[i])
556
+ logger.error(nonzers)
557
+ logger.error(relevant_idx)
558
+ logger.error(i)
559
+ raise ValueError("Something went wrong")
560
+ nonzers = nonzers[0]
561
+ nonzers = nonzers.item()
562
+ end_idxs_raw.append(nonzers)
563
+ end_idxs = torch.tensor(end_idxs_raw)
564
+ end_idxs = end_idxs - 1 - name_tok_len
565
+
566
+ for i in range(toks.shape[0]):
567
+ assert toks[i][end_idxs[i] + 1] != 0 and (
568
+ toks.shape[1] == end_idxs[i] + 2 or toks[i][end_idxs[i] + 2] == pad_token_id
569
+ ), (
570
+ toks[i],
571
+ end_idxs[i],
572
+ toks[i].shape,
573
+ "the END idxs aren't properly formatted",
574
+ )
575
+
576
+ return end_idxs
577
+
578
+ ALL_SEM = [
579
+ "S",
580
+ "IO",
581
+ "S2",
582
+ "end",
583
+ "S+1",
584
+ "and",
585
+ ] # , "verb", "starts", "S-1", "punct"] # Kevin's antic averages
586
+
587
+ def get_idx_dict(ioi_prompts, tokenizer, prepend_bos=False, toks=None):
588
+ (
589
+ IO_idxs,
590
+ S_idxs,
591
+ S2_idxs,
592
+ ) = get_name_idxs(
593
+ ioi_prompts,
594
+ tokenizer,
595
+ idx_types=["IO", "S", "S2"],
596
+ prepend_bos=prepend_bos,
597
+ )
598
+
599
+ end_idxs = get_end_idxs(
600
+ ioi_prompts,
601
+ tokenizer,
602
+ name_tok_len=1,
603
+ prepend_bos=prepend_bos,
604
+ toks=toks,
605
+ )
606
+
607
+ punct_idxs = get_word_idxs(ioi_prompts, [",", "."], tokenizer)
608
+
609
+ return {
610
+ "IO": IO_idxs,
611
+ "IO-1": IO_idxs - 1,
612
+ "IO+1": IO_idxs + 1,
613
+ "S": S_idxs,
614
+ "S-1": S_idxs - 1,
615
+ "S+1": S_idxs + 1,
616
+ "S2": S2_idxs,
617
+ "end": end_idxs,
618
+ "starts": torch.zeros_like(end_idxs),
619
+ "punct": punct_idxs,
620
+ }
621
+
622
+ # Some functions for experiments on Pointer Arithmetic
623
+
624
+ PREFIXES = [
625
+ " Afterwards,",
626
+ " Two friends met at a bar. Then,",
627
+ " After a long day,",
628
+ " After a long day,",
629
+ " Then,",
630
+ " Then,",
631
+ ]
632
+
633
+ def flip_prefixes(ioi_prompts):
634
+ ioi_prompts = copy.deepcopy(ioi_prompts)
635
+ for prompt in ioi_prompts:
636
+ if prompt["text"].startswith("The "):
637
+ prompt["text"] = "After the lunch, the" + prompt["text"][4:]
638
+ else:
639
+ io_idx = prompt["text"].index(prompt["IO"])
640
+ s_idx = prompt["text"].index(prompt["S"])
641
+ first_idx = min(io_idx, s_idx)
642
+ prompt["text"] = random.choice(PREFIXES) + " " + prompt["text"][first_idx:]
643
+
644
+ return ioi_prompts
645
+
646
+ def flip_names(ioi_prompts):
647
+ ioi_prompts = copy.deepcopy(ioi_prompts)
648
+ for prompt in ioi_prompts:
649
+ punct_idx = max(
650
+ [i for i, x in enumerate(list(prompt["text"])) if x in [",", "."]]
651
+ ) # only flip name in the first clause
652
+ io = prompt["IO"]
653
+ s = prompt["S"]
654
+ prompt["text"] = (
655
+ prompt["text"][:punct_idx]
656
+ .replace(io, "#")
657
+ .replace(s, "@")
658
+ .replace("#", s)
659
+ .replace("@", io)
660
+ ) + prompt["text"][punct_idx:]
661
+
662
+ return ioi_prompts
663
+
664
+ class IOIDataset:
665
+ def __init__(
666
+ self,
667
+ prompt_type: Union[str, List[str]], # if list, then it will be a list of templates
668
+ N=500,
669
+ model=None, # Required model parameter for TokenIDGenerator
670
+ prompts=None,
671
+ symmetric=False,
672
+ prefixes=None,
673
+ nb_templates=None,
674
+ ioi_prompts_for_word_idxs=None,
675
+ prepend_bos=False,
676
+ manual_word_idx=None,
677
+ seed=None,
678
+ ):
679
+ """
680
+ ioi_prompts_for_word_idxs:
681
+ if you want to use a different set of prompts to get the word indices, you can pass it here
682
+ (example use case: making a ABCA dataset)
683
+ """
684
+
685
+ assert seed is not None
686
+ random.seed(seed)
687
+
688
+ if not (
689
+ N == 1
690
+ or prepend_bos is False
691
+ or tokenizer.bos_token_id # noqa: F821 - pre-existing vendored bug
692
+ == tokenizer.eos_token_id # noqa: F821 - pre-existing vendored bug
693
+ ):
694
+ warnings.warn("Probably word_idx will be calculated incorrectly due to this formatting")
695
+ assert not (symmetric and prompt_type == "ABC")
696
+ assert (prompts is not None) or (not symmetric) or (N % 2 == 0), f"{symmetric} {N}"
697
+ assert nb_templates is None or (nb_templates % 2 == 0 or prompt_type != "mixed")
698
+ self.prompt_type = prompt_type
699
+
700
+ if nb_templates is None:
701
+ nb_templates = len(BABA_TEMPLATES)
702
+
703
+ if prompt_type == "ABBA":
704
+ self.templates = ABBA_TEMPLATES[:nb_templates].copy()
705
+ elif prompt_type == "BABA":
706
+ self.templates = BABA_TEMPLATES[:nb_templates].copy()
707
+ elif prompt_type == "mixed":
708
+ self.templates = (
709
+ BABA_TEMPLATES[: nb_templates // 2].copy()
710
+ + ABBA_TEMPLATES[: nb_templates // 2].copy()
711
+ )
712
+ random.shuffle(self.templates)
713
+ elif prompt_type == "ABC":
714
+ self.templates = ABC_TEMPLATES[:nb_templates].copy()
715
+ elif prompt_type == "BAC":
716
+ self.templates = BAC_TEMPLATES[:nb_templates].copy()
717
+ elif prompt_type == "ABC mixed":
718
+ self.templates = (
719
+ ABC_TEMPLATES[: nb_templates // 2].copy()
720
+ + BAC_TEMPLATES[: nb_templates // 2].copy()
721
+ )
722
+ random.shuffle(self.templates)
723
+ elif isinstance(prompt_type, list):
724
+ self.templates = prompt_type
725
+ else:
726
+ raise ValueError(prompt_type)
727
+
728
+ if model is None:
729
+ raise ValueError(
730
+ "Model is required for IOIDataset. "
731
+ "No default model to ensure model compatibility."
732
+ )
733
+ self.tokenizer = model.tokenizer
734
+ self.model = model
735
+
736
+ self.prefixes = prefixes
737
+ self.prompt_type = prompt_type
738
+ if prompts is None:
739
+ self.ioi_prompts = gen_prompt_uniform( # a list of dict of the form {"text": "Alice and Bob bla bla. Bob gave bla to Alice", "IO": "Alice", "S": "Bob"}
740
+ self.templates,
741
+ NAMES,
742
+ nouns_dict={"[PLACE]": PLACES, "[OBJECT]": OBJECTS},
743
+ N=N,
744
+ symmetric=symmetric,
745
+ prefixes=self.prefixes,
746
+ abc=(prompt_type in ["ABC", "ABC mixed", "BAC"]),
747
+ seed=(seed + 987654321) % 123456789,
748
+ )
749
+ else:
750
+ assert N == len(prompts), f"{N} and {len(prompts)}"
751
+ self.ioi_prompts = prompts
752
+
753
+ all_ids = [prompt["TEMPLATE_IDX"] for prompt in self.ioi_prompts]
754
+ all_ids_ar = np.array(all_ids)
755
+ self.groups = []
756
+ for id in list(set(all_ids)):
757
+ self.groups.append(np.where(all_ids_ar == id)[0])
758
+
759
+ small_groups = []
760
+ for group in self.groups:
761
+ if len(group) < 5:
762
+ small_groups.append(len(group))
763
+ if len(small_groups) > 0:
764
+ warnings.warn(f"Some groups have less than 5 prompts, they have lengths {small_groups}")
765
+
766
+ self.sentences = [
767
+ prompt["text"] for prompt in self.ioi_prompts
768
+ ] # a list of strings. Renamed as this should NOT be forward passed
769
+
770
+ self.templates_by_prompt = [] # for each prompt if it's ABBA or BABA
771
+ for i in range(N):
772
+ if self.sentences[i].index(self.ioi_prompts[i]["IO"]) < self.sentences[i].index(
773
+ self.ioi_prompts[i]["S"]
774
+ ):
775
+ self.templates_by_prompt.append("ABBA")
776
+ else:
777
+ self.templates_by_prompt.append("BABA")
778
+
779
+ texts = [
780
+ (self.tokenizer.bos_token if prepend_bos else "") + prompt["text"]
781
+ for prompt in self.ioi_prompts
782
+ ]
783
+ self.toks = torch.Tensor(self.tokenizer(texts, padding=True).input_ids).type(torch.int)
784
+
785
+ if ioi_prompts_for_word_idxs is None:
786
+ ioi_prompts_for_word_idxs = self.ioi_prompts
787
+ self.word_idx = get_idx_dict(
788
+ ioi_prompts_for_word_idxs,
789
+ self.tokenizer,
790
+ prepend_bos=prepend_bos,
791
+ toks=self.toks,
792
+ )
793
+ self.prepend_bos = prepend_bos
794
+ if manual_word_idx is not None:
795
+ self.word_idx = manual_word_idx
796
+
797
+ self.sem_tok_idx = {
798
+ k: v for k, v in self.word_idx.items() if k in ALL_SEM
799
+ } # the semantic indices that kevin uses
800
+ self.N = N
801
+ self.max_len = max(
802
+ [len(self.tokenizer(prompt["text"]).input_ids) for prompt in self.ioi_prompts]
803
+ )
804
+
805
+ # Use TokenIDGenerator for consistent token ID generation
806
+ from circuitkit.utils.token_utils import TokenIDGenerator
807
+
808
+ # Create a mock model object if only tokenizer is provided
809
+ if model is None:
810
+ # Create a minimal model-like object for TokenIDGenerator
811
+ class MockModel:
812
+ def __init__(self, tokenizer):
813
+ self.tokenizer = tokenizer
814
+ self.cfg = type("Config", (), {"model_name": "unknown"})()
815
+
816
+ mock_model = MockModel(self.tokenizer)
817
+ token_gen = TokenIDGenerator(mock_model)
818
+ else:
819
+ token_gen = TokenIDGenerator(model)
820
+
821
+ self.io_tokenIDs = token_gen.get_token_ids_batch([" " + p["IO"] for p in self.ioi_prompts])
822
+ self.s_tokenIDs = token_gen.get_token_ids_batch([" " + p["S"] for p in self.ioi_prompts])
823
+
824
+ self.tokenized_prompts = []
825
+
826
+ for i in range(self.N):
827
+ self.tokenized_prompts.append(
828
+ "|".join([self.tokenizer.decode(tok) for tok in self.toks[i]])
829
+ )
830
+
831
+ @classmethod
832
+ def construct_from_ioi_prompts_metadata(cls, templates, ioi_prompts_data, **kwargs):
833
+ """
834
+ Given a list of dictionaries (ioi_prompts_data)
835
+ {
836
+ "S": "Bob",
837
+ "IO": "Alice",
838
+ "TEMPLATE_IDX": 0
839
+ }
840
+
841
+ create and IOIDataset from these
842
+ """
843
+
844
+ prompts = []
845
+ for metadata in ioi_prompts_data:
846
+ cur_template = templates[metadata["TEMPLATE_IDX"]]
847
+ prompts.append(metadata)
848
+ prompts[-1]["text"] = (
849
+ cur_template.replace("[A]", metadata["IO"])
850
+ .replace("[B]", metadata["S"])
851
+ .replace("[PLACE]", metadata["[PLACE]"])
852
+ .replace("[OBJECT]", metadata["[OBJECT]"])
853
+ )
854
+ # prompts[-1]["[PLACE]"] = metadata["[PLACE]"]
855
+ # prompts[-1]["[OBJECT]"] = metadata["[OBJECT]"]
856
+ return IOIDataset(prompt_type=templates, prompts=prompts, **kwargs)
857
+
858
+ def gen_flipped_prompts(self, flip, seed=None):
859
+ """
860
+ Return a IOIDataset where the name to flip has been replaced by a random name.
861
+ """
862
+
863
+ assert seed is not None
864
+
865
+ assert isinstance(flip, tuple) or flip in [
866
+ "prefix",
867
+ ], f"{flip} is not a tuple. Probably change to ('IO', 'RAND') or equivalent?"
868
+
869
+ if flip == "prefix":
870
+ flipped_prompts = flip_prefixes(self.ioi_prompts)
871
+ else:
872
+ if flip in [("IO", "S1"), ("S", "IO")]:
873
+ flipped_prompts = gen_flipped_prompts(
874
+ self.ioi_prompts,
875
+ None,
876
+ flip,
877
+ seed=(seed + 12345) % 9876,
878
+ )
879
+ elif flip == ("S2", "IO"):
880
+ flipped_prompts = gen_flipped_prompts(
881
+ self.ioi_prompts,
882
+ None,
883
+ flip,
884
+ seed=(seed + 12345) % 6543,
885
+ )
886
+
887
+ else:
888
+ assert flip[1] == "RAND" and flip[0] in [
889
+ "S",
890
+ "RAND",
891
+ "S2",
892
+ "IO",
893
+ "S1",
894
+ "S+1",
895
+ ], flip
896
+ flipped_prompts = gen_flipped_prompts(
897
+ self.ioi_prompts, NAMES, flip, seed=(seed + 345467) % 5432
898
+ )
899
+
900
+ flipped_ioi_dataset = IOIDataset(
901
+ prompt_type=self.prompt_type,
902
+ N=self.N,
903
+ model=self.model,
904
+ prompts=flipped_prompts,
905
+ prefixes=self.prefixes,
906
+ ioi_prompts_for_word_idxs=flipped_prompts if flip[0] == "RAND" else None,
907
+ prepend_bos=self.prepend_bos,
908
+ manual_word_idx=self.word_idx,
909
+ seed=(seed + 23456) % 963,
910
+ )
911
+ return flipped_ioi_dataset
912
+
913
+ def copy(self):
914
+ copy_ioi_dataset = IOIDataset(
915
+ prompt_type=self.prompt_type,
916
+ N=self.N,
917
+ model=self.model,
918
+ prompts=self.ioi_prompts.copy(),
919
+ prefixes=self.prefixes.copy() if self.prefixes is not None else self.prefixes,
920
+ ioi_prompts_for_word_idxs=self.ioi_prompts.copy(),
921
+ )
922
+ return copy_ioi_dataset
923
+
924
+ def __getitem__(self, key):
925
+ sliced_prompts = self.ioi_prompts[key]
926
+ sliced_dataset = IOIDataset(
927
+ prompt_type=self.prompt_type,
928
+ N=len(sliced_prompts),
929
+ model=self.model,
930
+ prompts=sliced_prompts,
931
+ prefixes=self.prefixes,
932
+ prepend_bos=self.prepend_bos,
933
+ )
934
+ return sliced_dataset
935
+
936
+ def __setitem__(self, key, value):
937
+ raise TypeError(
938
+ "IOIDataset is immutable. To create a modified dataset, "
939
+ "construct a new IOIDataset() with the desired parameters."
940
+ )
941
+
942
+ def __delitem__(self, key):
943
+ raise TypeError(
944
+ "IOIDataset is immutable. To create a modified dataset, "
945
+ "construct a new IOIDataset() with the desired parameters."
946
+ )
947
+
948
+ def __len__(self):
949
+ return self.N
950
+
951
+ def tokenized_prompts(self):
952
+ return self.toks
953
+
954
+ # tests that the templates work as intended
955
+ # assert len(BABA_EARLY_IOS) == len(BABA_LATE_IOS), (len(BABA_EARLY_IOS), len(BABA_LATE_IOS))
956
+ # for i in range(len(BABA_EARLY_IOS)):
957
+ # d1 = IOIDataset(N=1, prompt_type=BABA_EARLY_IOS[i:i+1])
958
+ # d2 = IOIDataset(N=1, prompt_type=BABA_LATE_IOS[i:i+1])
959
+ # for tok in ["IO", "S"]: # occur one earlier and one later
960
+ # assert d1.word_idx[tok] + 1 == d2.word_idx[tok], (d1.word_idx[tok], d2.word_idx[tok])
961
+ # for tok in ["S2"]:
962
+ # assert d1.word_idx[tok] == d2.word_idx[tok], (d1.word_idx[tok], d2.word_idx[tok])