bergson 2.2.0__tar.gz → 2.2.2__tar.gz
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.
- {bergson-2.2.0 → bergson-2.2.2}/PKG-INFO +28 -34
- {bergson-2.2.0 → bergson-2.2.2}/README.md +27 -33
- {bergson-2.2.0 → bergson-2.2.2}/bergson/__init__.py +1 -1
- {bergson-2.2.0 → bergson-2.2.2}/bergson/config/config.py +10 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/eigenvectors.py +108 -29
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/hessian_approximations.py +48 -12
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/kfac.py +3 -1
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/sharded_computation.py +19 -7
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/cli.py +59 -49
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/data_stream.py +72 -35
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/fsdp.py +11 -5
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/metasmoothness.py +21 -8
- bergson-2.2.2/bergson/magic/shard_load.py +119 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/trainer.py +16 -4
- bergson-2.2.2/bergson/score/output_influence.py +154 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/score/score.py +117 -2
- {bergson-2.2.0 → bergson-2.2.2}/bergson/score/scorer.py +4 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/worker_utils.py +21 -8
- {bergson-2.2.0 → bergson-2.2.2}/bergson/validate.py +19 -31
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/PKG-INFO +28 -34
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/SOURCES.txt +5 -0
- {bergson-2.2.0 → bergson-2.2.2}/pyproject.toml +1 -1
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_ddp.py +2 -2
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_magic.py +55 -25
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_metasmoothness.py +3 -2
- bergson-2.2.2/tests/test_output_influence.py +191 -0
- bergson-2.2.2/tests/test_padding.py +50 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_per_query_magic.py +16 -18
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_per_token_lds.py +36 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_loss_padding.py +2 -2
- bergson-2.2.2/tests/test_shard_load.py +46 -0
- {bergson-2.2.0 → bergson-2.2.2}/LICENSE +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/__main__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/adam_preconditioner.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/approx_unrolling_math.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/pipeline.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/precompute_checkpoints.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/segment_aggregation.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/train_cfg_io.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/build.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/builder.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/commands.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/trackstar.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/trak.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collection.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/collector.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/dist_autocorrelation_gradient_collector.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/gradient_collectors.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/in_memory_collector.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/projection_matrix.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/config/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/config/config_io.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/config/validation.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/data.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/diagnose.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/distributed.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/format.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/gradients.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/apply_hessian.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/astra.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/autocorrelation.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/inversion.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/pipeline.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/preconditioner.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/shampoo.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/tkfac.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/huggingface.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/config.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/dtensor_patch.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/grad_accum.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/optim.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/rtl_tqdm.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/score_plot.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/swap.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/moe.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/process_autocorrelation.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/process_grads.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/query/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/query/attributor.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/query/faiss_index.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/query/query_index.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/facts.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/generate.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/recall.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/score/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/score/candidates.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/score/score_writer.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/__init__.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/batch_size.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/csv_writer.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/gradcheck.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/load_from_optimizer.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/logger.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/logging.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/math.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/peft.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/step_state.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/trainer_export.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/utils.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/dependency_links.txt +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/entry_points.txt +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/requires.txt +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/top_level.txt +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/setup.cfg +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_adam_state_loading.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_advantages.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_approx_unrolling_checkpoint_parse.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_attention_head_normalizers.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_attribute_tokens.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_attributor.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_bank_loss_cache.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_banked_model_loading.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_batch_size_invariance.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_build.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_builder.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_candidates.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_ckpt_avg_query_grads.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_cli.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_compute_lambda.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_config_runner.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_contrast.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_data.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_diagnose.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_batch_budget.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_cap.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_failure.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_magic.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_epoch_shuffle.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_faiss_ann_cpu_load.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_force_math_sdp.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_format.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_global_kfac.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_global_projection.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_grad_clipping.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_gradcheck.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_gradients.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_hessian.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_launch_devices.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_log_fn.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_logit_scale.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_moe.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_multi_query_validate.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_multinode.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_muon.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_normalizer.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_outer_product_gradients.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_pretokenized.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_inner_products.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_matrix.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_settings.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_batching_invariance.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_set_config.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_recall.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_reduce.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_save_mode_final.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_score.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_score_plot.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_score_writer_dist.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_fisher_normalization.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_optimizer_variants.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_resume_polarity.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_trainer_integration.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_step_matrix.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_step_state.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_tokenize.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_trainer_callback.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_trak.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_truncation.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_validate_filter.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_validate_routing.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_validation_config.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_wandb_logging.py +0 -0
- {bergson-2.2.0 → bergson-2.2.2}/tests/test_weighted_ce.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: bergson
|
|
3
|
-
Version: 2.2.
|
|
3
|
+
Version: 2.2.2
|
|
4
4
|
Summary: Tracing the memory of neural nets with data attribution
|
|
5
5
|
License: MIT License
|
|
6
6
|
Keywords: interpretability,explainable-ai
|
|
@@ -60,7 +60,7 @@ Requires-Dist: matplotlib; extra == "plot"
|
|
|
60
60
|
Dynamic: license-file
|
|
61
61
|
|
|
62
62
|
# Bergson
|
|
63
|
-
Bergson is a python library which provides scalable, state-of-the-art data attribution methods for large language models. Data attribution methods estimate the effect on a behavior of interest of removing data points from a model's training corpus. We support [EK-FAC](https://arxiv.org/abs/2308.03296) (2023), [TrackStar](https://arxiv.org/abs/2410.17413v3) (2024), [SOURCE](https://arxiv.org/abs/2405.12186) (2024), [MAGIC](https://arxiv.org/abs/2504.16430) (2025), gradient cosine similarity, and more.
|
|
63
|
+
Bergson is a python library which provides scalable, state-of-the-art data attribution methods for large language models. Data attribution methods estimate the effect on a behavior of interest of removing data points from a model's training corpus. We support [EK-FAC](https://arxiv.org/abs/2308.03296) (2023), [TrackStar](https://arxiv.org/abs/2410.17413v3) (2024), [SOURCE](https://arxiv.org/abs/2405.12186) (2024), [MAGIC](https://arxiv.org/abs/2504.16430) (2025), [ASTRA](https://arxiv.org/abs/2507.14740) (2025), gradient cosine similarity, and more.
|
|
64
64
|
|
|
65
65
|
We try to make research and application straightforward. You can reproduce any of your runs with a single line, use a few CLI flags to save and re-use useful intermediate artifacts, do everything using just the CLI or just code, train models here or use existing ones, tune and evaluate your methods, and scale them up to multi-node runs with 70B+ parameters. More information is available in the [Bergson](https://bergson.readthedocs.io) docs.
|
|
66
66
|
|
|
@@ -86,23 +86,23 @@ pip install -e .
|
|
|
86
86
|
|
|
87
87
|
Indicative performance of data attribution methods in the finetuning regime - see the [leaderboard](https://bergson.readthedocs.io/en/latest/leaderboard.html) for more methods.
|
|
88
88
|
|
|
89
|
-
[Linear datamodeling score](https://arxiv.org/abs/2303.14186) (LDS) is the accuracy of a method for producing global data rankings by influence. The query loss difference (QLD) shows how much model loss for a held-out query can be increased by retraining without the most highly ranked data by influence (here the top 1%), compared to a random removal baseline.
|
|
89
|
+
[Linear datamodeling score](https://arxiv.org/abs/2303.14186) (LDS) is the accuracy of a method for producing global data rankings by influence. The query loss difference (QLD) shows how much model loss for a held-out query can be increased by retraining without the most highly ranked data by influence (here the top 1%), compared to a random removal baseline. The per-token QLD removes the loss terms of the top 1% of training tokens, ranked by the forward-mode equivalent of each method's scores.
|
|
90
90
|
|
|
91
|
-
| Method |
|
|
92
|
-
|
|
93
|
-
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] |
|
|
94
|
-
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] |
|
|
95
|
-
|
|
|
96
|
-
|
|
|
97
|
-
|
|
|
98
|
-
| [
|
|
99
|
-
|
|
|
100
|
-
|
|
|
101
|
-
|
|
|
102
|
-
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] |
|
|
103
|
-
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] |
|
|
91
|
+
| Method | QLD | LDS | Per-token QLD |
|
|
92
|
+
|:---|:---:|:---:|:---:|
|
|
93
|
+
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] | 1.375 [1.348, 1.403] |
|
|
94
|
+
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] | 1.082 [1.047, 1.116] |
|
|
95
|
+
| Eigenvalue-corrected Shampoo | 0.071 [0.060, 0.082] | 0.517 [0.491, 0.539] | 0.922 [0.887, 0.957] |
|
|
96
|
+
| SOURCE (Adam, EK-FAC) | 0.071 [0.060, 0.084] | 0.473 [0.446, 0.498] | 0.908 [0.877, 0.942] |
|
|
97
|
+
| EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
|
|
98
|
+
| BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
|
|
99
|
+
| Semantic search ([Qwen3 8B](https://huggingface.co/Qwen/Qwen3-Embedding-8B)) | 0.049 [0.038, 0.061] | 0.132 [0.093, 0.169] | 0.008 [0.006, 0.010] |
|
|
100
|
+
| TrackStar (no optimizer correction, projection 64) | 0.045 [0.036, 0.055] | 0.270 [0.240, 0.295] | 0.489 [0.465, 0.515] |
|
|
101
|
+
| TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
|
|
102
|
+
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
|
|
103
|
+
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
|
|
104
104
|
|
|
105
|
-
Results for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
105
|
+
Results with 95% confidence intervals for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
106
106
|
|
|
107
107
|
## Functionality
|
|
108
108
|
|
|
@@ -110,7 +110,7 @@ Results for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dr
|
|
|
110
110
|
|
|
111
111
|
`bergson magic` runs a powerful attribution method that backpropagates through the training process to compute the gradient of a model behavior loss with respect to a weighting placed on each training item. It is powered by our twice-differentiable trainer, which can also be called directly using `bergson train`.
|
|
112
112
|
|
|
113
|
-
**Note: unrolled differentiation efficacy is proportional to [metasmoothness](https://bergson.readthedocs.io/en/latest/magic.html#metasmoothness)
|
|
113
|
+
**Note: unrolled differentiation efficacy is proportional to [metasmoothness](https://bergson.readthedocs.io/en/latest/magic.html#metasmoothness), which is low in some settings, including early pretraining steps. Check your run's estimated metasmoothness with `bergson metasmoothness`.**
|
|
114
114
|
|
|
115
115
|
`bergson approxunrolling` is an approximation of `bergson magic` that uses a handful of training checkpoints to run the multi-step SOURCE attribution pipeline. This is roughly equivalent to an influence function averaged over several checkpoints.
|
|
116
116
|
|
|
@@ -120,45 +120,39 @@ To build a train‑time gradient store, use our HF Trainer callback. This will i
|
|
|
120
120
|
|
|
121
121
|
Bergson supports on-disk gradient stores and on-the-fly queries, and per-token and per-sequence attribution.
|
|
122
122
|
|
|
123
|
-
`bergson trackstar`, `bergson ekfac`, and `bergson trak` all orchestrate multi-step attribution recipes over a model checkpoint. `bergson trak` also supports [ensembling](https://arxiv.org/abs/2303.14186) over independently trained models.
|
|
123
|
+
The CLI commands `bergson trackstar`, `bergson ekfac`, and `bergson trak` all orchestrate multi-step attribution recipes over a model checkpoint. `bergson trak` also supports [ensembling](https://arxiv.org/abs/2303.14186) over independently trained models.
|
|
124
124
|
|
|
125
125
|
At a lower level, you can build your own gradient store for efficient serial queries using `bergson build`. Collection-time gradient compression makes the store space-efficient, and a FAISS integration enables fast KNN search over large stores - see `bergson query`, or `Attributor` in the programmatic interface. For small queries and methods that don't use gradient compression (e.g., EK-FAC), you can score a dataset in a single pass using an in-memory query index of precomputed gradients. Dataset items may be scored using max, mean, and individual scoring strategies, enabling [LESS](https://arxiv.org/pdf/2402.04333)-style data filtering. See `bergson score` and `bergson build`.
|
|
126
126
|
|
|
127
127
|
Per-module and per-attention head gradients can be extracted from the store.
|
|
128
128
|
|
|
129
|
-
**Note: influence functions
|
|
129
|
+
**Note: influence functions are sensitive to the Hessian approximation damping hyperparameter in tiny models. Untuned hyperparameters can result in a suboptimal linear datamodeling score.**
|
|
130
130
|
|
|
131
131
|
### Evaluate
|
|
132
132
|
|
|
133
133
|
Use `bergson validate` and `bergson recall` to compute LDS and recall@k metrics respectively.
|
|
134
134
|
|
|
135
|
-
#
|
|
135
|
+
# Getting Started
|
|
136
136
|
|
|
137
|
-
There are many example YAMLs in the `examples` directory.
|
|
137
|
+
There are many example YAMLs in the `examples` directory, including various paper experiment replications. Use `bergson <yaml_path>` to run them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
|
|
138
138
|
|
|
139
139
|
```bash
|
|
140
140
|
bergson examples/magic/gpt2_wikitext_tiny.yaml
|
|
141
141
|
```
|
|
142
142
|
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
| | Notebook | Description |
|
|
146
|
-
|---|----------|-------------|
|
|
147
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
148
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
149
|
-
|
|
150
|
-
Construct and query an on-disk index of randomly projected gradients:
|
|
143
|
+
You can use the same fields used to specify experiments in the YAMLs to run experiments directly in the Bergson CLI. For example, to construct and query an on-disk index of randomly projected gradients from the CLI:
|
|
151
144
|
|
|
152
145
|
```bash
|
|
153
146
|
bergson build runs/index --model EleutherAI/pythia-14m --dataset NeelNanda/pile-10k --truncation --token_batch_size 4096 --projection_dim 16
|
|
154
147
|
bergson query --index runs/index --unit_norm
|
|
155
148
|
```
|
|
156
149
|
|
|
157
|
-
|
|
150
|
+
Or check out a notebook for programmatic usage:
|
|
158
151
|
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
152
|
+
| | Notebook | Description |
|
|
153
|
+
|---|----------|-------------|
|
|
154
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
155
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
162
156
|
|
|
163
157
|
# Development
|
|
164
158
|
|
|
@@ -1,5 +1,5 @@
|
|
|
1
1
|
# Bergson
|
|
2
|
-
Bergson is a python library which provides scalable, state-of-the-art data attribution methods for large language models. Data attribution methods estimate the effect on a behavior of interest of removing data points from a model's training corpus. We support [EK-FAC](https://arxiv.org/abs/2308.03296) (2023), [TrackStar](https://arxiv.org/abs/2410.17413v3) (2024), [SOURCE](https://arxiv.org/abs/2405.12186) (2024), [MAGIC](https://arxiv.org/abs/2504.16430) (2025), gradient cosine similarity, and more.
|
|
2
|
+
Bergson is a python library which provides scalable, state-of-the-art data attribution methods for large language models. Data attribution methods estimate the effect on a behavior of interest of removing data points from a model's training corpus. We support [EK-FAC](https://arxiv.org/abs/2308.03296) (2023), [TrackStar](https://arxiv.org/abs/2410.17413v3) (2024), [SOURCE](https://arxiv.org/abs/2405.12186) (2024), [MAGIC](https://arxiv.org/abs/2504.16430) (2025), [ASTRA](https://arxiv.org/abs/2507.14740) (2025), gradient cosine similarity, and more.
|
|
3
3
|
|
|
4
4
|
We try to make research and application straightforward. You can reproduce any of your runs with a single line, use a few CLI flags to save and re-use useful intermediate artifacts, do everything using just the CLI or just code, train models here or use existing ones, tune and evaluate your methods, and scale them up to multi-node runs with 70B+ parameters. More information is available in the [Bergson](https://bergson.readthedocs.io) docs.
|
|
5
5
|
|
|
@@ -25,23 +25,23 @@ pip install -e .
|
|
|
25
25
|
|
|
26
26
|
Indicative performance of data attribution methods in the finetuning regime - see the [leaderboard](https://bergson.readthedocs.io/en/latest/leaderboard.html) for more methods.
|
|
27
27
|
|
|
28
|
-
[Linear datamodeling score](https://arxiv.org/abs/2303.14186) (LDS) is the accuracy of a method for producing global data rankings by influence. The query loss difference (QLD) shows how much model loss for a held-out query can be increased by retraining without the most highly ranked data by influence (here the top 1%), compared to a random removal baseline.
|
|
28
|
+
[Linear datamodeling score](https://arxiv.org/abs/2303.14186) (LDS) is the accuracy of a method for producing global data rankings by influence. The query loss difference (QLD) shows how much model loss for a held-out query can be increased by retraining without the most highly ranked data by influence (here the top 1%), compared to a random removal baseline. The per-token QLD removes the loss terms of the top 1% of training tokens, ranked by the forward-mode equivalent of each method's scores.
|
|
29
29
|
|
|
30
|
-
| Method |
|
|
31
|
-
|
|
32
|
-
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] |
|
|
33
|
-
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] |
|
|
34
|
-
|
|
|
35
|
-
|
|
|
36
|
-
|
|
|
37
|
-
| [
|
|
38
|
-
|
|
|
39
|
-
|
|
|
40
|
-
|
|
|
41
|
-
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] |
|
|
42
|
-
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] |
|
|
30
|
+
| Method | QLD | LDS | Per-token QLD |
|
|
31
|
+
|:---|:---:|:---:|:---:|
|
|
32
|
+
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] | 1.375 [1.348, 1.403] |
|
|
33
|
+
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] | 1.082 [1.047, 1.116] |
|
|
34
|
+
| Eigenvalue-corrected Shampoo | 0.071 [0.060, 0.082] | 0.517 [0.491, 0.539] | 0.922 [0.887, 0.957] |
|
|
35
|
+
| SOURCE (Adam, EK-FAC) | 0.071 [0.060, 0.084] | 0.473 [0.446, 0.498] | 0.908 [0.877, 0.942] |
|
|
36
|
+
| EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
|
|
37
|
+
| BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
|
|
38
|
+
| Semantic search ([Qwen3 8B](https://huggingface.co/Qwen/Qwen3-Embedding-8B)) | 0.049 [0.038, 0.061] | 0.132 [0.093, 0.169] | 0.008 [0.006, 0.010] |
|
|
39
|
+
| TrackStar (no optimizer correction, projection 64) | 0.045 [0.036, 0.055] | 0.270 [0.240, 0.295] | 0.489 [0.465, 0.515] |
|
|
40
|
+
| TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
|
|
41
|
+
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
|
|
42
|
+
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
|
|
43
43
|
|
|
44
|
-
Results for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
44
|
+
Results with 95% confidence intervals for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
45
45
|
|
|
46
46
|
## Functionality
|
|
47
47
|
|
|
@@ -49,7 +49,7 @@ Results for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dr
|
|
|
49
49
|
|
|
50
50
|
`bergson magic` runs a powerful attribution method that backpropagates through the training process to compute the gradient of a model behavior loss with respect to a weighting placed on each training item. It is powered by our twice-differentiable trainer, which can also be called directly using `bergson train`.
|
|
51
51
|
|
|
52
|
-
**Note: unrolled differentiation efficacy is proportional to [metasmoothness](https://bergson.readthedocs.io/en/latest/magic.html#metasmoothness)
|
|
52
|
+
**Note: unrolled differentiation efficacy is proportional to [metasmoothness](https://bergson.readthedocs.io/en/latest/magic.html#metasmoothness), which is low in some settings, including early pretraining steps. Check your run's estimated metasmoothness with `bergson metasmoothness`.**
|
|
53
53
|
|
|
54
54
|
`bergson approxunrolling` is an approximation of `bergson magic` that uses a handful of training checkpoints to run the multi-step SOURCE attribution pipeline. This is roughly equivalent to an influence function averaged over several checkpoints.
|
|
55
55
|
|
|
@@ -59,45 +59,39 @@ To build a train‑time gradient store, use our HF Trainer callback. This will i
|
|
|
59
59
|
|
|
60
60
|
Bergson supports on-disk gradient stores and on-the-fly queries, and per-token and per-sequence attribution.
|
|
61
61
|
|
|
62
|
-
`bergson trackstar`, `bergson ekfac`, and `bergson trak` all orchestrate multi-step attribution recipes over a model checkpoint. `bergson trak` also supports [ensembling](https://arxiv.org/abs/2303.14186) over independently trained models.
|
|
62
|
+
The CLI commands `bergson trackstar`, `bergson ekfac`, and `bergson trak` all orchestrate multi-step attribution recipes over a model checkpoint. `bergson trak` also supports [ensembling](https://arxiv.org/abs/2303.14186) over independently trained models.
|
|
63
63
|
|
|
64
64
|
At a lower level, you can build your own gradient store for efficient serial queries using `bergson build`. Collection-time gradient compression makes the store space-efficient, and a FAISS integration enables fast KNN search over large stores - see `bergson query`, or `Attributor` in the programmatic interface. For small queries and methods that don't use gradient compression (e.g., EK-FAC), you can score a dataset in a single pass using an in-memory query index of precomputed gradients. Dataset items may be scored using max, mean, and individual scoring strategies, enabling [LESS](https://arxiv.org/pdf/2402.04333)-style data filtering. See `bergson score` and `bergson build`.
|
|
65
65
|
|
|
66
66
|
Per-module and per-attention head gradients can be extracted from the store.
|
|
67
67
|
|
|
68
|
-
**Note: influence functions
|
|
68
|
+
**Note: influence functions are sensitive to the Hessian approximation damping hyperparameter in tiny models. Untuned hyperparameters can result in a suboptimal linear datamodeling score.**
|
|
69
69
|
|
|
70
70
|
### Evaluate
|
|
71
71
|
|
|
72
72
|
Use `bergson validate` and `bergson recall` to compute LDS and recall@k metrics respectively.
|
|
73
73
|
|
|
74
|
-
#
|
|
74
|
+
# Getting Started
|
|
75
75
|
|
|
76
|
-
There are many example YAMLs in the `examples` directory.
|
|
76
|
+
There are many example YAMLs in the `examples` directory, including various paper experiment replications. Use `bergson <yaml_path>` to run them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
|
|
77
77
|
|
|
78
78
|
```bash
|
|
79
79
|
bergson examples/magic/gpt2_wikitext_tiny.yaml
|
|
80
80
|
```
|
|
81
81
|
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
| | Notebook | Description |
|
|
85
|
-
|---|----------|-------------|
|
|
86
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
87
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
88
|
-
|
|
89
|
-
Construct and query an on-disk index of randomly projected gradients:
|
|
82
|
+
You can use the same fields used to specify experiments in the YAMLs to run experiments directly in the Bergson CLI. For example, to construct and query an on-disk index of randomly projected gradients from the CLI:
|
|
90
83
|
|
|
91
84
|
```bash
|
|
92
85
|
bergson build runs/index --model EleutherAI/pythia-14m --dataset NeelNanda/pile-10k --truncation --token_batch_size 4096 --projection_dim 16
|
|
93
86
|
bergson query --index runs/index --unit_norm
|
|
94
87
|
```
|
|
95
88
|
|
|
96
|
-
|
|
89
|
+
Or check out a notebook for programmatic usage:
|
|
97
90
|
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
91
|
+
| | Notebook | Description |
|
|
92
|
+
|---|----------|-------------|
|
|
93
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
94
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
101
95
|
|
|
102
96
|
# Development
|
|
103
97
|
|
|
@@ -885,6 +885,16 @@ class ScoreConfig(Serializable):
|
|
|
885
885
|
capability (e.g. in influence functions). False for unrolled
|
|
886
886
|
differentiation."""
|
|
887
887
|
|
|
888
|
+
token_influence: Literal["gradient", "output"] = "gradient"
|
|
889
|
+
"""What each row scores with ``attribute_tokens``. ``gradient`` scores row
|
|
890
|
+
``t`` by the per-token gradient at position ``t``, which is position ``t``'s
|
|
891
|
+
effect on the loss of every later token. ``output`` scores row ``t`` by the
|
|
892
|
+
loss on token ``t + 1`` alone, the output token influence of Grosse et al.
|
|
893
|
+
(2023). Both kinds of row sum to the per-example score, so without
|
|
894
|
+
``attribute_tokens``, ``output`` only changes how that score is computed.
|
|
895
|
+
``output`` costs one forward-mode pass per query column, so aggregate the
|
|
896
|
+
query when you can, and needs an unprojected query and dot-product scoring."""
|
|
897
|
+
|
|
888
898
|
candidates: CandidateConfig = field(default_factory=CandidateConfig)
|
|
889
899
|
"""Score only the rows an earlier run ranked highest. The store then has one
|
|
890
900
|
row per candidate; ``candidates.npy`` beside it holds their training rows."""
|
|
@@ -85,8 +85,9 @@ class LambdaCollector(HookCollectorBase):
|
|
|
85
85
|
then computes outer products for diagonal correction terms.
|
|
86
86
|
|
|
87
87
|
Distributed, each module belongs to the rank that owned its covariances,
|
|
88
|
-
which
|
|
89
|
-
rank's positions for it; teardown
|
|
88
|
+
which takes its full eigenvectors from ``eigenvectors`` or joins them from
|
|
89
|
+
every rank's shard, and receives every rank's positions for it; teardown
|
|
90
|
+
saves the usual row shards.
|
|
90
91
|
"""
|
|
91
92
|
|
|
92
93
|
path: str
|
|
@@ -104,6 +105,11 @@ class LambdaCollector(HookCollectorBase):
|
|
|
104
105
|
num_documents: int = 1
|
|
105
106
|
"""Divides the summed corrections into a mean over documents."""
|
|
106
107
|
|
|
108
|
+
eigenvectors: tuple[dict[str, Tensor], dict[str, Tensor]] | None = None
|
|
109
|
+
"""The activation and gradient eigenvectors of the modules this rank owns,
|
|
110
|
+
in full, as returned by :func:`eigendecompose_owned`. Read from the shard
|
|
111
|
+
files when not given."""
|
|
112
|
+
|
|
107
113
|
def setup(self) -> None:
|
|
108
114
|
"""Load eigenvectors and initialize storage."""
|
|
109
115
|
self.shard_computer = ShardedMul()
|
|
@@ -121,6 +127,13 @@ class LambdaCollector(HookCollectorBase):
|
|
|
121
127
|
self._setup_owned(eigen_src)
|
|
122
128
|
return
|
|
123
129
|
|
|
130
|
+
if self.eigenvectors is not None:
|
|
131
|
+
self.eigen_a, self.eigen_g = (
|
|
132
|
+
{k: v.to(self.device, self.dtype) for k, v in vectors.items()}
|
|
133
|
+
for vectors in self.eigenvectors
|
|
134
|
+
)
|
|
135
|
+
return
|
|
136
|
+
|
|
124
137
|
# Load precomputed eigenvectors
|
|
125
138
|
self.eigen_a = load_file(
|
|
126
139
|
os.path.join(
|
|
@@ -141,7 +154,7 @@ class LambdaCollector(HookCollectorBase):
|
|
|
141
154
|
|
|
142
155
|
def _setup_owned(self, eigen_src: str) -> None:
|
|
143
156
|
"""Load the full eigenvectors of the modules this rank owns, joining the
|
|
144
|
-
row shards every rank wrote."""
|
|
157
|
+
row shards every rank wrote unless they were passed in."""
|
|
145
158
|
self.owners = assign_module_owners(self.target_info, self.world_size)
|
|
146
159
|
owned = [name for name, owner in self.owners.items() if owner == self.rank]
|
|
147
160
|
|
|
@@ -160,8 +173,14 @@ class LambdaCollector(HookCollectorBase):
|
|
|
160
173
|
for name, rows in full.items()
|
|
161
174
|
}
|
|
162
175
|
|
|
163
|
-
self.
|
|
164
|
-
|
|
176
|
+
if self.eigenvectors is not None:
|
|
177
|
+
self.eigen_a, self.eigen_g = (
|
|
178
|
+
{name: vectors[name].to(self.device, self.dtype) for name in owned}
|
|
179
|
+
for vectors in self.eigenvectors
|
|
180
|
+
)
|
|
181
|
+
else:
|
|
182
|
+
self.eigen_a = load("eigen_activation_sharded")
|
|
183
|
+
self.eigen_g = load("eigen_gradient_sharded")
|
|
165
184
|
self.correction_shapes = {
|
|
166
185
|
name: (out_dim, in_dim + collect_bias)
|
|
167
186
|
for name, (_, (out_dim, in_dim), collect_bias) in self.target_info.items()
|
|
@@ -322,6 +341,86 @@ def _compute_full_matrix(
|
|
|
322
341
|
return full_matrix
|
|
323
342
|
|
|
324
343
|
|
|
344
|
+
def _eigh(
|
|
345
|
+
key: str, matrix: Tensor, total_processed: Tensor, dtype: torch.dtype
|
|
346
|
+
) -> tuple[Tensor, Tensor]:
|
|
347
|
+
"""Eigenvalues and eigenvectors of ``matrix / total_processed``, computed in
|
|
348
|
+
fp64 on ``total_processed``'s device and returned there in ``dtype``."""
|
|
349
|
+
matrix_normalized = matrix.to(total_processed.device, torch.float64)
|
|
350
|
+
# Free the caller's copy early when it passed the only reference.
|
|
351
|
+
del matrix
|
|
352
|
+
matrix_normalized = matrix_normalized + matrix_normalized.T
|
|
353
|
+
matrix_normalized.div_(2 * total_processed)
|
|
354
|
+
|
|
355
|
+
if not torch.isfinite(matrix_normalized).all():
|
|
356
|
+
raise ValueError(
|
|
357
|
+
f"Covariance matrix for {key} contains NaNs or Infs. "
|
|
358
|
+
"Consider using fp32."
|
|
359
|
+
)
|
|
360
|
+
|
|
361
|
+
try:
|
|
362
|
+
eigenvalues, eigenvectors = torch.linalg.eigh(matrix_normalized)
|
|
363
|
+
except Exception as e:
|
|
364
|
+
raise RuntimeError(f"Eigendecomposition failed for {key}") from e
|
|
365
|
+
|
|
366
|
+
del matrix_normalized
|
|
367
|
+
return eigenvalues.to(dtype), eigenvectors.to(dtype).contiguous()
|
|
368
|
+
|
|
369
|
+
|
|
370
|
+
def eigendecompose_owned(
|
|
371
|
+
covariances: dict[str, Tensor],
|
|
372
|
+
shapes: dict[str, tuple[int, ...]],
|
|
373
|
+
owners: dict[str, int] | None,
|
|
374
|
+
total_processed: int | Tensor,
|
|
375
|
+
output_path: str,
|
|
376
|
+
dtype: torch.dtype,
|
|
377
|
+
) -> tuple[dict[str, Tensor], dict[str, Tensor]]:
|
|
378
|
+
"""Eigendecompose the covariances this rank owns, in full, from memory.
|
|
379
|
+
|
|
380
|
+
``owners`` maps every module in ``shapes`` to its rank, or is ``None`` in a
|
|
381
|
+
single process. Saves this rank's row shard of every module's eigenvectors
|
|
382
|
+
to ``output_path`` in ``dtype``, the layout :func:`compute_eigendecomposition`
|
|
383
|
+
writes.
|
|
384
|
+
|
|
385
|
+
Returns the eigenvectors of the modules this rank owns, in full, on this
|
|
386
|
+
rank's device, and this rank's row shard of every module's eigenvalues on
|
|
387
|
+
the CPU. Empties ``covariances`` as it goes, so each covariance is freed once
|
|
388
|
+
its eigenvectors exist.
|
|
389
|
+
"""
|
|
390
|
+
rank = dist.get_rank() if owners is not None else 0
|
|
391
|
+
device = get_device(rank)
|
|
392
|
+
total_processed = torch.as_tensor(total_processed, device=device)
|
|
393
|
+
|
|
394
|
+
eigenvectors: dict[str, Tensor] = {}
|
|
395
|
+
eigenvalues: dict[str, Tensor] = {}
|
|
396
|
+
for key in tqdm(
|
|
397
|
+
list(covariances),
|
|
398
|
+
desc=f"Rank {rank}: Computing eigenvectors",
|
|
399
|
+
position=rank,
|
|
400
|
+
leave=False,
|
|
401
|
+
):
|
|
402
|
+
values, vectors = _eigh(key, covariances.pop(key), total_processed, dtype)
|
|
403
|
+
eigenvectors[key] = vectors
|
|
404
|
+
eigenvalues[key] = values
|
|
405
|
+
|
|
406
|
+
shards = eigenvectors
|
|
407
|
+
value_shards = {key: values.cpu() for key, values in eigenvalues.items()}
|
|
408
|
+
if owners is not None:
|
|
409
|
+
shards = owned_to_row_shards(eigenvectors, shapes, owners, dtype, device)
|
|
410
|
+
value_shards = owned_to_row_shards(
|
|
411
|
+
eigenvalues,
|
|
412
|
+
{key: shape[:1] for key, shape in shapes.items()},
|
|
413
|
+
owners,
|
|
414
|
+
dtype,
|
|
415
|
+
device,
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
os.makedirs(output_path, exist_ok=True)
|
|
419
|
+
save_file(shards, os.path.join(output_path, f"shard_{rank}.safetensors"))
|
|
420
|
+
get_logger().info(f"Saved eigenvectors to {output_path}")
|
|
421
|
+
return eigenvectors, value_shards
|
|
422
|
+
|
|
423
|
+
|
|
325
424
|
def compute_eigendecomposition(
|
|
326
425
|
covariance_path: str,
|
|
327
426
|
total_processed: int | Tensor,
|
|
@@ -383,30 +482,10 @@ def compute_eigendecomposition(
|
|
|
383
482
|
world_size=world_size,
|
|
384
483
|
)
|
|
385
484
|
|
|
386
|
-
|
|
387
|
-
|
|
388
|
-
|
|
389
|
-
|
|
390
|
-
if not torch.isfinite(matrix_normalized).all():
|
|
391
|
-
raise ValueError(
|
|
392
|
-
f"Covariance matrix for {key} contains NaNs or Infs. "
|
|
393
|
-
"Consider using fp32."
|
|
394
|
-
)
|
|
395
|
-
|
|
396
|
-
try:
|
|
397
|
-
eigenvalues, eigenvectors = torch.linalg.eigh(matrix_normalized)
|
|
398
|
-
except Exception as e:
|
|
399
|
-
raise RuntimeError(f"Eigendecomposition failed for {key}") from e
|
|
400
|
-
|
|
401
|
-
# TODO: Maybe possible to avoid CPU transfer here?
|
|
402
|
-
eigenvectors = eigenvectors.to(original_dtype).to(device="cpu").contiguous()
|
|
403
|
-
covariance_eigenvectors[key] = eigenvectors
|
|
404
|
-
covariance_eigenvalues[key] = (
|
|
405
|
-
eigenvalues.to(original_dtype).to(device="cpu").contiguous()
|
|
406
|
-
)
|
|
407
|
-
covariance_eigenvalues[key] = (
|
|
408
|
-
eigenvalues.to(original_dtype).to(device="cpu").contiguous()
|
|
409
|
-
)
|
|
485
|
+
eigenvalues, eigenvectors = _eigh(key, matrix, total_processed, original_dtype)
|
|
486
|
+
del matrix
|
|
487
|
+
covariance_eigenvectors[key] = eigenvectors.cpu()
|
|
488
|
+
covariance_eigenvalues[key] = eigenvalues.cpu()
|
|
410
489
|
|
|
411
490
|
covariance_eigenvectors = _gather_and_shard_along_dim_0(
|
|
412
491
|
input_dict=covariance_eigenvectors,
|
|
@@ -9,6 +9,7 @@ import torch.distributed as dist
|
|
|
9
9
|
from datasets import Dataset
|
|
10
10
|
from safetensors import safe_open
|
|
11
11
|
from safetensors.torch import save_file
|
|
12
|
+
from torch import Tensor
|
|
12
13
|
from transformers import PreTrainedModel
|
|
13
14
|
|
|
14
15
|
from bergson.collector.collector import (
|
|
@@ -27,6 +28,7 @@ from bergson.hessians.autocorrelation import (
|
|
|
27
28
|
from bergson.hessians.eigenvectors import (
|
|
28
29
|
LambdaCollector,
|
|
29
30
|
compute_eigendecomposition,
|
|
31
|
+
eigendecompose_owned,
|
|
30
32
|
save_uncorrected_eigenvalues,
|
|
31
33
|
)
|
|
32
34
|
from bergson.hessians.kfac import CovarianceCollector
|
|
@@ -315,7 +317,11 @@ def fit_factored_hessians(
|
|
|
315
317
|
"path": path,
|
|
316
318
|
}
|
|
317
319
|
|
|
318
|
-
collect_hessians(**kwargs)
|
|
320
|
+
collector = collect_hessians(**kwargs)
|
|
321
|
+
if not do_eigendecomposition and isinstance(collector, CovarianceCollector):
|
|
322
|
+
# Only the eigendecomposition reads the covariances kept in memory.
|
|
323
|
+
collector.A_cov_dict.clear()
|
|
324
|
+
collector.S_cov_dict.clear()
|
|
319
325
|
_release_device_memory()
|
|
320
326
|
|
|
321
327
|
dist.barrier() if dist.is_initialized() else None
|
|
@@ -332,14 +338,35 @@ def fit_factored_hessians(
|
|
|
332
338
|
weights_only=False,
|
|
333
339
|
)
|
|
334
340
|
|
|
335
|
-
|
|
336
|
-
|
|
337
|
-
|
|
338
|
-
|
|
339
|
-
|
|
340
|
-
|
|
341
|
-
|
|
342
|
-
|
|
341
|
+
eigenvectors = None
|
|
342
|
+
if isinstance(collector, CovarianceCollector):
|
|
343
|
+
# Each rank still holds the covariances it owns, so skip the files.
|
|
344
|
+
eigenvectors_a, eigenvalues_a = eigendecompose_owned(
|
|
345
|
+
collector.A_cov_dict,
|
|
346
|
+
collector.A_shapes,
|
|
347
|
+
collector.owners,
|
|
348
|
+
total_processed,
|
|
349
|
+
os.path.join(path, "eigen_activation_sharded"),
|
|
350
|
+
collector.dtype,
|
|
351
|
+
)
|
|
352
|
+
eigenvectors_g, eigenvalues_g = eigendecompose_owned(
|
|
353
|
+
collector.S_cov_dict,
|
|
354
|
+
collector.S_shapes,
|
|
355
|
+
collector.owners,
|
|
356
|
+
total_processed,
|
|
357
|
+
os.path.join(path, "eigen_gradient_sharded"),
|
|
358
|
+
collector.dtype,
|
|
359
|
+
)
|
|
360
|
+
eigenvectors = (eigenvectors_a, eigenvectors_g)
|
|
361
|
+
else:
|
|
362
|
+
eigenvalues_a = compute_eigendecomposition(
|
|
363
|
+
os.path.join(path, "activation_sharded"),
|
|
364
|
+
total_processed=total_processed,
|
|
365
|
+
)
|
|
366
|
+
eigenvalues_g = compute_eigendecomposition(
|
|
367
|
+
os.path.join(path, "gradient_sharded"),
|
|
368
|
+
total_processed=total_processed,
|
|
369
|
+
)
|
|
343
370
|
|
|
344
371
|
dist.barrier() if dist.is_initialized() else None
|
|
345
372
|
|
|
@@ -354,7 +381,12 @@ def fit_factored_hessians(
|
|
|
354
381
|
)
|
|
355
382
|
|
|
356
383
|
if hessian_cfg.ev_correction:
|
|
357
|
-
collect_hessians(
|
|
384
|
+
collect_hessians(
|
|
385
|
+
**kwargs,
|
|
386
|
+
ev_correction=True,
|
|
387
|
+
num_documents=len(data),
|
|
388
|
+
eigenvectors=eigenvectors,
|
|
389
|
+
)
|
|
358
390
|
_release_device_memory()
|
|
359
391
|
|
|
360
392
|
|
|
@@ -380,11 +412,13 @@ def collect_hessians(
|
|
|
380
412
|
output_subdir: str = "eigenvalue_correction_sharded",
|
|
381
413
|
path: str | None = None,
|
|
382
414
|
num_documents: int = 1,
|
|
383
|
-
|
|
415
|
+
eigenvectors: tuple[dict[str, Tensor], dict[str, Tensor]] | None = None,
|
|
416
|
+
) -> HookCollectorBase:
|
|
384
417
|
"""
|
|
385
418
|
Compute Hessian approximations using the hooks specified in the collector.
|
|
386
419
|
If ev_correction is True, uses LambdaCollector to compute eigenvalue corrections.
|
|
387
|
-
``path`` overrides where the collector writes.
|
|
420
|
+
``path`` overrides where the collector writes. ``eigenvectors`` are passed to
|
|
421
|
+
the LambdaCollector. Returns the finished collector.
|
|
388
422
|
"""
|
|
389
423
|
|
|
390
424
|
hessian_dtype = convert_precision_to_torch(hessian_cfg.hessian_dtype)
|
|
@@ -405,6 +439,7 @@ def collect_hessians(
|
|
|
405
439
|
eigen_path=eigen_path,
|
|
406
440
|
output_subdir=output_subdir,
|
|
407
441
|
num_documents=num_documents,
|
|
442
|
+
eigenvectors=eigenvectors,
|
|
408
443
|
)
|
|
409
444
|
desc += " (eigenvalue correction)"
|
|
410
445
|
else:
|
|
@@ -421,3 +456,4 @@ def collect_hessians(
|
|
|
421
456
|
computer.forward_backward = fwd_bwd_hessian_factory(index_cfg, hessian_cfg)
|
|
422
457
|
|
|
423
458
|
computer.run_with_collector_hooks(desc=desc)
|
|
459
|
+
return collector
|
|
@@ -31,6 +31,8 @@ class CovarianceCollector(HookCollectorBase):
|
|
|
31
31
|
|
|
32
32
|
Distributed, each module's covariances belong to one rank, which receives
|
|
33
33
|
every rank's positions for that module; teardown saves the usual row shards.
|
|
34
|
+
Teardown leaves the covariances this rank owns, in full, on its device in
|
|
35
|
+
``A_cov_dict`` and ``S_cov_dict`` for the eigendecomposition.
|
|
34
36
|
"""
|
|
35
37
|
|
|
36
38
|
dtype: torch.dtype
|
|
@@ -152,4 +154,4 @@ class CovarianceCollector(HookCollectorBase):
|
|
|
152
154
|
else:
|
|
153
155
|
shards = covariances
|
|
154
156
|
save_file(shards, os.path.join(path, f"shard_{self.rank}.safetensors"))
|
|
155
|
-
|
|
157
|
+
del shards
|
|
@@ -74,19 +74,31 @@ def owned_to_row_shards(
|
|
|
74
74
|
device: str | torch.device,
|
|
75
75
|
) -> dict[str, Tensor]:
|
|
76
76
|
"""This rank's row shard of every matrix, on the CPU, from ``owned``, the
|
|
77
|
-
matrices this rank owns in full. One matrix at a time is sent from its owner
|
|
77
|
+
matrices this rank owns in full. One matrix at a time is sent from its owner,
|
|
78
|
+
which sends each rank only its rows."""
|
|
78
79
|
rank, world_size = dist.get_rank(), dist.get_world_size()
|
|
79
80
|
shards = {}
|
|
80
81
|
for name, shape in shapes.items():
|
|
81
82
|
owner = owners[name]
|
|
83
|
+
# Rank 0's shard is the largest; the others are padded to its size so
|
|
84
|
+
# every rank receives the same shape.
|
|
85
|
+
_, rows = shard_bounds(shape[0], 0, world_size)
|
|
86
|
+
received = torch.empty((rows, *shape[1:]), device=device, dtype=dtype)
|
|
87
|
+
blocks = None
|
|
82
88
|
if rank == owner:
|
|
83
|
-
full = owned[name].contiguous()
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
89
|
+
full = owned[name].to(device, dtype).contiguous()
|
|
90
|
+
blocks = []
|
|
91
|
+
for r in range(world_size):
|
|
92
|
+
block = full[slice(*shard_bounds(shape[0], r, world_size))]
|
|
93
|
+
if len(block) < rows:
|
|
94
|
+
block = torch.cat(
|
|
95
|
+
[block, block.new_zeros(rows - len(block), *shape[1:])]
|
|
96
|
+
)
|
|
97
|
+
blocks.append(block)
|
|
98
|
+
dist.scatter(received, blocks, src=owner)
|
|
87
99
|
start, end = shard_bounds(shape[0], rank, world_size)
|
|
88
|
-
shards[name] =
|
|
89
|
-
del
|
|
100
|
+
shards[name] = received[: end - start].cpu()
|
|
101
|
+
del received, blocks
|
|
90
102
|
return shards
|
|
91
103
|
|
|
92
104
|
|