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.
- circuitkit/__init__.py +128 -0
- circuitkit/__main__.py +9 -0
- circuitkit/analysis/__init__.py +19 -0
- circuitkit/analysis/cross_method_jaccard.py +116 -0
- circuitkit/analysis/metrics.py +54 -0
- circuitkit/analysis/scores.py +44 -0
- circuitkit/api.py +2682 -0
- circuitkit/applications/__init__.py +70 -0
- circuitkit/applications/arch_registry.py +315 -0
- circuitkit/applications/arch_utils.py +302 -0
- circuitkit/applications/common_utils/__init__.py +15 -0
- circuitkit/applications/common_utils/_covariance.py +223 -0
- circuitkit/applications/common_utils/_device.py +33 -0
- circuitkit/applications/common_utils/_metrics.py +429 -0
- circuitkit/applications/common_utils/_tokenization.py +510 -0
- circuitkit/applications/common_utils/benchmark_analysis.py +400 -0
- circuitkit/applications/common_utils/cure_clue.py +338 -0
- circuitkit/applications/common_utils/hallucination_detection.py +497 -0
- circuitkit/applications/common_utils/linear_probe.py +294 -0
- circuitkit/applications/editing/__init__.py +30 -0
- circuitkit/applications/editing/cake.py +237 -0
- circuitkit/applications/editing/circuit_guided_editing.py +487 -0
- circuitkit/applications/editing/fine_tune_editing.py +327 -0
- circuitkit/applications/editing/knowledge_editing.py +598 -0
- circuitkit/applications/editing/knowledge_editing_enhanced.py +863 -0
- circuitkit/applications/editing/mcircke.py +263 -0
- circuitkit/applications/editing/memit_wrapper.py +947 -0
- circuitkit/applications/editing/rome_wrapper.py +770 -0
- circuitkit/applications/finetuning/__init__.py +17 -0
- circuitkit/applications/finetuning/benchmark_peft.py +451 -0
- circuitkit/applications/finetuning/circuit_tuning.py +377 -0
- circuitkit/applications/finetuning/healing_metrics.py +304 -0
- circuitkit/applications/finetuning/peft_methods.py +563 -0
- circuitkit/applications/finetuning/soft_healing.py +717 -0
- circuitkit/applications/pruning/__init__.py +13 -0
- circuitkit/applications/pruning/eval_utils.py +347 -0
- circuitkit/applications/pruning/examples/__init__.py +0 -0
- circuitkit/applications/pruning/examples/prune.py +787 -0
- circuitkit/applications/pruning/examples/prune_llama.py +661 -0
- circuitkit/applications/pruning/examples/prune_qwen.py +591 -0
- circuitkit/applications/pruning/finetune_utils.py +477 -0
- circuitkit/applications/pruning/importance.py +97 -0
- circuitkit/applications/pruning/neuron_pruner.py +36 -0
- circuitkit/applications/pruning/node_pruner.py +186 -0
- circuitkit/applications/pruning/pruner.py +529 -0
- circuitkit/applications/pruning/score_extractor.py +541 -0
- circuitkit/applications/pruning/selectors/__init__.py +0 -0
- circuitkit/applications/pruning/selectors/multi_granular_selector.py +119 -0
- circuitkit/applications/pruning/selectors/taylor_selector.py +112 -0
- circuitkit/applications/pruning/weight_pruner.py +652 -0
- circuitkit/applications/quantization/__init__.py +19 -0
- circuitkit/applications/quantization/examples/__init__.py +0 -0
- circuitkit/applications/quantization/examples/quantize_llama.py +745 -0
- circuitkit/applications/quantization/examples/quantize_qwen.py +726 -0
- circuitkit/applications/quantization/llmcompressor_quantize.py +388 -0
- circuitkit/applications/quantization/quant_utils.py +825 -0
- circuitkit/applications/quantization/score_extractor.py +465 -0
- circuitkit/applications/quantization/selectors/__init__.py +0 -0
- circuitkit/applications/quantization/selectors/awq_selector.py +109 -0
- circuitkit/applications/quantization/selectors/tacq_selector.py +145 -0
- circuitkit/applications/selective_finetuning/__init__.py +0 -0
- circuitkit/applications/selective_finetuning/examples/__init__.py +0 -0
- circuitkit/applications/selective_finetuning/examples/finetune_llama.py +668 -0
- circuitkit/applications/selective_finetuning/examples/finetune_qwen.py +654 -0
- circuitkit/applications/selective_finetuning/finetune_utils.py +643 -0
- circuitkit/applications/selective_finetuning/score_loader.py +585 -0
- circuitkit/applications/selective_finetuning/selector.py +616 -0
- circuitkit/applications/steering/__init__.py +31 -0
- circuitkit/applications/steering/steering.py +791 -0
- circuitkit/applications/steering/steering_enhanced.py +556 -0
- circuitkit/applications/steering/weight_steering.py +407 -0
- circuitkit/artifacts/__init__.py +25 -0
- circuitkit/artifacts/circuit_artifact.py +559 -0
- circuitkit/artifacts/converters.py +405 -0
- circuitkit/artifacts/scores.py +195 -0
- circuitkit/backends/__init__.py +96 -0
- circuitkit/backends/acdc/__init__.py +0 -0
- circuitkit/backends/acdc/artifact_export.py +140 -0
- circuitkit/backends/acdc/data.py +183 -0
- circuitkit/backends/acdc/model_utils/__init__.py +0 -0
- circuitkit/backends/acdc/model_utils/micro_model_utils.py +143 -0
- circuitkit/backends/acdc/model_utils/transformer_lens_utils.py +232 -0
- circuitkit/backends/acdc/prune.py +123 -0
- circuitkit/backends/acdc/prune_algos/ACDC.py +136 -0
- circuitkit/backends/acdc/prune_algos/__init__.py +0 -0
- circuitkit/backends/acdc/prune_algos/mask_gradient.py +129 -0
- circuitkit/backends/acdc/prune_algos/prune_algos.py +33 -0
- circuitkit/backends/acdc/tasks/__init__.py +3 -0
- circuitkit/backends/acdc/tasks/docstring_prompts.py +837 -0
- circuitkit/backends/acdc/tasks/docstring_utils.py +88 -0
- circuitkit/backends/acdc/tasks/induction_utils.py +122 -0
- circuitkit/backends/acdc/tasks/ioi_dataset.py +156 -0
- circuitkit/backends/acdc/tasks/ioi_utils.py +96 -0
- circuitkit/backends/acdc/types.py +248 -0
- circuitkit/backends/acdc/utils/__init__.py +0 -0
- circuitkit/backends/acdc/utils/ablation_activations.py +160 -0
- circuitkit/backends/acdc/utils/custom_tqdm.py +14 -0
- circuitkit/backends/acdc/utils/graph_utils.py +497 -0
- circuitkit/backends/acdc/utils/misc.py +68 -0
- circuitkit/backends/acdc/utils/patch_wrapper.py +106 -0
- circuitkit/backends/acdc/utils/patchable_model.py +143 -0
- circuitkit/backends/acdc/utils/task_utils.py +31 -0
- circuitkit/backends/acdc/utils/tensor_ops.py +118 -0
- circuitkit/backends/acdc/visualize.py +253 -0
- circuitkit/backends/cdt/__init__.py +30 -0
- circuitkit/backends/cdt/adapter.py +245 -0
- circuitkit/backends/cdt/propagation.py +357 -0
- circuitkit/backends/cdt/pyfunctions/__init__.py +14 -0
- circuitkit/backends/cdt/pyfunctions/cdt_ablations.py +215 -0
- circuitkit/backends/cdt/pyfunctions/cdt_basic.py +235 -0
- circuitkit/backends/cdt/pyfunctions/cdt_core.py +418 -0
- circuitkit/backends/cdt/pyfunctions/cdt_from_source_nodes.py +237 -0
- circuitkit/backends/cdt/pyfunctions/cdt_source_to_target.py +685 -0
- circuitkit/backends/cdt/pyfunctions/faithfulness_ablations.py +251 -0
- circuitkit/backends/cdt/pyfunctions/general.py +314 -0
- circuitkit/backends/cdt/pyfunctions/ioi_dataset.py +958 -0
- circuitkit/backends/cdt/pyfunctions/local_importance.py +809 -0
- circuitkit/backends/cdt/pyfunctions/pathology.py +460 -0
- circuitkit/backends/cdt/pyfunctions/toy_model.py +190 -0
- circuitkit/backends/cdt/pyfunctions/wrappers.py +159 -0
- circuitkit/backends/eap/__init__.py +2 -0
- circuitkit/backends/eap/artifact_export.py +137 -0
- circuitkit/backends/eap/attribute.py +784 -0
- circuitkit/backends/eap/attribute_node.py +1795 -0
- circuitkit/backends/eap/circuit_kit_adapter.py +121 -0
- circuitkit/backends/eap/eap_utils.py +582 -0
- circuitkit/backends/eap/evaluate.py +762 -0
- circuitkit/backends/eap/graph.py +1569 -0
- circuitkit/backends/eap/metrics.py +793 -0
- circuitkit/backends/eap/py.typed +0 -0
- circuitkit/backends/eap/visualization.py +101 -0
- circuitkit/backends/ibcircuit/__init__.py +0 -0
- circuitkit/backends/ibcircuit/artifact_export.py +127 -0
- circuitkit/backends/ibcircuit/ib_noise.py +216 -0
- circuitkit/backends/ibcircuit/ib_utils.py +194 -0
- circuitkit/backends/ibcircuit/model_wrapper.py +537 -0
- circuitkit/backends/ibcircuit/trainer.py +603 -0
- circuitkit/benchmarks/__init__.py +47 -0
- circuitkit/benchmarks/baselines/__init__.py +20 -0
- circuitkit/benchmarks/baselines/gptq.py +200 -0
- circuitkit/benchmarks/baselines/magnitude.py +204 -0
- circuitkit/benchmarks/baselines/random.py +142 -0
- circuitkit/benchmarks/baselines/sparsegpt.py +236 -0
- circuitkit/benchmarks/baselines/wanda.py +263 -0
- circuitkit/benchmarks/benchmark.py +764 -0
- circuitkit/benchmarks/reporting.py +639 -0
- circuitkit/circuit.py +390 -0
- circuitkit/cli/__init__.py +1 -0
- circuitkit/cli/config.py +74 -0
- circuitkit/cli/debug.py +279 -0
- circuitkit/cli/main.py +2208 -0
- circuitkit/cli/utils.py +351 -0
- circuitkit/corruption/__init__.py +61 -0
- circuitkit/corruption/base.py +132 -0
- circuitkit/corruption/color_swap.py +170 -0
- circuitkit/corruption/distractor.py +319 -0
- circuitkit/corruption/distractor_variation.py +313 -0
- circuitkit/corruption/effectiveness.py +297 -0
- circuitkit/corruption/entity_swap.py +288 -0
- circuitkit/corruption/negation.py +364 -0
- circuitkit/corruption/paraphrase.py +390 -0
- circuitkit/corruption/pipeline.py +333 -0
- circuitkit/corruption/position_shift.py +106 -0
- circuitkit/corruption/role_swap.py +381 -0
- circuitkit/corruption/token_swap.py +257 -0
- circuitkit/corruption/validators.py +570 -0
- circuitkit/corruption/voice_swap.py +393 -0
- circuitkit/data/__init__.py +8 -0
- circuitkit/data/adapters/__init__.py +20 -0
- circuitkit/data/adapters/base.py +123 -0
- circuitkit/data/adapters/code.py +106 -0
- circuitkit/data/adapters/conversational.py +167 -0
- circuitkit/data/adapters/forget_retain.py +152 -0
- circuitkit/data/adapters/instruction.py +125 -0
- circuitkit/data/adapters/math.py +133 -0
- circuitkit/data/adapters/mcq.py +194 -0
- circuitkit/data/adapters/pairwise.py +182 -0
- circuitkit/data/adapters/safety_prompt.py +258 -0
- circuitkit/data/auto_detect.py +187 -0
- circuitkit/data/clean_only.py +124 -0
- circuitkit/data/corruption/__init__.py +38 -0
- circuitkit/data/corruption/base.py +239 -0
- circuitkit/data/corruption/benign_rewrite.py +127 -0
- circuitkit/data/corruption/code_syntax_corrupt.py +98 -0
- circuitkit/data/corruption/entity_swap.py +113 -0
- circuitkit/data/corruption/final_answer_swap.py +231 -0
- circuitkit/data/corruption/instruction_swap.py +149 -0
- circuitkit/data/corruption/llm_counterfactual.py +161 -0
- circuitkit/data/corruption/logical_negation.py +96 -0
- circuitkit/data/corruption/math_step_corrupt.py +88 -0
- circuitkit/data/corruption/mcq_choice_swap.py +106 -0
- circuitkit/data/corruption/operand_swap.py +103 -0
- circuitkit/data/corruption/profession_swap.py +125 -0
- circuitkit/data/corruption/resample.py +75 -0
- circuitkit/data/corruption/template.py +195 -0
- circuitkit/data/corruption/template_utils.py +328 -0
- circuitkit/data/corruption/token_swap.py +90 -0
- circuitkit/data/dataset_schema.py +169 -0
- circuitkit/data/eap_dataset.py +98 -0
- circuitkit/data/invariance_groups/__init__.py +33 -0
- circuitkit/data/invariance_groups/builder.py +323 -0
- circuitkit/data/invariance_groups/schema.py +274 -0
- circuitkit/data/normalized.py +259 -0
- circuitkit/data/normalized_task.py +594 -0
- circuitkit/data/task_data/__init__.py +11 -0
- circuitkit/data/task_data/core/TLACDCCorrespondence.py +263 -0
- circuitkit/data/task_data/core/TLACDCEdge.py +113 -0
- circuitkit/data/task_data/core/TLACDCExperiment.py +1052 -0
- circuitkit/data/task_data/core/TLACDCInterpNode.py +96 -0
- circuitkit/data/task_data/core/__init__.py +12 -0
- circuitkit/data/task_data/core/acdc_utils.py +614 -0
- circuitkit/data/task_data/generation/__init__.py +11 -0
- circuitkit/data/task_data/generation/cache.py +267 -0
- circuitkit/data/task_data/generation/manager.py +562 -0
- circuitkit/data/task_data/generation/utils.py +323 -0
- circuitkit/data/task_data/storage/__init__.py +18 -0
- circuitkit/data/task_data/storage/greaterthan/greaterthan_32_ffd33106.json +23 -0
- circuitkit/data/task_data/storage/ioi/ioi_16_8c879ddb.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_32_a432ca4a.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_500_1f7e7324.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_64_3bba747e.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_64_f4a164db.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_8_e87df42e.json +43 -0
- circuitkit/data/task_data/tasks/__init__.py +10 -0
- circuitkit/data/task_data/tasks/binary_align/generate_binary_align.py +1167 -0
- circuitkit/data/task_data/tasks/binary_align/jailbreak_binary.csv +335 -0
- circuitkit/data/task_data/tasks/binary_align/safe_binary.csv +335 -0
- circuitkit/data/task_data/tasks/capital_country/__init__.py +5 -0
- circuitkit/data/task_data/tasks/capital_country/utils.py +395 -0
- circuitkit/data/task_data/tasks/docstring/__init__.py +5 -0
- circuitkit/data/task_data/tasks/docstring/prompts.py +1175 -0
- circuitkit/data/task_data/tasks/docstring/utils.py +282 -0
- circuitkit/data/task_data/tasks/double_io/__init__.py +0 -0
- circuitkit/data/task_data/tasks/double_io/double_io_dataset.py +485 -0
- circuitkit/data/task_data/tasks/gender_bias/__init__.py +5 -0
- circuitkit/data/task_data/tasks/gender_bias/utils.py +396 -0
- circuitkit/data/task_data/tasks/gender_bias/utils2.py +143 -0
- circuitkit/data/task_data/tasks/greaterthan/__init__.py +5 -0
- circuitkit/data/task_data/tasks/greaterthan/utils.py +534 -0
- circuitkit/data/task_data/tasks/hypernymy/__init__.py +5 -0
- circuitkit/data/task_data/tasks/hypernymy/utils.py +326 -0
- circuitkit/data/task_data/tasks/induction/__init__.py +5 -0
- circuitkit/data/task_data/tasks/induction/utils.py +222 -0
- circuitkit/data/task_data/tasks/ioi/__init__.py +8 -0
- circuitkit/data/task_data/tasks/ioi/ioi_dataset.py +962 -0
- circuitkit/data/task_data/tasks/ioi/utils.py +656 -0
- circuitkit/data/task_data/tasks/sva/__init__.py +5 -0
- circuitkit/data/task_data/tasks/sva/utils.py +132 -0
- circuitkit/data/task_data/tasks/wmdp/wmdp_utils.py +296 -0
- circuitkit/data/template.py +392 -0
- circuitkit/data/wikitext_calibration.py +164 -0
- circuitkit/data/worthiness.py +746 -0
- circuitkit/evaluation/__init__.py +83 -0
- circuitkit/evaluation/checkpoint_benchmark.py +857 -0
- circuitkit/evaluation/evaluate.py +929 -0
- circuitkit/evaluation/full.py +556 -0
- circuitkit/evaluation/hf_checkpoint.py +1219 -0
- circuitkit/evaluation/intervention_faithfulness.py +198 -0
- circuitkit/evaluation/lm_eval_simple.py +223 -0
- circuitkit/evaluation/lm_harness.py +681 -0
- circuitkit/evaluation/master_grid.py +307 -0
- circuitkit/evaluation/mmlu_eval.py +208 -0
- circuitkit/evaluation/pillars/__init__.py +28 -0
- circuitkit/evaluation/pillars/ablation.py +404 -0
- circuitkit/evaluation/pillars/baselines.py +902 -0
- circuitkit/evaluation/pillars/causal_patching.py +371 -0
- circuitkit/evaluation/pillars/generalization.py +623 -0
- circuitkit/evaluation/pillars/intervention_reliability.py +325 -0
- circuitkit/evaluation/pillars/robustness.py +854 -0
- circuitkit/evaluation/pillars/stability.py +571 -0
- circuitkit/evaluation/report.py +318 -0
- circuitkit/evaluation/reports/__init__.py +20 -0
- circuitkit/evaluation/reports/aggregator.py +533 -0
- circuitkit/evaluation/reports/robustness_report.py +331 -0
- circuitkit/evaluation/reports/stability_report.py +298 -0
- circuitkit/evaluation/stability_discovery.py +439 -0
- circuitkit/evaluation/transfer.py +510 -0
- circuitkit/evaluation/transfer_analysis.py +315 -0
- circuitkit/evaluation/transfer_visualizer.py +327 -0
- circuitkit/evaluation/weight_based_eval.py +227 -0
- circuitkit/pipeline.py +1000 -0
- circuitkit/quick.py +1157 -0
- circuitkit/selection/__init__.py +54 -0
- circuitkit/selection/cdt_selector.py +64 -0
- circuitkit/selection/eap_gp_selector.py +67 -0
- circuitkit/selection/eap_selector.py +83 -0
- circuitkit/selection/gptq_selector.py +164 -0
- circuitkit/selection/ibcircuit_selector.py +120 -0
- circuitkit/selection/magnitude_selector.py +28 -0
- circuitkit/selection/random_selector.py +16 -0
- circuitkit/selection/relp_selector.py +66 -0
- circuitkit/selection/wanda_selector.py +174 -0
- circuitkit/tasks/__init__.py +40 -0
- circuitkit/tasks/_algorithm_families.py +106 -0
- circuitkit/tasks/_chat.py +262 -0
- circuitkit/tasks/auto_schema.py +526 -0
- circuitkit/tasks/bootstrap.py +93 -0
- circuitkit/tasks/builtins/__init__.py +44 -0
- circuitkit/tasks/builtins/boolq.py +501 -0
- circuitkit/tasks/builtins/capital_country.py +261 -0
- circuitkit/tasks/builtins/double_io.py +380 -0
- circuitkit/tasks/builtins/gender_bias.py +284 -0
- circuitkit/tasks/builtins/glue.py +683 -0
- circuitkit/tasks/builtins/greater_than.py +471 -0
- circuitkit/tasks/builtins/gsm8k.py +563 -0
- circuitkit/tasks/builtins/hypernymy.py +262 -0
- circuitkit/tasks/builtins/ifeval.py +116 -0
- circuitkit/tasks/builtins/ioi.py +447 -0
- circuitkit/tasks/builtins/ioi_acdc.py +323 -0
- circuitkit/tasks/builtins/ioi_legacy.py +473 -0
- circuitkit/tasks/builtins/mmlu.py +1578 -0
- circuitkit/tasks/builtins/sva.py +257 -0
- circuitkit/tasks/builtins/truthfulqa.py +520 -0
- circuitkit/tasks/builtins/winogrande.py +647 -0
- circuitkit/tasks/builtins/winogrande_mc.py +484 -0
- circuitkit/tasks/builtins/wmdp.py +1120 -0
- circuitkit/tasks/generic.py +1405 -0
- circuitkit/tasks/hf_factory.py +480 -0
- circuitkit/tasks/inspect.py +95 -0
- circuitkit/tasks/registry.py +72 -0
- circuitkit/tasks/safety_datasets.py +139 -0
- circuitkit/tasks/specs.py +246 -0
- circuitkit/tasks/type_specs/__init__.py +30 -0
- circuitkit/tasks/type_specs/classification_spec.py +49 -0
- circuitkit/tasks/type_specs/generation_spec.py +47 -0
- circuitkit/tasks/type_specs/mcq_spec.py +49 -0
- circuitkit/tasks/type_specs/qa_spec.py +129 -0
- circuitkit/tasks/type_specs/summarization_spec.py +45 -0
- circuitkit/tasks/type_specs/translation_spec.py +45 -0
- circuitkit/tasks/validator.py +481 -0
- circuitkit/tasks/yaml_loader.py +381 -0
- circuitkit/tooling/__init__.py +7 -0
- circuitkit/tooling/validate_environment.py +61 -0
- circuitkit/utils/__init__.py +0 -0
- circuitkit/utils/artifacts.py +39 -0
- circuitkit/utils/async_processing.py +371 -0
- circuitkit/utils/bootstrap.py +194 -0
- circuitkit/utils/config.py +330 -0
- circuitkit/utils/corruption_validation.py +369 -0
- circuitkit/utils/dataset_cache.py +303 -0
- circuitkit/utils/debug.py +345 -0
- circuitkit/utils/debugging.py +354 -0
- circuitkit/utils/device.py +51 -0
- circuitkit/utils/distributed.py +531 -0
- circuitkit/utils/exceptions.py +355 -0
- circuitkit/utils/logging.py +382 -0
- circuitkit/utils/memory.py +191 -0
- circuitkit/utils/optimization.py +316 -0
- circuitkit/utils/profiling.py +414 -0
- circuitkit/utils/token_utils.py +159 -0
- circuitkit/visualize/__init__.py +79 -0
- circuitkit/visualize/comparison.py +499 -0
- circuitkit/visualize/d3_template.py +1064 -0
- circuitkit/visualize/editor.py +385 -0
- circuitkit/visualize/feature_saliency.py +408 -0
- circuitkit/visualize/gallery.py +392 -0
- circuitkit/visualize/graph_viz.py +884 -0
- circuitkit/visualize/jupyter_suite.py +223 -0
- circuitkit/visualize/plotter.py +246 -0
- circuitkit/visualize/saliency.py +402 -0
- circuitkit/visualize/streamlit_app.py +471 -0
- circuitkit/visualize/theme.py +331 -0
- circuitkit-0.1.0.dist-info/METADATA +192 -0
- circuitkit-0.1.0.dist-info/RECORD +368 -0
- circuitkit-0.1.0.dist-info/WHEEL +5 -0
- circuitkit-0.1.0.dist-info/entry_points.txt +2 -0
- circuitkit-0.1.0.dist-info/licenses/LICENSE.md +73 -0
- circuitkit-0.1.0.dist-info/top_level.txt +1 -0
circuitkit/__init__.py
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
1
|
+
"""
|
|
2
|
+
CircuitKit: A comprehensive toolkit for circuit discovery in transformer models.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
# Version information
|
|
6
|
+
__version__ = "0.1.0"
|
|
7
|
+
__author__ = "Pratinav Seth, Hem Gosalia, Aditya Kasliwal, Vinay Kumar Sankarapu"
|
|
8
|
+
__description__ = "Unified Discover, Evaluate, Intervene toolkit for mechanistic interpretability"
|
|
9
|
+
_API_EXPORTS = {"discover_circuit", "evaluate_circuit", "load_circuit"}
|
|
10
|
+
|
|
11
|
+
# Flat front-door API (circuitkit.quick) — lazily imported so `import circuitkit`
|
|
12
|
+
# stays fast and torch-free until one of these is actually accessed.
|
|
13
|
+
_QUICK_EXPORTS = {
|
|
14
|
+
"load_model",
|
|
15
|
+
"discover",
|
|
16
|
+
"faithfulness",
|
|
17
|
+
"prune",
|
|
18
|
+
"quantize",
|
|
19
|
+
"export_checkpoint",
|
|
20
|
+
"benchmark",
|
|
21
|
+
"load_scores",
|
|
22
|
+
"selective_finetune",
|
|
23
|
+
"visualize_circuit",
|
|
24
|
+
}
|
|
25
|
+
_CIRCUIT_EXPORTS = {"Circuit"}
|
|
26
|
+
_PIPELINE_EXPORTS = {"Pipeline"}
|
|
27
|
+
|
|
28
|
+
# Main exports
|
|
29
|
+
__all__ = [
|
|
30
|
+
# Core dict-config API
|
|
31
|
+
"discover_circuit",
|
|
32
|
+
"evaluate_circuit",
|
|
33
|
+
"load_circuit",
|
|
34
|
+
# Flat front-door API
|
|
35
|
+
"load_model",
|
|
36
|
+
"discover",
|
|
37
|
+
"faithfulness",
|
|
38
|
+
"prune",
|
|
39
|
+
"quantize",
|
|
40
|
+
"export_checkpoint",
|
|
41
|
+
"benchmark",
|
|
42
|
+
"Circuit",
|
|
43
|
+
"load_scores",
|
|
44
|
+
"selective_finetune",
|
|
45
|
+
"visualize_circuit",
|
|
46
|
+
"Pipeline",
|
|
47
|
+
# Task management
|
|
48
|
+
"get_task",
|
|
49
|
+
"list_tasks",
|
|
50
|
+
"register_task",
|
|
51
|
+
# Version info
|
|
52
|
+
"__version__",
|
|
53
|
+
"__author__",
|
|
54
|
+
"__description__",
|
|
55
|
+
]
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def __getattr__(name):
|
|
59
|
+
"""Lazily import heavy API helpers when accessed from the package root."""
|
|
60
|
+
if name in _API_EXPORTS:
|
|
61
|
+
from . import api
|
|
62
|
+
|
|
63
|
+
value = getattr(api, name)
|
|
64
|
+
globals()[name] = value
|
|
65
|
+
return value
|
|
66
|
+
if name in _QUICK_EXPORTS:
|
|
67
|
+
from . import quick
|
|
68
|
+
|
|
69
|
+
value = getattr(quick, name)
|
|
70
|
+
globals()[name] = value
|
|
71
|
+
return value
|
|
72
|
+
if name in _CIRCUIT_EXPORTS:
|
|
73
|
+
from .circuit import Circuit
|
|
74
|
+
|
|
75
|
+
globals()["Circuit"] = Circuit
|
|
76
|
+
return Circuit
|
|
77
|
+
if name in _PIPELINE_EXPORTS:
|
|
78
|
+
from .pipeline import Pipeline
|
|
79
|
+
|
|
80
|
+
globals()["Pipeline"] = Pipeline
|
|
81
|
+
return Pipeline
|
|
82
|
+
if name == "visualize":
|
|
83
|
+
# Deprecated pre-1.0 alias (H1 in the 1.0.0 audit): ``visualize`` was
|
|
84
|
+
# renamed to ``visualize_circuit`` in 1.0.0 and is now a subpackage.
|
|
85
|
+
# Return the (callable) subpackage so ``ck.visualize(...)`` keeps working
|
|
86
|
+
# and warns on use — see circuitkit/visualize/__init__.py. The warning
|
|
87
|
+
# fires on the call, not on attribute access, so it survives the
|
|
88
|
+
# submodule being imported (which shadows any parent __getattr__ shim).
|
|
89
|
+
from . import visualize as _visualize
|
|
90
|
+
|
|
91
|
+
return _visualize
|
|
92
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def __dir__():
|
|
96
|
+
"""Expose lazily-imported names for autocomplete / dir()."""
|
|
97
|
+
return sorted(set(globals()) | set(__all__))
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def _ensure_builtin_tasks():
|
|
101
|
+
"""Register built-in tasks only when task helpers are used."""
|
|
102
|
+
from .tasks.bootstrap import _bootstrap_builtin_tasks
|
|
103
|
+
|
|
104
|
+
return _bootstrap_builtin_tasks()
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def get_task(name):
|
|
108
|
+
"""Get a built-in or registered task by name."""
|
|
109
|
+
_ensure_builtin_tasks()
|
|
110
|
+
from .tasks.registry import get_task as _get_task
|
|
111
|
+
|
|
112
|
+
return _get_task(name)
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def list_tasks():
|
|
116
|
+
"""List registered task names."""
|
|
117
|
+
_ensure_builtin_tasks()
|
|
118
|
+
from .tasks.registry import list_tasks as _list_tasks
|
|
119
|
+
|
|
120
|
+
return _list_tasks()
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def register_task(spec):
|
|
124
|
+
"""Register a custom task specification."""
|
|
125
|
+
_ensure_builtin_tasks()
|
|
126
|
+
from .tasks.registry import register_task as _register_task
|
|
127
|
+
|
|
128
|
+
return _register_task(spec)
|
circuitkit/__main__.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
"""
|
|
2
|
+
CircuitKit Analysis Module
|
|
3
|
+
|
|
4
|
+
Provides analysis tools for circuits including metrics, scoring, and statistical analysis.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from .cross_method_jaccard import CrossMethodJaccardResult, cross_method_jaccard
|
|
8
|
+
from .metrics import * # noqa: F401,F403 - intentional API re-export
|
|
9
|
+
from .scores import * # noqa: F401,F403 - intentional API re-export
|
|
10
|
+
|
|
11
|
+
__all__ = [ # noqa: F405 - names provided via star imports above
|
|
12
|
+
# Metrics
|
|
13
|
+
"compute_metrics",
|
|
14
|
+
# Scores
|
|
15
|
+
"compute_scores",
|
|
16
|
+
# Cross-method comparison (EMNLP 2026, Section 5)
|
|
17
|
+
"cross_method_jaccard",
|
|
18
|
+
"CrossMethodJaccardResult",
|
|
19
|
+
]
|
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
"""Cross-method circuit comparison via top-K Jaccard.
|
|
2
|
+
|
|
3
|
+
Used by the EMNLP 2026 submission's Finding B (Section 5) to
|
|
4
|
+
quantify component-level disagreement across discovery methods
|
|
5
|
+
discovering circuits for the same task on the same model.
|
|
6
|
+
|
|
7
|
+
Example
|
|
8
|
+
-------
|
|
9
|
+
>>> from circuitkit.analysis.cross_method_jaccard import (
|
|
10
|
+
... cross_method_jaccard,
|
|
11
|
+
... )
|
|
12
|
+
>>> # circuits: dict mapping method name to a CircuitScores artifact
|
|
13
|
+
>>> # (or a list of node names, or a dict node_name -> score)
|
|
14
|
+
>>> result = cross_method_jaccard(circuits, top_k=10)
|
|
15
|
+
>>> # result.matrix: numpy-style nested list of pairwise Jaccards
|
|
16
|
+
>>> # result.top_nodes: dict method -> top_k node names
|
|
17
|
+
>>> # result.range: (min, max) pairwise Jaccard over off-diagonal cells
|
|
18
|
+
|
|
19
|
+
The Jaccard at top-K is the standard descriptive statistic for the
|
|
20
|
+
"do methods agree on the circuit?" question reported in the EMNLP
|
|
21
|
+
paper, derived from the workshop result on GPT-2 IOI where the same
|
|
22
|
+
analysis showed Jaccard 0.11 to 1.00 across 8 EAP-family variants.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
from __future__ import annotations
|
|
26
|
+
|
|
27
|
+
from dataclasses import dataclass
|
|
28
|
+
from typing import Dict, List, Tuple, Union
|
|
29
|
+
|
|
30
|
+
CircuitLike = Union[Dict[str, float], List[str], object]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _extract_top_k_nodes(circuit: CircuitLike, top_k: int) -> List[str]:
|
|
34
|
+
"""Coerce a circuit-like input into a top-K node-name list.
|
|
35
|
+
|
|
36
|
+
Accepts:
|
|
37
|
+
- dict[str, float]: node_name -> score; top-K by absolute score
|
|
38
|
+
- list[str]: assumed to be a pre-ranked list; truncated to top-K
|
|
39
|
+
- object with `.node_scores` attribute (e.g. CircuitScores): dict-like
|
|
40
|
+
"""
|
|
41
|
+
if hasattr(circuit, "node_scores"):
|
|
42
|
+
node_scores = circuit.node_scores
|
|
43
|
+
else:
|
|
44
|
+
node_scores = circuit
|
|
45
|
+
if isinstance(node_scores, dict):
|
|
46
|
+
ranked = sorted(node_scores.items(), key=lambda kv: abs(float(kv[1])), reverse=True)
|
|
47
|
+
return [n for n, _ in ranked[:top_k]]
|
|
48
|
+
if isinstance(node_scores, list):
|
|
49
|
+
return list(node_scores[:top_k])
|
|
50
|
+
raise TypeError(f"Unsupported circuit type for jaccard: {type(circuit)!r}")
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _jaccard(a: List[str], b: List[str]) -> float:
|
|
54
|
+
sa, sb = set(a), set(b)
|
|
55
|
+
if not sa and not sb:
|
|
56
|
+
return 1.0
|
|
57
|
+
return len(sa & sb) / len(sa | sb)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
@dataclass
|
|
61
|
+
class CrossMethodJaccardResult:
|
|
62
|
+
methods: List[str]
|
|
63
|
+
top_k: int
|
|
64
|
+
matrix: List[List[float]]
|
|
65
|
+
top_nodes: Dict[str, List[str]]
|
|
66
|
+
range: Tuple[float, float]
|
|
67
|
+
n_pairs: int
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def cross_method_jaccard(
|
|
71
|
+
circuits: Dict[str, CircuitLike],
|
|
72
|
+
top_k: int = 10,
|
|
73
|
+
) -> CrossMethodJaccardResult:
|
|
74
|
+
"""Compute the symmetric pairwise Jaccard matrix between method
|
|
75
|
+
circuits at top-K nodes.
|
|
76
|
+
|
|
77
|
+
Parameters
|
|
78
|
+
----------
|
|
79
|
+
circuits : dict
|
|
80
|
+
Maps method name to either a CircuitScores artifact, a dict of
|
|
81
|
+
node_name to score, or a pre-ranked list of node names.
|
|
82
|
+
top_k : int
|
|
83
|
+
Truncation depth for each method's circuit. Default 10 to match
|
|
84
|
+
the EMNLP paper's protocol.
|
|
85
|
+
|
|
86
|
+
Returns
|
|
87
|
+
-------
|
|
88
|
+
CrossMethodJaccardResult
|
|
89
|
+
Includes the pairwise Jaccard matrix, the top-K node sets per
|
|
90
|
+
method, and the off-diagonal range (min, max).
|
|
91
|
+
"""
|
|
92
|
+
methods = sorted(circuits.keys())
|
|
93
|
+
n = len(methods)
|
|
94
|
+
top_nodes = {m: _extract_top_k_nodes(circuits[m], top_k) for m in methods}
|
|
95
|
+
matrix = [[1.0] * n for _ in range(n)]
|
|
96
|
+
off_diag: List[float] = []
|
|
97
|
+
for i, mi in enumerate(methods):
|
|
98
|
+
for j, mj in enumerate(methods):
|
|
99
|
+
if i <= j:
|
|
100
|
+
continue
|
|
101
|
+
j_val = _jaccard(top_nodes[mi], top_nodes[mj])
|
|
102
|
+
matrix[i][j] = j_val
|
|
103
|
+
matrix[j][i] = j_val
|
|
104
|
+
off_diag.append(j_val)
|
|
105
|
+
rng = (min(off_diag), max(off_diag)) if off_diag else (0.0, 0.0)
|
|
106
|
+
return CrossMethodJaccardResult(
|
|
107
|
+
methods=methods,
|
|
108
|
+
top_k=top_k,
|
|
109
|
+
matrix=matrix,
|
|
110
|
+
top_nodes=top_nodes,
|
|
111
|
+
range=rng,
|
|
112
|
+
n_pairs=len(off_diag),
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
__all__ = ["cross_method_jaccard", "CrossMethodJaccardResult"]
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
# FILE: circuitkit/analysis/metrics.py
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def calculate_faithfulness(original_output, pruned_output):
|
|
5
|
+
"""
|
|
6
|
+
Calculate faithfulness metric between original and pruned outputs.
|
|
7
|
+
Uses multiple metrics for comprehensive evaluation.
|
|
8
|
+
"""
|
|
9
|
+
import torch
|
|
10
|
+
import torch.nn.functional as F
|
|
11
|
+
|
|
12
|
+
# Ensure tensors are the same shape
|
|
13
|
+
if original_output.shape != pruned_output.shape:
|
|
14
|
+
min_size = min(original_output.numel(), pruned_output.numel())
|
|
15
|
+
original_output = original_output.flatten()[:min_size]
|
|
16
|
+
pruned_output = pruned_output.flatten()[:min_size]
|
|
17
|
+
|
|
18
|
+
# L2 norm difference
|
|
19
|
+
l2_diff = (original_output - pruned_output).pow(2).sum().item()
|
|
20
|
+
|
|
21
|
+
# KL divergence
|
|
22
|
+
try:
|
|
23
|
+
kl_div = F.kl_div(
|
|
24
|
+
F.log_softmax(pruned_output, dim=-1),
|
|
25
|
+
F.softmax(original_output, dim=-1),
|
|
26
|
+
reduction="sum",
|
|
27
|
+
).item()
|
|
28
|
+
except (RuntimeError, ValueError):
|
|
29
|
+
kl_div = float("inf")
|
|
30
|
+
|
|
31
|
+
# Cosine similarity
|
|
32
|
+
cos_sim = F.cosine_similarity(
|
|
33
|
+
original_output.flatten().unsqueeze(0), pruned_output.flatten().unsqueeze(0)
|
|
34
|
+
).item()
|
|
35
|
+
|
|
36
|
+
# Relative difference
|
|
37
|
+
rel_diff = torch.abs(original_output - pruned_output).sum().item() / (
|
|
38
|
+
torch.abs(original_output).sum().item() + 1e-8
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
return {
|
|
42
|
+
"l2_difference": l2_diff,
|
|
43
|
+
"kl_divergence": kl_div,
|
|
44
|
+
"cosine_similarity": cos_sim,
|
|
45
|
+
"relative_difference": rel_diff,
|
|
46
|
+
"faithfulness_score": 1.0 - min(rel_diff, 1.0), # Higher is better
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def calculate_complexity(graph):
|
|
51
|
+
"""
|
|
52
|
+
Calculates the complexity of a circuit, e.g., by node or edge count.
|
|
53
|
+
"""
|
|
54
|
+
return {"node_count": graph.number_of_nodes(), "edge_count": graph.number_of_edges()}
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
# FILE: circuitkit/analysis/scores.py
|
|
2
|
+
|
|
3
|
+
from collections import defaultdict
|
|
4
|
+
|
|
5
|
+
from ..backends.acdc.types import PruneScores
|
|
6
|
+
from ..backends.acdc.utils.patchable_model import PatchableModel
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def calculate_node_scores_from_edges(
|
|
10
|
+
p_model: PatchableModel, edge_prune_scores: PruneScores
|
|
11
|
+
) -> dict[str, float]:
|
|
12
|
+
"""
|
|
13
|
+
Calculates an importance score for each source node based on its outgoing edges.
|
|
14
|
+
|
|
15
|
+
The score for a node is the average of the absolute scores of all its
|
|
16
|
+
outgoing edges. Nodes with no outgoing edges will not be included.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
p_model: The patchable model, used to access the graph of nodes and edges.
|
|
20
|
+
edge_prune_scores: A dictionary mapping destination modules to tensors of
|
|
21
|
+
edge importance scores.
|
|
22
|
+
|
|
23
|
+
Returns:
|
|
24
|
+
A dictionary mapping source node names to their calculated importance score.
|
|
25
|
+
"""
|
|
26
|
+
# Use absolute values of scores as importance can be positive or negative in EAP
|
|
27
|
+
abs_edge_scores = {mod: scores.abs() for mod, scores in edge_prune_scores.items()}
|
|
28
|
+
|
|
29
|
+
# Group edge scores by their source node
|
|
30
|
+
outgoing_scores_by_node = defaultdict(list)
|
|
31
|
+
for edge in p_model.edges:
|
|
32
|
+
# edge.prune_score looks up the score for this specific edge
|
|
33
|
+
score = edge.prune_score(abs_edge_scores).item()
|
|
34
|
+
outgoing_scores_by_node[edge.src.name].append(score)
|
|
35
|
+
|
|
36
|
+
# Calculate the average score for each node. Use Python's built-in sum/len
|
|
37
|
+
# rather than np.mean() so values stay plain Python floats — numpy scalars
|
|
38
|
+
# (numpy.float64) are rejected by torch.load with weights_only=True.
|
|
39
|
+
node_scores = {}
|
|
40
|
+
for node_name, scores in outgoing_scores_by_node.items():
|
|
41
|
+
if scores:
|
|
42
|
+
node_scores[node_name] = sum(scores) / len(scores)
|
|
43
|
+
|
|
44
|
+
return node_scores
|