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