bergson 2.2.2__tar.gz → 2.2.3__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 (178) hide show
  1. {bergson-2.2.2 → bergson-2.2.3}/PKG-INFO +2 -2
  2. {bergson-2.2.2 → bergson-2.2.3}/README.md +1 -1
  3. {bergson-2.2.2 → bergson-2.2.3}/bergson/__init__.py +1 -1
  4. {bergson-2.2.2 → bergson-2.2.3}/bergson/process_autocorrelation.py +41 -30
  5. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/worker_utils.py +4 -6
  6. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/PKG-INFO +2 -2
  7. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/SOURCES.txt +1 -0
  8. {bergson-2.2.2 → bergson-2.2.3}/pyproject.toml +4 -1
  9. bergson-2.2.3/tests/test_process_autocorrelation.py +62 -0
  10. {bergson-2.2.2 → bergson-2.2.3}/LICENSE +0 -0
  11. {bergson-2.2.2 → bergson-2.2.3}/bergson/__main__.py +0 -0
  12. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/adam_preconditioner.py +0 -0
  13. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/approx_unrolling_math.py +0 -0
  14. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/pipeline.py +0 -0
  15. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/precompute_checkpoints.py +0 -0
  16. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/segment_aggregation.py +0 -0
  17. {bergson-2.2.2 → bergson-2.2.3}/bergson/approx_unrolling/train_cfg_io.py +0 -0
  18. {bergson-2.2.2 → bergson-2.2.3}/bergson/build.py +0 -0
  19. {bergson-2.2.2 → bergson-2.2.3}/bergson/builder.py +0 -0
  20. {bergson-2.2.2 → bergson-2.2.3}/bergson/cli/__init__.py +0 -0
  21. {bergson-2.2.2 → bergson-2.2.3}/bergson/cli/commands.py +0 -0
  22. {bergson-2.2.2 → bergson-2.2.3}/bergson/cli/trackstar.py +0 -0
  23. {bergson-2.2.2 → bergson-2.2.3}/bergson/cli/trak.py +0 -0
  24. {bergson-2.2.2 → bergson-2.2.3}/bergson/collection.py +0 -0
  25. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/__init__.py +0 -0
  26. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/collector.py +0 -0
  27. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/dist_autocorrelation_gradient_collector.py +0 -0
  28. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/gradient_collectors.py +0 -0
  29. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/in_memory_collector.py +0 -0
  30. {bergson-2.2.2 → bergson-2.2.3}/bergson/collector/projection_matrix.py +0 -0
  31. {bergson-2.2.2 → bergson-2.2.3}/bergson/config/__init__.py +0 -0
  32. {bergson-2.2.2 → bergson-2.2.3}/bergson/config/config.py +0 -0
  33. {bergson-2.2.2 → bergson-2.2.3}/bergson/config/config_io.py +0 -0
  34. {bergson-2.2.2 → bergson-2.2.3}/bergson/config/validation.py +0 -0
  35. {bergson-2.2.2 → bergson-2.2.3}/bergson/data.py +0 -0
  36. {bergson-2.2.2 → bergson-2.2.3}/bergson/diagnose.py +0 -0
  37. {bergson-2.2.2 → bergson-2.2.3}/bergson/distributed.py +0 -0
  38. {bergson-2.2.2 → bergson-2.2.3}/bergson/format.py +0 -0
  39. {bergson-2.2.2 → bergson-2.2.3}/bergson/gradients.py +0 -0
  40. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/apply_hessian.py +0 -0
  41. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/astra.py +0 -0
  42. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/autocorrelation.py +0 -0
  43. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/eigenvectors.py +0 -0
  44. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/hessian_approximations.py +0 -0
  45. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/inversion.py +0 -0
  46. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/kfac.py +0 -0
  47. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/pipeline.py +0 -0
  48. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/preconditioner.py +0 -0
  49. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/shampoo.py +0 -0
  50. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/sharded_computation.py +0 -0
  51. {bergson-2.2.2 → bergson-2.2.3}/bergson/hessians/tkfac.py +0 -0
  52. {bergson-2.2.2 → bergson-2.2.3}/bergson/huggingface.py +0 -0
  53. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/__init__.py +0 -0
  54. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/cli.py +0 -0
  55. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/config.py +0 -0
  56. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/data_stream.py +0 -0
  57. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/dtensor_patch.py +0 -0
  58. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/fsdp.py +0 -0
  59. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/grad_accum.py +0 -0
  60. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/metasmoothness.py +0 -0
  61. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/optim.py +0 -0
  62. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/rtl_tqdm.py +0 -0
  63. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/score_plot.py +0 -0
  64. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/shard_load.py +0 -0
  65. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/swap.py +0 -0
  66. {bergson-2.2.2 → bergson-2.2.3}/bergson/magic/trainer.py +0 -0
  67. {bergson-2.2.2 → bergson-2.2.3}/bergson/moe.py +0 -0
  68. {bergson-2.2.2 → bergson-2.2.3}/bergson/process_grads.py +0 -0
  69. {bergson-2.2.2 → bergson-2.2.3}/bergson/query/__init__.py +0 -0
  70. {bergson-2.2.2 → bergson-2.2.3}/bergson/query/attributor.py +0 -0
  71. {bergson-2.2.2 → bergson-2.2.3}/bergson/query/faiss_index.py +0 -0
  72. {bergson-2.2.2 → bergson-2.2.3}/bergson/query/query_index.py +0 -0
  73. {bergson-2.2.2 → bergson-2.2.3}/bergson/recall/__init__.py +0 -0
  74. {bergson-2.2.2 → bergson-2.2.3}/bergson/recall/facts.py +0 -0
  75. {bergson-2.2.2 → bergson-2.2.3}/bergson/recall/generate.py +0 -0
  76. {bergson-2.2.2 → bergson-2.2.3}/bergson/recall/recall.py +0 -0
  77. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/__init__.py +0 -0
  78. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/candidates.py +0 -0
  79. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/output_influence.py +0 -0
  80. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/score.py +0 -0
  81. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/score_writer.py +0 -0
  82. {bergson-2.2.2 → bergson-2.2.3}/bergson/score/scorer.py +0 -0
  83. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/__init__.py +0 -0
  84. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/batch_size.py +0 -0
  85. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/csv_writer.py +0 -0
  86. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/gradcheck.py +0 -0
  87. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/load_from_optimizer.py +0 -0
  88. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/logger.py +0 -0
  89. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/logging.py +0 -0
  90. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/math.py +0 -0
  91. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/peft.py +0 -0
  92. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/step_state.py +0 -0
  93. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/trainer_export.py +0 -0
  94. {bergson-2.2.2 → bergson-2.2.3}/bergson/utils/utils.py +0 -0
  95. {bergson-2.2.2 → bergson-2.2.3}/bergson/validate.py +0 -0
  96. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/dependency_links.txt +0 -0
  97. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/entry_points.txt +0 -0
  98. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/requires.txt +0 -0
  99. {bergson-2.2.2 → bergson-2.2.3}/bergson.egg-info/top_level.txt +0 -0
  100. {bergson-2.2.2 → bergson-2.2.3}/setup.cfg +0 -0
  101. {bergson-2.2.2 → bergson-2.2.3}/tests/test_adam_state_loading.py +0 -0
  102. {bergson-2.2.2 → bergson-2.2.3}/tests/test_advantages.py +0 -0
  103. {bergson-2.2.2 → bergson-2.2.3}/tests/test_approx_unrolling_checkpoint_parse.py +0 -0
  104. {bergson-2.2.2 → bergson-2.2.3}/tests/test_attention_head_normalizers.py +0 -0
  105. {bergson-2.2.2 → bergson-2.2.3}/tests/test_attribute_tokens.py +0 -0
  106. {bergson-2.2.2 → bergson-2.2.3}/tests/test_attributor.py +0 -0
  107. {bergson-2.2.2 → bergson-2.2.3}/tests/test_bank_loss_cache.py +0 -0
  108. {bergson-2.2.2 → bergson-2.2.3}/tests/test_banked_model_loading.py +0 -0
  109. {bergson-2.2.2 → bergson-2.2.3}/tests/test_batch_size_invariance.py +0 -0
  110. {bergson-2.2.2 → bergson-2.2.3}/tests/test_build.py +0 -0
  111. {bergson-2.2.2 → bergson-2.2.3}/tests/test_builder.py +0 -0
  112. {bergson-2.2.2 → bergson-2.2.3}/tests/test_candidates.py +0 -0
  113. {bergson-2.2.2 → bergson-2.2.3}/tests/test_ckpt_avg_query_grads.py +0 -0
  114. {bergson-2.2.2 → bergson-2.2.3}/tests/test_cli.py +0 -0
  115. {bergson-2.2.2 → bergson-2.2.3}/tests/test_compute_lambda.py +0 -0
  116. {bergson-2.2.2 → bergson-2.2.3}/tests/test_config_runner.py +0 -0
  117. {bergson-2.2.2 → bergson-2.2.3}/tests/test_contrast.py +0 -0
  118. {bergson-2.2.2 → bergson-2.2.3}/tests/test_data.py +0 -0
  119. {bergson-2.2.2 → bergson-2.2.3}/tests/test_ddp.py +0 -0
  120. {bergson-2.2.2 → bergson-2.2.3}/tests/test_diagnose.py +0 -0
  121. {bergson-2.2.2 → bergson-2.2.3}/tests/test_distributed_batch_budget.py +0 -0
  122. {bergson-2.2.2 → bergson-2.2.3}/tests/test_distributed_cap.py +0 -0
  123. {bergson-2.2.2 → bergson-2.2.3}/tests/test_distributed_failure.py +0 -0
  124. {bergson-2.2.2 → bergson-2.2.3}/tests/test_distributed_magic.py +0 -0
  125. {bergson-2.2.2 → bergson-2.2.3}/tests/test_epoch_shuffle.py +0 -0
  126. {bergson-2.2.2 → bergson-2.2.3}/tests/test_faiss_ann_cpu_load.py +0 -0
  127. {bergson-2.2.2 → bergson-2.2.3}/tests/test_force_math_sdp.py +0 -0
  128. {bergson-2.2.2 → bergson-2.2.3}/tests/test_format.py +0 -0
  129. {bergson-2.2.2 → bergson-2.2.3}/tests/test_global_kfac.py +0 -0
  130. {bergson-2.2.2 → bergson-2.2.3}/tests/test_global_projection.py +0 -0
  131. {bergson-2.2.2 → bergson-2.2.3}/tests/test_grad_clipping.py +0 -0
  132. {bergson-2.2.2 → bergson-2.2.3}/tests/test_gradcheck.py +0 -0
  133. {bergson-2.2.2 → bergson-2.2.3}/tests/test_gradients.py +0 -0
  134. {bergson-2.2.2 → bergson-2.2.3}/tests/test_hessian.py +0 -0
  135. {bergson-2.2.2 → bergson-2.2.3}/tests/test_launch_devices.py +0 -0
  136. {bergson-2.2.2 → bergson-2.2.3}/tests/test_log_fn.py +0 -0
  137. {bergson-2.2.2 → bergson-2.2.3}/tests/test_logit_scale.py +0 -0
  138. {bergson-2.2.2 → bergson-2.2.3}/tests/test_magic.py +0 -0
  139. {bergson-2.2.2 → bergson-2.2.3}/tests/test_metasmoothness.py +0 -0
  140. {bergson-2.2.2 → bergson-2.2.3}/tests/test_moe.py +0 -0
  141. {bergson-2.2.2 → bergson-2.2.3}/tests/test_multi_query_validate.py +0 -0
  142. {bergson-2.2.2 → bergson-2.2.3}/tests/test_multinode.py +0 -0
  143. {bergson-2.2.2 → bergson-2.2.3}/tests/test_muon.py +0 -0
  144. {bergson-2.2.2 → bergson-2.2.3}/tests/test_normalizer.py +0 -0
  145. {bergson-2.2.2 → bergson-2.2.3}/tests/test_outer_product_gradients.py +0 -0
  146. {bergson-2.2.2 → bergson-2.2.3}/tests/test_output_influence.py +0 -0
  147. {bergson-2.2.2 → bergson-2.2.3}/tests/test_padding.py +0 -0
  148. {bergson-2.2.2 → bergson-2.2.3}/tests/test_per_query_magic.py +0 -0
  149. {bergson-2.2.2 → bergson-2.2.3}/tests/test_per_token_lds.py +0 -0
  150. {bergson-2.2.2 → bergson-2.2.3}/tests/test_pretokenized.py +0 -0
  151. {bergson-2.2.2 → bergson-2.2.3}/tests/test_projection_inner_products.py +0 -0
  152. {bergson-2.2.2 → bergson-2.2.3}/tests/test_projection_matrix.py +0 -0
  153. {bergson-2.2.2 → bergson-2.2.3}/tests/test_projection_settings.py +0 -0
  154. {bergson-2.2.2 → bergson-2.2.3}/tests/test_query_batching_invariance.py +0 -0
  155. {bergson-2.2.2 → bergson-2.2.3}/tests/test_query_loss_padding.py +0 -0
  156. {bergson-2.2.2 → bergson-2.2.3}/tests/test_query_set_config.py +0 -0
  157. {bergson-2.2.2 → bergson-2.2.3}/tests/test_recall.py +0 -0
  158. {bergson-2.2.2 → bergson-2.2.3}/tests/test_reduce.py +0 -0
  159. {bergson-2.2.2 → bergson-2.2.3}/tests/test_save_mode_final.py +0 -0
  160. {bergson-2.2.2 → bergson-2.2.3}/tests/test_score.py +0 -0
  161. {bergson-2.2.2 → bergson-2.2.3}/tests/test_score_plot.py +0 -0
  162. {bergson-2.2.2 → bergson-2.2.3}/tests/test_score_writer_dist.py +0 -0
  163. {bergson-2.2.2 → bergson-2.2.3}/tests/test_shard_load.py +0 -0
  164. {bergson-2.2.2 → bergson-2.2.3}/tests/test_source_fisher_normalization.py +0 -0
  165. {bergson-2.2.2 → bergson-2.2.3}/tests/test_source_optimizer_variants.py +0 -0
  166. {bergson-2.2.2 → bergson-2.2.3}/tests/test_source_resume_polarity.py +0 -0
  167. {bergson-2.2.2 → bergson-2.2.3}/tests/test_source_trainer_integration.py +0 -0
  168. {bergson-2.2.2 → bergson-2.2.3}/tests/test_step_matrix.py +0 -0
  169. {bergson-2.2.2 → bergson-2.2.3}/tests/test_step_state.py +0 -0
  170. {bergson-2.2.2 → bergson-2.2.3}/tests/test_tokenize.py +0 -0
  171. {bergson-2.2.2 → bergson-2.2.3}/tests/test_trainer_callback.py +0 -0
  172. {bergson-2.2.2 → bergson-2.2.3}/tests/test_trak.py +0 -0
  173. {bergson-2.2.2 → bergson-2.2.3}/tests/test_truncation.py +0 -0
  174. {bergson-2.2.2 → bergson-2.2.3}/tests/test_validate_filter.py +0 -0
  175. {bergson-2.2.2 → bergson-2.2.3}/tests/test_validate_routing.py +0 -0
  176. {bergson-2.2.2 → bergson-2.2.3}/tests/test_validation_config.py +0 -0
  177. {bergson-2.2.2 → bergson-2.2.3}/tests/test_wandb_logging.py +0 -0
  178. {bergson-2.2.2 → bergson-2.2.3}/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.2
