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.
Files changed (177) hide show
  1. {bergson-2.2.0 → bergson-2.2.2}/PKG-INFO +28 -34
  2. {bergson-2.2.0 → bergson-2.2.2}/README.md +27 -33
  3. {bergson-2.2.0 → bergson-2.2.2}/bergson/__init__.py +1 -1
  4. {bergson-2.2.0 → bergson-2.2.2}/bergson/config/config.py +10 -0
  5. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/eigenvectors.py +108 -29
  6. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/hessian_approximations.py +48 -12
  7. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/kfac.py +3 -1
  8. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/sharded_computation.py +19 -7
  9. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/cli.py +59 -49
  10. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/data_stream.py +72 -35
  11. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/fsdp.py +11 -5
  12. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/metasmoothness.py +21 -8
  13. bergson-2.2.2/bergson/magic/shard_load.py +119 -0
  14. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/trainer.py +16 -4
  15. bergson-2.2.2/bergson/score/output_influence.py +154 -0
  16. {bergson-2.2.0 → bergson-2.2.2}/bergson/score/score.py +117 -2
  17. {bergson-2.2.0 → bergson-2.2.2}/bergson/score/scorer.py +4 -0
  18. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/worker_utils.py +21 -8
  19. {bergson-2.2.0 → bergson-2.2.2}/bergson/validate.py +19 -31
  20. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/PKG-INFO +28 -34
  21. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/SOURCES.txt +5 -0
  22. {bergson-2.2.0 → bergson-2.2.2}/pyproject.toml +1 -1
  23. {bergson-2.2.0 → bergson-2.2.2}/tests/test_ddp.py +2 -2
  24. {bergson-2.2.0 → bergson-2.2.2}/tests/test_magic.py +55 -25
  25. {bergson-2.2.0 → bergson-2.2.2}/tests/test_metasmoothness.py +3 -2
  26. bergson-2.2.2/tests/test_output_influence.py +191 -0
  27. bergson-2.2.2/tests/test_padding.py +50 -0
  28. {bergson-2.2.0 → bergson-2.2.2}/tests/test_per_query_magic.py +16 -18
  29. {bergson-2.2.0 → bergson-2.2.2}/tests/test_per_token_lds.py +36 -0
  30. {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_loss_padding.py +2 -2
  31. bergson-2.2.2/tests/test_shard_load.py +46 -0
  32. {bergson-2.2.0 → bergson-2.2.2}/LICENSE +0 -0
  33. {bergson-2.2.0 → bergson-2.2.2}/bergson/__main__.py +0 -0
  34. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/adam_preconditioner.py +0 -0
  35. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/approx_unrolling_math.py +0 -0
  36. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/pipeline.py +0 -0
  37. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/precompute_checkpoints.py +0 -0
  38. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/segment_aggregation.py +0 -0
  39. {bergson-2.2.0 → bergson-2.2.2}/bergson/approx_unrolling/train_cfg_io.py +0 -0
  40. {bergson-2.2.0 → bergson-2.2.2}/bergson/build.py +0 -0
  41. {bergson-2.2.0 → bergson-2.2.2}/bergson/builder.py +0 -0
  42. {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/__init__.py +0 -0
  43. {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/commands.py +0 -0
  44. {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/trackstar.py +0 -0
  45. {bergson-2.2.0 → bergson-2.2.2}/bergson/cli/trak.py +0 -0
  46. {bergson-2.2.0 → bergson-2.2.2}/bergson/collection.py +0 -0
  47. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/__init__.py +0 -0
  48. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/collector.py +0 -0
  49. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/dist_autocorrelation_gradient_collector.py +0 -0
  50. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/gradient_collectors.py +0 -0
  51. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/in_memory_collector.py +0 -0
  52. {bergson-2.2.0 → bergson-2.2.2}/bergson/collector/projection_matrix.py +0 -0
  53. {bergson-2.2.0 → bergson-2.2.2}/bergson/config/__init__.py +0 -0
  54. {bergson-2.2.0 → bergson-2.2.2}/bergson/config/config_io.py +0 -0
  55. {bergson-2.2.0 → bergson-2.2.2}/bergson/config/validation.py +0 -0
  56. {bergson-2.2.0 → bergson-2.2.2}/bergson/data.py +0 -0
  57. {bergson-2.2.0 → bergson-2.2.2}/bergson/diagnose.py +0 -0
  58. {bergson-2.2.0 → bergson-2.2.2}/bergson/distributed.py +0 -0
  59. {bergson-2.2.0 → bergson-2.2.2}/bergson/format.py +0 -0
  60. {bergson-2.2.0 → bergson-2.2.2}/bergson/gradients.py +0 -0
  61. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/apply_hessian.py +0 -0
  62. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/astra.py +0 -0
  63. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/autocorrelation.py +0 -0
  64. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/inversion.py +0 -0
  65. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/pipeline.py +0 -0
  66. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/preconditioner.py +0 -0
  67. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/shampoo.py +0 -0
  68. {bergson-2.2.0 → bergson-2.2.2}/bergson/hessians/tkfac.py +0 -0
  69. {bergson-2.2.0 → bergson-2.2.2}/bergson/huggingface.py +0 -0
  70. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/__init__.py +0 -0
  71. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/config.py +0 -0
  72. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/dtensor_patch.py +0 -0
  73. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/grad_accum.py +0 -0
  74. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/optim.py +0 -0
  75. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/rtl_tqdm.py +0 -0
  76. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/score_plot.py +0 -0
  77. {bergson-2.2.0 → bergson-2.2.2}/bergson/magic/swap.py +0 -0
  78. {bergson-2.2.0 → bergson-2.2.2}/bergson/moe.py +0 -0
  79. {bergson-2.2.0 → bergson-2.2.2}/bergson/process_autocorrelation.py +0 -0
  80. {bergson-2.2.0 → bergson-2.2.2}/bergson/process_grads.py +0 -0
  81. {bergson-2.2.0 → bergson-2.2.2}/bergson/query/__init__.py +0 -0
  82. {bergson-2.2.0 → bergson-2.2.2}/bergson/query/attributor.py +0 -0
  83. {bergson-2.2.0 → bergson-2.2.2}/bergson/query/faiss_index.py +0 -0
  84. {bergson-2.2.0 → bergson-2.2.2}/bergson/query/query_index.py +0 -0
  85. {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/__init__.py +0 -0
  86. {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/facts.py +0 -0
  87. {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/generate.py +0 -0
  88. {bergson-2.2.0 → bergson-2.2.2}/bergson/recall/recall.py +0 -0
  89. {bergson-2.2.0 → bergson-2.2.2}/bergson/score/__init__.py +0 -0
  90. {bergson-2.2.0 → bergson-2.2.2}/bergson/score/candidates.py +0 -0
  91. {bergson-2.2.0 → bergson-2.2.2}/bergson/score/score_writer.py +0 -0
  92. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/__init__.py +0 -0
  93. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/batch_size.py +0 -0
  94. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/csv_writer.py +0 -0
  95. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/gradcheck.py +0 -0
  96. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/load_from_optimizer.py +0 -0
  97. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/logger.py +0 -0
  98. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/logging.py +0 -0
  99. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/math.py +0 -0
  100. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/peft.py +0 -0
  101. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/step_state.py +0 -0
  102. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/trainer_export.py +0 -0
  103. {bergson-2.2.0 → bergson-2.2.2}/bergson/utils/utils.py +0 -0
  104. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/dependency_links.txt +0 -0
  105. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/entry_points.txt +0 -0
  106. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/requires.txt +0 -0
  107. {bergson-2.2.0 → bergson-2.2.2}/bergson.egg-info/top_level.txt +0 -0
  108. {bergson-2.2.0 → bergson-2.2.2}/setup.cfg +0 -0
  109. {bergson-2.2.0 → bergson-2.2.2}/tests/test_adam_state_loading.py +0 -0
  110. {bergson-2.2.0 → bergson-2.2.2}/tests/test_advantages.py +0 -0
  111. {bergson-2.2.0 → bergson-2.2.2}/tests/test_approx_unrolling_checkpoint_parse.py +0 -0
  112. {bergson-2.2.0 → bergson-2.2.2}/tests/test_attention_head_normalizers.py +0 -0
  113. {bergson-2.2.0 → bergson-2.2.2}/tests/test_attribute_tokens.py +0 -0
  114. {bergson-2.2.0 → bergson-2.2.2}/tests/test_attributor.py +0 -0
  115. {bergson-2.2.0 → bergson-2.2.2}/tests/test_bank_loss_cache.py +0 -0
  116. {bergson-2.2.0 → bergson-2.2.2}/tests/test_banked_model_loading.py +0 -0
  117. {bergson-2.2.0 → bergson-2.2.2}/tests/test_batch_size_invariance.py +0 -0
  118. {bergson-2.2.0 → bergson-2.2.2}/tests/test_build.py +0 -0
  119. {bergson-2.2.0 → bergson-2.2.2}/tests/test_builder.py +0 -0
  120. {bergson-2.2.0 → bergson-2.2.2}/tests/test_candidates.py +0 -0
  121. {bergson-2.2.0 → bergson-2.2.2}/tests/test_ckpt_avg_query_grads.py +0 -0
  122. {bergson-2.2.0 → bergson-2.2.2}/tests/test_cli.py +0 -0
  123. {bergson-2.2.0 → bergson-2.2.2}/tests/test_compute_lambda.py +0 -0
  124. {bergson-2.2.0 → bergson-2.2.2}/tests/test_config_runner.py +0 -0
  125. {bergson-2.2.0 → bergson-2.2.2}/tests/test_contrast.py +0 -0
  126. {bergson-2.2.0 → bergson-2.2.2}/tests/test_data.py +0 -0
  127. {bergson-2.2.0 → bergson-2.2.2}/tests/test_diagnose.py +0 -0
  128. {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_batch_budget.py +0 -0
  129. {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_cap.py +0 -0
  130. {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_failure.py +0 -0
  131. {bergson-2.2.0 → bergson-2.2.2}/tests/test_distributed_magic.py +0 -0
  132. {bergson-2.2.0 → bergson-2.2.2}/tests/test_epoch_shuffle.py +0 -0
  133. {bergson-2.2.0 → bergson-2.2.2}/tests/test_faiss_ann_cpu_load.py +0 -0
  134. {bergson-2.2.0 → bergson-2.2.2}/tests/test_force_math_sdp.py +0 -0
  135. {bergson-2.2.0 → bergson-2.2.2}/tests/test_format.py +0 -0
  136. {bergson-2.2.0 → bergson-2.2.2}/tests/test_global_kfac.py +0 -0
  137. {bergson-2.2.0 → bergson-2.2.2}/tests/test_global_projection.py +0 -0
  138. {bergson-2.2.0 → bergson-2.2.2}/tests/test_grad_clipping.py +0 -0
  139. {bergson-2.2.0 → bergson-2.2.2}/tests/test_gradcheck.py +0 -0
  140. {bergson-2.2.0 → bergson-2.2.2}/tests/test_gradients.py +0 -0
  141. {bergson-2.2.0 → bergson-2.2.2}/tests/test_hessian.py +0 -0
  142. {bergson-2.2.0 → bergson-2.2.2}/tests/test_launch_devices.py +0 -0
  143. {bergson-2.2.0 → bergson-2.2.2}/tests/test_log_fn.py +0 -0
  144. {bergson-2.2.0 → bergson-2.2.2}/tests/test_logit_scale.py +0 -0
  145. {bergson-2.2.0 → bergson-2.2.2}/tests/test_moe.py +0 -0
  146. {bergson-2.2.0 → bergson-2.2.2}/tests/test_multi_query_validate.py +0 -0
  147. {bergson-2.2.0 → bergson-2.2.2}/tests/test_multinode.py +0 -0
  148. {bergson-2.2.0 → bergson-2.2.2}/tests/test_muon.py +0 -0
  149. {bergson-2.2.0 → bergson-2.2.2}/tests/test_normalizer.py +0 -0
  150. {bergson-2.2.0 → bergson-2.2.2}/tests/test_outer_product_gradients.py +0 -0
  151. {bergson-2.2.0 → bergson-2.2.2}/tests/test_pretokenized.py +0 -0
  152. {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_inner_products.py +0 -0
  153. {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_matrix.py +0 -0
  154. {bergson-2.2.0 → bergson-2.2.2}/tests/test_projection_settings.py +0 -0
  155. {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_batching_invariance.py +0 -0
  156. {bergson-2.2.0 → bergson-2.2.2}/tests/test_query_set_config.py +0 -0
  157. {bergson-2.2.0 → bergson-2.2.2}/tests/test_recall.py +0 -0
  158. {bergson-2.2.0 → bergson-2.2.2}/tests/test_reduce.py +0 -0
  159. {bergson-2.2.0 → bergson-2.2.2}/tests/test_save_mode_final.py +0 -0
  160. {bergson-2.2.0 → bergson-2.2.2}/tests/test_score.py +0 -0
  161. {bergson-2.2.0 → bergson-2.2.2}/tests/test_score_plot.py +0 -0
  162. {bergson-2.2.0 → bergson-2.2.2}/tests/test_score_writer_dist.py +0 -0
  163. {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_fisher_normalization.py +0 -0
  164. {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_optimizer_variants.py +0 -0
  165. {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_resume_polarity.py +0 -0
  166. {bergson-2.2.0 → bergson-2.2.2}/tests/test_source_trainer_integration.py +0 -0
  167. {bergson-2.2.0 → bergson-2.2.2}/tests/test_step_matrix.py +0 -0
  168. {bergson-2.2.0 → bergson-2.2.2}/tests/test_step_state.py +0 -0
  169. {bergson-2.2.0 → bergson-2.2.2}/tests/test_tokenize.py +0 -0
  170. {bergson-2.2.0 → bergson-2.2.2}/tests/test_trainer_callback.py +0 -0
  171. {bergson-2.2.0 → bergson-2.2.2}/tests/test_trak.py +0 -0
  172. {bergson-2.2.0 → bergson-2.2.2}/tests/test_truncation.py +0 -0
  173. {bergson-2.2.0 → bergson-2.2.2}/tests/test_validate_filter.py +0 -0
  174. {bergson-2.2.0 → bergson-2.2.2}/tests/test_validate_routing.py +0 -0
  175. {bergson-2.2.0 → bergson-2.2.2}/tests/test_validation_config.py +0 -0
  176. {bergson-2.2.0 → bergson-2.2.2}/tests/test_wandb_logging.py +0 -0
  177. {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.0
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 | Proponent QLD [95% CI] | LDS [95% CI] |
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
- | EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] |
96
- | KFAC | 0.067 [0.056, 0.080] | 0.420 [0.391, 0.446] |
97
- | BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] |
98
- | [Qwen3-Embedding-8B](https://huggingface.co/spaces/mteb/leaderboard) semantic search | 0.049 [0.038, 0.061] | 0.132 [0.093, 0.169] |
99
- | TrackStar (no optimizer correction, projection 64) | 0.045 [0.036, 0.055] | 0.270 [0.240, 0.295] |
100
- | TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] |
101
- | SOURCE (Adam) | 0.024 [0.018, 0.030] | 0.154 [0.126, 0.181] |
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). Untuned training hyperparameters can result in disappointing linear datamodeling scores. Check your run's estimated metasmoothness with `bergson 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 can be sensitive to both Hessian approximation inversion hyperparameters (in tiny models) and metasmoothness. Untuned hyperparameters can result in a disappointing linear datamodeling score.**
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
- # Examples
135
+ # Getting Started
136
136
 
137
- There are many example YAMLs in the `examples` directory. use `bergson <yaml_path>` to try them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
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
- Or check out a notebook for programmatic usage:
144
-
145
- | | Notebook | Description |
146
- |---|----------|-------------|
147
- | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
- | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
- Collect TrackStar attribution scores for an I.I.D sample query:
150
+ Or check out a notebook for programmatic usage:
158
151
 
159
- ```bash
160
- bergson trackstar runs/trackstar --model EleutherAI/pythia-14m --query.dataset NeelNanda/pile-10k --data.dataset NeelNanda/pile-10k --data.truncation --token_batch_size 4096 --query.truncation --query.split "train[:20]"
161
- ```
152
+ | | Notebook | Description |
153
+ |---|----------|-------------|
154
+ | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
+ | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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 | Proponent QLD [95% CI] | LDS [95% CI] |
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
- | EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] |
35
- | KFAC | 0.067 [0.056, 0.080] | 0.420 [0.391, 0.446] |
36
- | BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] |
37
- | [Qwen3-Embedding-8B](https://huggingface.co/spaces/mteb/leaderboard) semantic search | 0.049 [0.038, 0.061] | 0.132 [0.093, 0.169] |
38
- | TrackStar (no optimizer correction, projection 64) | 0.045 [0.036, 0.055] | 0.270 [0.240, 0.295] |
39
- | TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] |
40
- | SOURCE (Adam) | 0.024 [0.018, 0.030] | 0.154 [0.126, 0.181] |
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). Untuned training hyperparameters can result in disappointing linear datamodeling scores. Check your run's estimated metasmoothness with `bergson 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 can be sensitive to both Hessian approximation inversion hyperparameters (in tiny models) and metasmoothness. Untuned hyperparameters can result in a disappointing linear datamodeling score.**
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
- # Examples
74
+ # Getting Started
75
75
 
76
- There are many example YAMLs in the `examples` directory. use `bergson <yaml_path>` to try them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
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
- Or check out a notebook for programmatic usage:
83
-
84
- | | Notebook | Description |
85
- |---|----------|-------------|
86
- | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
- | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
- Collect TrackStar attribution scores for an I.I.D sample query:
89
+ Or check out a notebook for programmatic usage:
97
90
 
98
- ```bash
99
- bergson trackstar runs/trackstar --model EleutherAI/pythia-14m --query.dataset NeelNanda/pile-10k --data.dataset NeelNanda/pile-10k --data.truncation --token_batch_size 4096 --query.truncation --query.split "train[:20]"
100
- ```
91
+ | | Notebook | Description |
92
+ |---|----------|-------------|
93
+ | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
+ | [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](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
 
@@ -1,4 +1,4 @@
1
- __version__ = "2.2.0"
1
+ __version__ = "2.2.2"
2
2
 
3
3
  import logging
4
4
 
@@ -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 loads its full eigenvectors from every rank's shard and receives every
89
- rank's positions for it; teardown saves the usual row shards.
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.eigen_a = load("eigen_activation_sharded")
164
- self.eigen_g = load("eigen_gradient_sharded")
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
- # original_dtype = matrix.dtype
387
- matrix_normalized = matrix.to(torch.float64) / total_processed
388
- matrix_normalized = (matrix_normalized + matrix_normalized.T).div(2)
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
- eigenvalues_a = compute_eigendecomposition(
336
- os.path.join(path, "activation_sharded"),
337
- total_processed=total_processed,
338
- )
339
- eigenvalues_g = compute_eigendecomposition(
340
- os.path.join(path, "gradient_sharded"),
341
- total_processed=total_processed,
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(**kwargs, ev_correction=True, num_documents=len(data))
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
- covariances.clear()
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
- else:
85
- full = torch.empty(shape, device=device, dtype=dtype)
86
- dist.broadcast(full, src=owner)
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] = full[start:end].cpu()
89
- del full
100
+ shards[name] = received[: end - start].cpu()
101
+ del received, blocks
90
102
  return shards
91
103
 
92
104