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