3
+ Version: 2.2.3
4
4
  Summary: Tracing the memory of neural nets with data attribution
5
5
  License: MIT License
6
6
  Keywords: interpretability,explainable-ai
@@ -97,7 +97,7 @@ Indicative performance of data attribution methods in the finetuning regime - se
97
97
  | EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
98
98
  | BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
99
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] |
100
+ | TrackStar (no optimizer) | 0.045 [0.036, 0.055] | 0.276 [0.252, 0.299] | 0.492 [0.467, 0.517] |
101
101
  | TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
102
102
  | Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
103
103
  | Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
@@ -36,7 +36,7 @@ Indicative performance of data attribution methods in the finetuning regime - se
36
36
  | EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
37
37
  | BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
38
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] |
39
+ | TrackStar (no optimizer) | 0.045 [0.036, 0.055] | 0.276 [0.252, 0.299] | 0.492 [0.467, 0.517] |
40
40
  | TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
41
41
  | Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
42
42
  | Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
@@ -1,4 +1,4 @@
1
- __version__ = "2.2.2"
1
+ __version__ = "2.2.3"
2
2
 
3
3
  import logging
4
4
 
@@ -4,6 +4,14 @@ import torch.distributed as dist
4
4
  from bergson.gradients import GradientProcessor
5
5
 
6
6
 
7
+ def _eigh(prec: torch.Tensor, device, dtype) -> tuple[torch.Tensor, torch.Tensor]:
8
+ eigvals, eigvecs = torch.linalg.eigh(prec.to(dtype=torch.float64, device=device))
9
+ return (
10
+ eigvals.to(dtype=dtype).contiguous().cpu(),
11
+ eigvecs.to(dtype=dtype).contiguous().cpu(),
12
+ )
13
+
14
+
7
15
  def process_autocorrelation_matrices(
8
16
  processor: GradientProcessor,
9
17
  hessians: dict[str, torch.Tensor],
@@ -20,9 +28,11 @@ def process_autocorrelation_matrices(
20
28
  count when ``attribute_tokens`` produced per-token rows — so the
21
29
  normalized matrix is the second moment over whichever unit the rows
22
30
  represent.
23
- """
24
- hessians_eigen = {}
25
31
 
32
+ A rank may hold a partial Gram over its own rows for every module, or the
33
+ full Gram for the modules it owns. Each module's Gram is summed onto one
34
+ of the ranks that hold it, which eigendecomposes the total.
35
+ """
26
36
  device = next(iter(hessians.values())).device
27
37
  dtype = next(iter(hessians.values())).dtype
28
38
 
@@ -32,57 +42,58 @@ def process_autocorrelation_matrices(
32
42
  for name, prec in hessians.items():
33
43
  hessians[name] = (prec / num_rows).cpu()
34
44
 
35
- if rank == 0:
36
- print("Computing hessian eigen decompositions...")
37
-
38
- for name in hessians.keys():
39
- prec = hessians[name].to(dtype=torch.float64, device=device)
40
- eigvals, eigvecs = torch.linalg.eigh(prec)
41
- hessians_eigen[name] = (
42
- eigvals.to(dtype=dtype).contiguous().cpu(),
43
- eigvecs.to(dtype=dtype).contiguous().cpu(),
44
- )
45
-
46
45
  if not dist.is_initialized():
46
+ if rank == 0:
47
+ print("Computing hessian eigen decompositions...")
47
48
  processor.hessians = hessians
48
- processor.hessians_eigen = hessians_eigen
49
+ processor.hessians_eigen = {
50
+ name: _eigh(prec, device, dtype) for name, prec in hessians.items()
51
+ }
49
52
  return
50
53
 
51
54
  if rank == 0:
52
- print("Gathering hessians...")
55
+ print("Reducing hessians and computing eigen decompositions...")
53
56
 
54
57
  cpu_group: dist.ProcessGroup = dist.new_group(backend="gloo") # type: ignore[assignment]
55
-
56
- for name, grad_size in grad_sizes.items():
58
+ world_size = dist.get_world_size()
59
+ held: list[list[str]] = [[] for _ in range(world_size)]
60
+ dist.all_gather_object(held, sorted(hessians), group=cpu_group)
61
+
62
+ totals, hessians_eigen = {}, {}
63
+ for i, (name, grad_size) in enumerate(grad_sizes.items()):
64
+ holders = [r for r in range(world_size) if name in held[r]] or [0]
65
+ owner = holders[i % len(holders)]
57
66
  if name in hessians:
58
- local_prec = hessians[name]
59
- del hessians[name]
67
+ local_prec = hessians.pop(name)
60
68
  else:
61
69
  local_prec = torch.zeros([grad_size, grad_size], dtype=dtype, device="cpu")
62
70
 
63
- dist.reduce(local_prec, dst=0, op=dist.ReduceOp.SUM, group=cpu_group)
71
+ dist.reduce(local_prec, dst=owner, op=dist.ReduceOp.SUM, group=cpu_group)
64
72
 
65
- if rank == 0:
66
- hessians[name] = local_prec
73
+ if rank == owner:
74
+ totals[name] = local_prec
75
+ hessians_eigen[name] = _eigh(local_prec, device, dtype)
67
76
 
68
77
  if rank == 0:
69
- processor.hessians = hessians
70
-
71
- print("Gathering eigen decompositions...")
78
+ print("Gathering hessians and eigen decompositions...")
72
79
 
73
80
  for name, grad_size in grad_sizes.items():
74
81
  prec_size = torch.Size([grad_size, grad_size])
75
- if name not in hessians_eigen:
82
+ prec = totals.pop(name, None)
83
+ if name in hessians_eigen:
84
+ eigval, eigvec = hessians_eigen[name]
85
+ else:
86
+ prec = torch.zeros(prec_size, dtype=dtype)
76
87
  eigval = torch.zeros(prec_size[0], dtype=dtype)
77
88
  eigvec = torch.zeros(prec_size, dtype=dtype)
78
- else:
79
- eigval, eigvec = hessians_eigen[name]
80
89
 
81
- dist.reduce(eigval, dst=0, op=dist.ReduceOp.SUM, group=cpu_group)
82
- dist.reduce(eigvec, dst=0, op=dist.ReduceOp.SUM, group=cpu_group)
90
+ for t in (prec, eigval, eigvec):
91
+ dist.reduce(t, dst=0, op=dist.ReduceOp.SUM, group=cpu_group)
83
92
 
84
93
  if rank == 0:
94
+ hessians[name] = prec
85
95
  hessians_eigen[name] = (eigval, eigvec)
86
96
 
87
97
  if rank == 0:
98
+ processor.hessians = hessians
88
99
  processor.hessians_eigen = hessians_eigen
@@ -217,13 +217,11 @@ def setup_model_and_peft(
217
217
  model_kwargs.update(simple_parse_kwargs_string(cfg.model_kwargs))
218
218
 
219
219
  if meta_init:
220
- config = AutoConfig.from_pretrained(base_model_path, revision=cfg.revision)
220
+ config = AutoConfig.from_pretrained(
221
+ base_model_path, revision=cfg.revision, **model_kwargs
222
+ )
221
223
  with torch.device("meta"):
222
- model = AutoModelForCausalLM.from_config(
223
- config,
224
- dtype=dtype,
225
- **model_kwargs,
226
- )
224
+ model = AutoModelForCausalLM.from_config(config, dtype=dtype)
227
225
  else:
228
226
  model = AutoModelForCausalLM.from_pretrained(
229
227
  base_model_path,
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bergson
3
- Version: 2.2.2
3
+ Version: 2.2.3
4
4
  Summary: Tracing the memory of neural nets with data attribution
5
5
  License: MIT License
6
6
  Keywords: interpretability,explainable-ai
@@ -97,7 +97,7 @@ Indicative performance of data attribution methods in the finetuning regime - se
97
97
  | EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
98
98
  | BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
99
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] |
100
+ | TrackStar (no optimizer) | 0.045 [0.036, 0.055] | 0.276 [0.252, 0.299] | 0.492 [0.467, 0.517] |
101
101
  | TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
102
102
  | Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
103
103
  | Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
@@ -145,6 +145,7 @@ tests/test_padding.py
145
145
  tests/test_per_query_magic.py
146
146
  tests/test_per_token_lds.py
147
147
  tests/test_pretokenized.py
148
+ tests/test_process_autocorrelation.py
148
149
  tests/test_projection_inner_products.py
149
150
  tests/test_projection_matrix.py
150
151
  tests/test_projection_settings.py
@@ -37,7 +37,7 @@ dependencies = [
37
37
  "bitsandbytes",
38
38
  "huggingface-hub>=1.13.0",
39
39
  ]
40
- version = "2.2.2"
40
+ version = "2.2.3"
41
41
  [project.optional-dependencies]
42
42
  dev = [
43
43
  "pre-commit",
@@ -93,6 +93,9 @@ reportPrivateImportUsage = false
93
93
  [tool.setuptools.packages.find]
94
94
  include = ["bergson*"]
95
95
 
96
+ [tool.uv.pip]
97
+ torch-backend = "auto"
98
+
96
99
  [tool.black]
97
100
  target-version = ["py310", "py311", "py312"]
98
101
 
@@ -0,0 +1,62 @@
1
+ """The eigendecomposition saved across ranks must be that of the summed Gram."""
2
+
3
+ import socket
4
+
5
+ import pytest
6
+ import torch
7
+ import torch.distributed as dist
8
+ import torch.multiprocessing as mp
9
+
10
+ from bergson.gradients import GradientProcessor
11
+ from bergson.process_autocorrelation import process_autocorrelation_matrices
12
+
13
+ SIZES = {"a": 5, "b": 3, "c": 4}
14
+ NUM_ROWS = 12
15
+
16
+
17
+ def rank_rows(rank: int, name: str) -> torch.Tensor:
18
+ gen = torch.Generator().manual_seed(100 * rank + ord(name))
19
+ return torch.randn(NUM_ROWS // 2, SIZES[name], generator=gen, dtype=torch.float64)
20
+
21
+
22
+ def full_gram(name: str, world_size: int) -> torch.Tensor:
23
+ rows = torch.cat([rank_rows(r, name) for r in range(world_size)])
24
+ return rows.mT @ rows / NUM_ROWS
25
+
26
+
27
+ def run(rank: int, world_size: int, port: int, owned: bool, out: dict):
28
+ dist.init_process_group(
29
+ "gloo", init_method=f"tcp://localhost:{port}", rank=rank, world_size=world_size
30
+ )
31
+ if owned:
32
+ # Each rank holds the full Gram of the modules it owns.
33
+ names = [n for i, n in enumerate(SIZES) if i % world_size == rank]
34
+ hessians = {n: full_gram(n, world_size) * NUM_ROWS for n in names}
35
+ else:
36
+ # Each rank holds a partial Gram over its own rows for every module.
37
+ hessians = {n: rank_rows(rank, n).mT @ rank_rows(rank, n) for n in SIZES}
38
+
39
+ processor = GradientProcessor()
40
+ process_autocorrelation_matrices(processor, hessians, NUM_ROWS, SIZES, rank)
41
+ if rank == 0:
42
+ out.update(
43
+ hessians=dict(processor.hessians), eigen=dict(processor.hessians_eigen)
44
+ )
45
+ dist.destroy_process_group()
46
+
47
+
48
+ @pytest.mark.parametrize("owned", [False, True])
49
+ def test_eigen_matches_summed_gram(owned: bool):
50
+ world_size = 2
51
+ with socket.socket() as s:
52
+ s.bind(("", 0))
53
+ port = s.getsockname()[1]
54
+ out = mp.Manager().dict()
55
+ mp.spawn(run, args=(world_size, port, owned, out), nprocs=world_size)
56
+
57
+ for name in SIZES:
58
+ expected = full_gram(name, world_size)
59
+ torch.testing.assert_close(out["hessians"][name], expected)
60
+ eigvals, eigvecs = out["eigen"][name]
61
+ torch.testing.assert_close(eigvals, torch.linalg.eigvalsh(expected))
62
+ torch.testing.assert_close(eigvecs @ torch.diag(eigvals) @ eigvecs.mT, expected)
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes