bergson 2.2.1__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.1 → bergson-2.2.2}/PKG-INFO +25 -31
  2. {bergson-2.2.1 → bergson-2.2.2}/README.md +26 -32
  3. {bergson-2.2.1 → bergson-2.2.2}/bergson/__init__.py +1 -1
  4. {bergson-2.2.1 → bergson-2.2.2}/bergson/config/config.py +10 -0
  5. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/cli.py +59 -49
  6. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/data_stream.py +72 -35
  7. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/fsdp.py +11 -5
  8. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/metasmoothness.py +21 -8
  9. bergson-2.2.2/bergson/magic/shard_load.py +119 -0
  10. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/trainer.py +16 -4
  11. bergson-2.2.2/bergson/score/output_influence.py +154 -0
  12. {bergson-2.2.1 → bergson-2.2.2}/bergson/score/score.py +117 -2
  13. {bergson-2.2.1 → bergson-2.2.2}/bergson/score/scorer.py +4 -0
  14. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/worker_utils.py +21 -8
  15. {bergson-2.2.1 → bergson-2.2.2}/bergson/validate.py +19 -31
  16. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/PKG-INFO +25 -31
  17. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/SOURCES.txt +5 -0
  18. {bergson-2.2.1 → bergson-2.2.2}/pyproject.toml +1 -1
  19. {bergson-2.2.1 → bergson-2.2.2}/tests/test_ddp.py +2 -2
  20. {bergson-2.2.1 → bergson-2.2.2}/tests/test_magic.py +55 -25
  21. {bergson-2.2.1 → bergson-2.2.2}/tests/test_metasmoothness.py +3 -2
  22. bergson-2.2.2/tests/test_output_influence.py +191 -0
  23. bergson-2.2.2/tests/test_padding.py +50 -0
  24. {bergson-2.2.1 → bergson-2.2.2}/tests/test_per_query_magic.py +16 -18
  25. {bergson-2.2.1 → bergson-2.2.2}/tests/test_per_token_lds.py +36 -0
  26. {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_loss_padding.py +2 -2
  27. bergson-2.2.2/tests/test_shard_load.py +46 -0
  28. {bergson-2.2.1 → bergson-2.2.2}/LICENSE +0 -0
  29. {bergson-2.2.1 → bergson-2.2.2}/bergson/__main__.py +0 -0
  30. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/adam_preconditioner.py +0 -0
  31. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/approx_unrolling_math.py +0 -0
  32. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/pipeline.py +0 -0
  33. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/precompute_checkpoints.py +0 -0
  34. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/segment_aggregation.py +0 -0
  35. {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/train_cfg_io.py +0 -0
  36. {bergson-2.2.1 → bergson-2.2.2}/bergson/build.py +0 -0
  37. {bergson-2.2.1 → bergson-2.2.2}/bergson/builder.py +0 -0
  38. {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/__init__.py +0 -0
  39. {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/commands.py +0 -0
  40. {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/trackstar.py +0 -0
  41. {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/trak.py +0 -0
  42. {bergson-2.2.1 → bergson-2.2.2}/bergson/collection.py +0 -0
  43. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/__init__.py +0 -0
  44. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/collector.py +0 -0
  45. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/dist_autocorrelation_gradient_collector.py +0 -0
  46. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/gradient_collectors.py +0 -0
  47. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/in_memory_collector.py +0 -0
  48. {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/projection_matrix.py +0 -0
  49. {bergson-2.2.1 → bergson-2.2.2}/bergson/config/__init__.py +0 -0
  50. {bergson-2.2.1 → bergson-2.2.2}/bergson/config/config_io.py +0 -0
  51. {bergson-2.2.1 → bergson-2.2.2}/bergson/config/validation.py +0 -0
  52. {bergson-2.2.1 → bergson-2.2.2}/bergson/data.py +0 -0
  53. {bergson-2.2.1 → bergson-2.2.2}/bergson/diagnose.py +0 -0
  54. {bergson-2.2.1 → bergson-2.2.2}/bergson/distributed.py +0 -0
  55. {bergson-2.2.1 → bergson-2.2.2}/bergson/format.py +0 -0
  56. {bergson-2.2.1 → bergson-2.2.2}/bergson/gradients.py +0 -0
  57. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/apply_hessian.py +0 -0
  58. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/astra.py +0 -0
  59. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/autocorrelation.py +0 -0
  60. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/eigenvectors.py +0 -0
  61. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/hessian_approximations.py +0 -0
  62. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/inversion.py +0 -0
  63. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/kfac.py +0 -0
  64. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/pipeline.py +0 -0
  65. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/preconditioner.py +0 -0
  66. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/shampoo.py +0 -0
  67. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/sharded_computation.py +0 -0
  68. {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/tkfac.py +0 -0
  69. {bergson-2.2.1 → bergson-2.2.2}/bergson/huggingface.py +0 -0
  70. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/__init__.py +0 -0
  71. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/config.py +0 -0
  72. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/dtensor_patch.py +0 -0
  73. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/grad_accum.py +0 -0
  74. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/optim.py +0 -0
  75. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/rtl_tqdm.py +0 -0
  76. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/score_plot.py +0 -0
  77. {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/swap.py +0 -0
  78. {bergson-2.2.1 → bergson-2.2.2}/bergson/moe.py +0 -0
  79. {bergson-2.2.1 → bergson-2.2.2}/bergson/process_autocorrelation.py +0 -0
  80. {bergson-2.2.1 → bergson-2.2.2}/bergson/process_grads.py +0 -0
  81. {bergson-2.2.1 → bergson-2.2.2}/bergson/query/__init__.py +0 -0
  82. {bergson-2.2.1 → bergson-2.2.2}/bergson/query/attributor.py +0 -0
  83. {bergson-2.2.1 → bergson-2.2.2}/bergson/query/faiss_index.py +0 -0
  84. {bergson-2.2.1 → bergson-2.2.2}/bergson/query/query_index.py +0 -0
  85. {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/__init__.py +0 -0
  86. {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/facts.py +0 -0
  87. {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/generate.py +0 -0
  88. {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/recall.py +0 -0
  89. {bergson-2.2.1 → bergson-2.2.2}/bergson/score/__init__.py +0 -0
  90. {bergson-2.2.1 → bergson-2.2.2}/bergson/score/candidates.py +0 -0
  91. {bergson-2.2.1 → bergson-2.2.2}/bergson/score/score_writer.py +0 -0
  92. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/__init__.py +0 -0
  93. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/batch_size.py +0 -0
  94. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/csv_writer.py +0 -0
  95. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/gradcheck.py +0 -0
  96. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/load_from_optimizer.py +0 -0
  97. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/logger.py +0 -0
  98. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/logging.py +0 -0
  99. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/math.py +0 -0
  100. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/peft.py +0 -0
  101. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/step_state.py +0 -0
  102. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/trainer_export.py +0 -0
  103. {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/utils.py +0 -0
  104. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/dependency_links.txt +0 -0
  105. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/entry_points.txt +0 -0
  106. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/requires.txt +0 -0
  107. {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/top_level.txt +0 -0
  108. {bergson-2.2.1 → bergson-2.2.2}/setup.cfg +0 -0
  109. {bergson-2.2.1 → bergson-2.2.2}/tests/test_adam_state_loading.py +0 -0
  110. {bergson-2.2.1 → bergson-2.2.2}/tests/test_advantages.py +0 -0
  111. {bergson-2.2.1 → bergson-2.2.2}/tests/test_approx_unrolling_checkpoint_parse.py +0 -0
  112. {bergson-2.2.1 → bergson-2.2.2}/tests/test_attention_head_normalizers.py +0 -0
  113. {bergson-2.2.1 → bergson-2.2.2}/tests/test_attribute_tokens.py +0 -0
  114. {bergson-2.2.1 → bergson-2.2.2}/tests/test_attributor.py +0 -0
  115. {bergson-2.2.1 → bergson-2.2.2}/tests/test_bank_loss_cache.py +0 -0
  116. {bergson-2.2.1 → bergson-2.2.2}/tests/test_banked_model_loading.py +0 -0
  117. {bergson-2.2.1 → bergson-2.2.2}/tests/test_batch_size_invariance.py +0 -0
  118. {bergson-2.2.1 → bergson-2.2.2}/tests/test_build.py +0 -0
  119. {bergson-2.2.1 → bergson-2.2.2}/tests/test_builder.py +0 -0
  120. {bergson-2.2.1 → bergson-2.2.2}/tests/test_candidates.py +0 -0
  121. {bergson-2.2.1 → bergson-2.2.2}/tests/test_ckpt_avg_query_grads.py +0 -0
  122. {bergson-2.2.1 → bergson-2.2.2}/tests/test_cli.py +0 -0
  123. {bergson-2.2.1 → bergson-2.2.2}/tests/test_compute_lambda.py +0 -0
  124. {bergson-2.2.1 → bergson-2.2.2}/tests/test_config_runner.py +0 -0
  125. {bergson-2.2.1 → bergson-2.2.2}/tests/test_contrast.py +0 -0
  126. {bergson-2.2.1 → bergson-2.2.2}/tests/test_data.py +0 -0
  127. {bergson-2.2.1 → bergson-2.2.2}/tests/test_diagnose.py +0 -0
  128. {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_batch_budget.py +0 -0
  129. {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_cap.py +0 -0
  130. {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_failure.py +0 -0
  131. {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_magic.py +0 -0
  132. {bergson-2.2.1 → bergson-2.2.2}/tests/test_epoch_shuffle.py +0 -0
  133. {bergson-2.2.1 → bergson-2.2.2}/tests/test_faiss_ann_cpu_load.py +0 -0
  134. {bergson-2.2.1 → bergson-2.2.2}/tests/test_force_math_sdp.py +0 -0
  135. {bergson-2.2.1 → bergson-2.2.2}/tests/test_format.py +0 -0
  136. {bergson-2.2.1 → bergson-2.2.2}/tests/test_global_kfac.py +0 -0
  137. {bergson-2.2.1 → bergson-2.2.2}/tests/test_global_projection.py +0 -0
  138. {bergson-2.2.1 → bergson-2.2.2}/tests/test_grad_clipping.py +0 -0
  139. {bergson-2.2.1 → bergson-2.2.2}/tests/test_gradcheck.py +0 -0
  140. {bergson-2.2.1 → bergson-2.2.2}/tests/test_gradients.py +0 -0
  141. {bergson-2.2.1 → bergson-2.2.2}/tests/test_hessian.py +0 -0
  142. {bergson-2.2.1 → bergson-2.2.2}/tests/test_launch_devices.py +0 -0
  143. {bergson-2.2.1 → bergson-2.2.2}/tests/test_log_fn.py +0 -0
  144. {bergson-2.2.1 → bergson-2.2.2}/tests/test_logit_scale.py +0 -0
  145. {bergson-2.2.1 → bergson-2.2.2}/tests/test_moe.py +0 -0
  146. {bergson-2.2.1 → bergson-2.2.2}/tests/test_multi_query_validate.py +0 -0
  147. {bergson-2.2.1 → bergson-2.2.2}/tests/test_multinode.py +0 -0
  148. {bergson-2.2.1 → bergson-2.2.2}/tests/test_muon.py +0 -0
  149. {bergson-2.2.1 → bergson-2.2.2}/tests/test_normalizer.py +0 -0
  150. {bergson-2.2.1 → bergson-2.2.2}/tests/test_outer_product_gradients.py +0 -0
  151. {bergson-2.2.1 → bergson-2.2.2}/tests/test_pretokenized.py +0 -0
  152. {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_inner_products.py +0 -0
  153. {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_matrix.py +0 -0
  154. {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_settings.py +0 -0
  155. {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_batching_invariance.py +0 -0
  156. {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_set_config.py +0 -0
  157. {bergson-2.2.1 → bergson-2.2.2}/tests/test_recall.py +0 -0
  158. {bergson-2.2.1 → bergson-2.2.2}/tests/test_reduce.py +0 -0
  159. {bergson-2.2.1 → bergson-2.2.2}/tests/test_save_mode_final.py +0 -0
  160. {bergson-2.2.1 → bergson-2.2.2}/tests/test_score.py +0 -0
  161. {bergson-2.2.1 → bergson-2.2.2}/tests/test_score_plot.py +0 -0
  162. {bergson-2.2.1 → bergson-2.2.2}/tests/test_score_writer_dist.py +0 -0
  163. {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_fisher_normalization.py +0 -0
  164. {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_optimizer_variants.py +0 -0
  165. {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_resume_polarity.py +0 -0
  166. {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_trainer_integration.py +0 -0
  167. {bergson-2.2.1 → bergson-2.2.2}/tests/test_step_matrix.py +0 -0
  168. {bergson-2.2.1 → bergson-2.2.2}/tests/test_step_state.py +0 -0
  169. {bergson-2.2.1 → bergson-2.2.2}/tests/test_tokenize.py +0 -0
  170. {bergson-2.2.1 → bergson-2.2.2}/tests/test_trainer_callback.py +0 -0
  171. {bergson-2.2.1 → bergson-2.2.2}/tests/test_trak.py +0 -0
  172. {bergson-2.2.1 → bergson-2.2.2}/tests/test_truncation.py +0 -0
  173. {bergson-2.2.1 → bergson-2.2.2}/tests/test_validate_filter.py +0 -0
  174. {bergson-2.2.1 → bergson-2.2.2}/tests/test_validate_routing.py +0 -0
  175. {bergson-2.2.1 → bergson-2.2.2}/tests/test_validation_config.py +0 -0
  176. {bergson-2.2.1 → bergson-2.2.2}/tests/test_wandb_logging.py +0 -0
  177. {bergson-2.2.1 → 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.1
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
@@ -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
 
@@ -120,7 +120,7 @@ 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
 
@@ -132,33 +132,27 @@ Per-module and per-attention head gradients can be extracted from the store.
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
 
@@ -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.
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] |
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.
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
+
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
+
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
 
@@ -59,7 +59,7 @@ 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
 
@@ -71,33 +71,27 @@ Per-module and per-attention head gradients can be extracted from the store.
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.1"
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."""
@@ -52,7 +52,12 @@ from ..utils.utils import (
52
52
  from ..utils.worker_utils import setup_data_pipeline
53
53
  from ..validate import build_contrast_stream, validate_scores
54
54
  from .config import MagicConfig
55
- from .data_stream import DataStream, mask_padded_rows, pad_dataset_to_batch_size
55
+ from .data_stream import (
56
+ DataStream,
57
+ Padding,
58
+ mask_padded_rows,
59
+ pad_dataset_to_batch_size,
60
+ )
56
61
  from .grad_accum import accumulate_grads
57
62
  from .score_plot import plot_score_trajectory
58
63
  from .trainer import BackwardState, TrainerState, prepare_trainer, write_lr_history
@@ -171,8 +176,7 @@ def compute_per_query_magic_scores(
171
176
  run_cfg: "MagicConfig",
172
177
  world_size: int,
173
178
  global_rank: int,
174
- pad_count: int,
175
- weight_pad_count: int,
179
+ padding: Padding,
176
180
  ) -> torch.Tensor:
177
181
  """Per-query MAGIC scores: one backward per query, sharing the forward.
178
182
 
@@ -243,7 +247,7 @@ def compute_per_query_magic_scores(
243
247
  one = query_dataset.select([qi])
244
248
  if "doc_ids" in one.column_names:
245
249
  one = one.remove_columns("doc_ids")
246
- one, n_one, one_pad, one_wpad = pad_dataset_to_batch_size(
250
+ one, n_one, one_padding = pad_dataset_to_batch_size(
247
251
  one, world_size, 1, f"Query {qi}", global_rank
248
252
  )
249
253
  qstream = DataStream(
@@ -253,8 +257,7 @@ def compute_per_query_magic_scores(
253
257
  input_key=run_cfg.query.prompt_column,
254
258
  weight_shape=(n_one,),
255
259
  )
256
- if one_pad:
257
- qstream.weights.data[-one_wpad:] = 0.0
260
+ one_padding.zero_weights(qstream.weights.data)
258
261
  assert_ckpts_exist()
259
262
  qgrads, _ = compute_query_gradients(
260
263
  fwd_state, model, qstream, "mean", run_cfg.fsdp, run_cfg.grad_accum_steps
@@ -288,9 +291,7 @@ def compute_per_query_magic_scores(
288
291
  if world_size > 1:
289
292
  dist.all_reduce(bwd_state.weight_grads, op=dist.ReduceOp.SUM)
290
293
 
291
- s = bwd_state.weight_grads.detach().cpu()
292
- if pad_count:
293
- s = s[:-weight_pad_count] if s.ndim == 1 else s[:-pad_count]
294
+ s = padding.trim(bwd_state.weight_grads.detach().cpu())
294
295
  if main:
295
296
  # Atomic write
296
297
  torch.save(s, qpath + ".tmp")
@@ -357,7 +358,7 @@ def save_magic_scores(
357
358
  run_path: str,
358
359
  scores: torch.Tensor,
359
360
  train_dataset: Dataset,
360
- pad_count: int,
361
+ padding: Padding,
361
362
  per_token: bool,
362
363
  ) -> str:
363
364
  """Write MAGIC scores as a score directory under ``<run_path>/scores``.
@@ -375,9 +376,19 @@ def save_magic_scores(
375
376
 
376
377
  num_token_grads = compute_num_token_grads(train_dataset)
377
378
  doc_ids = np.asarray(train_dataset["doc_ids"], dtype=np.int64)
378
- if pad_count:
379
- num_token_grads = num_token_grads[:-pad_count]
380
- doc_ids = doc_ids[:-pad_count]
379
+ if "example_ids" in train_dataset.column_names:
380
+ # The grid has one row per example; the stream's rows (shuffled
381
+ # epochs, then padding past the grid) fold onto it.
382
+ example_ids = np.asarray(train_dataset["example_ids"], dtype=np.int64)
383
+ keep = example_ids < len(scores)
384
+ per_example_ntg = np.zeros(len(scores), dtype=np.int64)
385
+ per_example_ntg[example_ids[keep]] = num_token_grads[keep]
386
+ per_example_doc_ids = np.zeros((len(scores), doc_ids.shape[1]), np.int64)
387
+ per_example_doc_ids[example_ids[keep]] = doc_ids[keep]
388
+ num_token_grads, doc_ids = per_example_ntg, per_example_doc_ids
389
+ elif padding:
390
+ num_token_grads = num_token_grads[: -padding.num_rows]
391
+ doc_ids = doc_ids[: -padding.num_rows]
381
392
 
382
393
  offsets = np.zeros(len(num_token_grads) + 1, dtype=np.int64)
383
394
  np.cumsum(num_token_grads, out=offsets[1:])
@@ -395,6 +406,14 @@ def save_magic_scores(
395
406
  return str(path)
396
407
 
397
408
 
409
+ def attach_example_ids(dataset: Dataset) -> Dataset:
410
+ """Number the rows so per-token weights and scores keep one row per
411
+ example across the shuffled epochs. No-op if ``example_ids`` is present."""
412
+ if "example_ids" in dataset.column_names:
413
+ return dataset
414
+ return dataset.add_column("example_ids", list(range(len(dataset))))
415
+
416
+
398
417
  def shuffled_epochs(dataset: Dataset, seed: int, num_epochs: int) -> Dataset:
399
418
  """Concatenate `num_epochs` independently shuffled copies of `dataset`.
400
419
 
@@ -444,24 +463,26 @@ def worker(
444
463
  # Ensure total effective batch size is divisible by world size
445
464
  assert run_cfg.batch_size % world_size == 0
446
465
 
447
- # Pad train dataset to be divisible by batch_size (weight=0 for padding)
448
- train_dataset, num_train_docs, pad_count, weight_pad_count = (
449
- pad_dataset_to_batch_size(
450
- train_dataset, run_cfg.batch_size, num_train_docs, "Train", global_rank
451
- )
452
- )
453
-
454
466
  # Plain magic runs enter with score_path="" (scores are computed below).
455
467
  per_token = (isinstance(run_cfg, MagicConfig) and run_cfg.attribute_tokens) or (
456
468
  score_path and scores_are_per_token(score_path)
457
469
  )
470
+ if per_token:
471
+ train_dataset = attach_example_ids(train_dataset)
472
+
473
+ # Pad train dataset to be divisible by batch_size (weight=0 for padding)
474
+ train_dataset, num_train_docs, padding = pad_dataset_to_batch_size(
475
+ train_dataset, run_cfg.batch_size, num_train_docs, "Train", global_rank
476
+ )
477
+
458
478
  if per_token:
459
479
  seq_len = run_cfg.data.chunk_length
460
480
  if seq_len <= 0:
461
481
  seq_len = max(train_dataset["length"])
462
482
  print(f"Using max sequence length {seq_len} for per-token attribution")
463
483
 
464
- w_shape = (len(train_dataset), seq_len)
484
+ # One weight row per example; the pad rows share one synthetic row.
485
+ w_shape = (max(train_dataset["example_ids"]) + 1, seq_len)
465
486
  else:
466
487
  w_shape = (num_train_docs,)
467
488
 
@@ -472,11 +493,7 @@ def worker(
472
493
  input_key=run_cfg.data.prompt_column,
473
494
  weight_shape=w_shape,
474
495
  )
475
- if pad_count:
476
- if stream.weights.ndim == 1:
477
- stream.weights.data[-weight_pad_count:] = 0.0
478
- else:
479
- stream.weights.data[-pad_count:] = 0.0
496
+ padding.zero_weights(stream.weights.data)
480
497
 
481
498
  log_fn = None
482
499
  if run_cfg.wandb_project and global_rank == 0:
@@ -564,10 +581,8 @@ def worker(
564
581
 
565
582
  # Pad query dataset to be divisible by batch_size (weight=0 for padding)
566
583
  num_real_query_docs = num_query_docs
567
- query_dataset, num_query_docs, query_pad_count, query_weight_pad_count = (
568
- pad_dataset_to_batch_size(
569
- query_dataset, run_cfg.batch_size, num_query_docs, "Query", global_rank
570
- )
584
+ query_dataset, num_query_docs, query_padding = pad_dataset_to_batch_size(
585
+ query_dataset, run_cfg.batch_size, num_query_docs, "Query", global_rank
571
586
  )
572
587
  if len(query_dataset) < run_cfg.batch_size:
573
588
  raise ValueError(
@@ -584,9 +599,8 @@ def worker(
584
599
  input_key=run_cfg.query.prompt_column,
585
600
  weight_shape=(num_query_docs,),
586
601
  )
587
- if query_pad_count:
588
- # query_stream.weights is always 1D (weight_shape=(num_query_docs,))
589
- query_stream.weights.data[-query_weight_pad_count:] = 0.0
602
+ # query_stream.weights is always 1D (weight_shape=(num_query_docs,))
603
+ query_padding.zero_weights(query_stream.weights.data)
590
604
 
591
605
  query_grads, baseline = compute_query_gradients(
592
606
  fwd_state,
@@ -635,15 +649,14 @@ def worker(
635
649
  run_cfg,
636
650
  world_size,
637
651
  global_rank,
638
- pad_count,
639
- weight_pad_count,
652
+ padding,
640
653
  )
641
654
  multi_query = True
642
655
  if global_rank == 0:
643
656
  print(f"Baseline loss: {baseline}")
644
657
  print(f"Score summary: {describe(scores.flatten())}")
645
658
  score_path = save_magic_scores(
646
- run_cfg.run_path, scores, train_dataset, pad_count, bool(per_token)
659
+ run_cfg.run_path, scores, train_dataset, padding, bool(per_token)
647
660
  )
648
661
  plot_score_trajectory(
649
662
  scores.float().numpy(),
@@ -688,12 +701,7 @@ def worker(
688
701
  if world_size > 1:
689
702
  dist.all_reduce(bwd_state.weight_grads, op=dist.ReduceOp.SUM)
690
703
 
691
- scores = bwd_state.weight_grads.cpu()
692
- if pad_count:
693
- if scores.ndim == 1:
694
- scores = scores[:-weight_pad_count]
695
- else:
696
- scores = scores[:-pad_count]
704
+ scores = padding.trim(bwd_state.weight_grads.cpu())
697
705
 
698
706
  if global_rank == 0:
699
707
  print(f"Baseline loss: {baseline}")
@@ -702,7 +710,7 @@ def worker(
702
710
  print(f"Score summary: {summ}")
703
711
 
704
712
  score_path = save_magic_scores(
705
- run_cfg.run_path, scores, train_dataset, pad_count, bool(per_token)
713
+ run_cfg.run_path, scores, train_dataset, padding, bool(per_token)
706
714
  )
707
715
  plot_score_trajectory(
708
716
  scores.float().numpy(),
@@ -731,9 +739,8 @@ def worker(
731
739
  model=model,
732
740
  baseline=baseline,
733
741
  num_query_docs=num_query_docs,
734
- query_weight_pad_count=query_weight_pad_count,
735
- pad_count=pad_count,
736
- weight_pad_count=weight_pad_count,
742
+ query_padding=query_padding,
743
+ padding=padding,
737
744
  )
738
745
 
739
746
 
@@ -766,15 +773,18 @@ def run_magic(
766
773
  save_run_config(run_cfg, run_path)
767
774
 
768
775
  # HF datasets caches are not safe for concurrent writers, so the main node
769
- # must finish populating the cache before others read from it.
770
- barrier = run_path / ".preprocess_done" if multi_node else None
776
+ # must finish populating the cache before others read from it. The job id
777
+ # keeps a stale barrier from releasing the wait early.
778
+ job_id = os.environ.get("SLURM_JOB_ID", "")
779
+ barrier = run_path / f".preprocess_done{job_id}" if multi_node else None
771
780
  if barrier is not None and not is_main_node:
772
- run_path.mkdir(parents=True, exist_ok=True)
781
+ # Don't create run_path here to avoid multi-node hang.
773
782
  while not barrier.exists():
774
783
  time.sleep(0.5)
775
784
 
776
785
  train_ds, train_n = setup_data_pipeline(run_cfg)
777
786
  train_ds = attach_doc_ids_if_missing(train_ds)
787
+ train_ds = attach_example_ids(train_ds)
778
788
 
779
789
  train_ds = shuffled_epochs(train_ds, run_cfg.seed, max(1, run_cfg.num_epochs))
780
790
 
@@ -1,6 +1,9 @@
1
+ from dataclasses import dataclass
2
+
1
3
  import torch
2
4
  import torch.distributed as dist
3
- from datasets import Dataset
5
+ from datasets import Dataset, concatenate_datasets
6
+ from torch import Tensor
4
7
 
5
8
  from ..data import pad_and_tensor
6
9
 
@@ -18,61 +21,88 @@ def mask_padded_rows(batch: dict) -> tuple[dict, int]:
18
21
  return batch, int((batch["labels"][:, 1:] != -100).sum())
19
22
 
20
23
 
24
+ @dataclass(frozen=True)
25
+ class Padding:
26
+ """How much of the padded dataset is padding, and how to keep padding from
27
+ affecting influence scores.
28
+
29
+ A run weights either documents, as a vector with one entry per document,
30
+ or tokens, as a grid with one row per example. ``num_docs`` counts the
31
+ vector entries that are pure padding and ``num_examples`` the grid rows.
32
+ Each is one if the pad rows share a synthetic id, ``doc_ids`` for the
33
+ first and ``example_ids`` for the second, and ``num_rows`` if they do not;
34
+ the two are decided separately. Pass the weights to :meth:`zero_weights`
35
+ and the scores to :meth:`trim`, which use whichever count matches.
36
+ """
37
+
38
+ num_rows: int = 0
39
+ num_docs: int = 0
40
+ num_examples: int = 0
41
+
42
+ def __bool__(self) -> bool:
43
+ return bool(self.num_rows)
44
+
45
+ def _trailing_pad_len(self, t: Tensor) -> int:
46
+ return self.num_docs if t.ndim == 1 else self.num_examples
47
+
48
+ def zero_weights(self, weights: Tensor) -> None:
49
+ """Zero the pad rows' weights in place."""
50
+ if n := self._trailing_pad_len(weights):
51
+ weights.data[-n:] = 0.0
52
+
53
+ def trim(self, scores: Tensor) -> Tensor:
54
+ """Return ``scores`` without the entries the pad rows produced."""
55
+ n = self._trailing_pad_len(scores)
56
+ return scores[:-n] if n else scores
57
+
58
+
21
59
  def pad_dataset_to_batch_size(
22
60
  dataset: Dataset,
23
61
  batch_size: int,
24
62
  num_docs: int,
25
63
  label: str,
26
64
  global_rank: int,
27
- ) -> tuple[Dataset, int, int, int]:
65
+ ) -> tuple[Dataset, int, Padding]:
28
66
  """Pad dataset to be divisible by batch_size by repeating the last example.
29
67
 
30
- Returns (padded_dataset, num_docs, pad_count, weight_pad_count).
31
-
32
- `pad_count` is the number of rows appended to the dataset (0 if unchanged).
33
- `weight_pad_count` is the number of trailing entries of a *1D* per-doc
34
- weight tensor that should be zeroed to silence the pad rows' training
35
- contribution.
36
-
37
- - If the dataset has a "doc_ids" column, `.select(total - 1, ...)` copies
38
- the last doc's doc_ids into every pad row. Zeroing the last `pad_count`
39
- entries of a weights-indexed-by-doc_id tensor would silence real docs,
40
- so we instead route pad rows to a fresh synthetic doc id (=num_docs),
41
- bump num_docs by 1, and set `weight_pad_count = 1`.
42
- - Otherwise rows are self-identified docs: num_docs becomes the padded
43
- length and `weight_pad_count = pad_count` zeros the pad rows directly.
44
-
45
- In per-token (2D) mode callers should zero `weights[-pad_count:]` instead
46
- — `weight_pad_count` applies only to 1D per-doc weights.
68
+ Repeating a row copies its ``doc_ids``, which would add the pad rows'
69
+ scores to the last document; they are given a separate document id
70
+ instead, so one zeroed weight entry covers all of them. An ``example_ids``
71
+ column, which per-token weights index by, is re-pointed the same way.
47
72
  """
48
73
  remainder = len(dataset) % batch_size
49
74
  if not remainder:
50
- return dataset, num_docs, 0, 0
75
+ return dataset, num_docs, Padding()
51
76
 
52
77
  pad_count = batch_size - remainder
53
78
  total = len(dataset)
54
- pad_indices = list(range(total)) + [total - 1] * pad_count
55
- dataset = dataset.select(pad_indices)
79
+ last = dataset[total - 1]
80
+
81
+ if "example_ids" in dataset.column_names:
82
+ last = {**last, "example_ids": max(dataset["example_ids"]) + 1}
83
+ num_examples = 1
84
+ else:
85
+ num_examples = pad_count
56
86
 
57
87
  if "doc_ids" in dataset.column_names:
58
- synthetic_doc_id = num_docs
59
- new_doc_ids = [
60
- row if i < total else [synthetic_doc_id] * len(row)
61
- for i, row in enumerate(dataset["doc_ids"])
62
- ]
63
- dataset = dataset.remove_columns("doc_ids").add_column("doc_ids", new_doc_ids)
88
+ last = {**last, "doc_ids": [num_docs] * len(last["doc_ids"])}
64
89
  num_docs += 1
65
- weight_pad_count = 1
90
+ padding = Padding(pad_count, num_docs=1, num_examples=num_examples)
66
91
  else:
67
- num_docs = len(dataset)
68
- weight_pad_count = pad_count
69
-
92
+ num_docs = total + pad_count
93
+ padding = Padding(pad_count, num_docs=pad_count, num_examples=num_examples)
94
+
95
+ # Built in memory: the pad rows are a handful of copies of the last row.
96
+ pad_rows = Dataset.from_dict(
97
+ {k: [v] * pad_count for k, v in last.items()}, features=dataset.features
98
+ )
99
+ dataset = concatenate_datasets([dataset, pad_rows])
70
100
  if global_rank == 0:
71
101
  print(
72
102
  f"{label}: padded {pad_count}/{total + pad_count} examples "
73
103
  f"(weight=0) to fill last batch"
74
104
  )
75
- return dataset, num_docs, pad_count, weight_pad_count
105
+ return dataset, num_docs, padding
76
106
 
77
107
 
78
108
  class DataStream:
@@ -130,8 +160,15 @@ class DataStream:
130
160
  # If the weights are 1D, we assume they correspond to documents and look for
131
161
  # "doc_ids" in the batch to index them. If they're 2D, they correspond to tokens
132
162
  if self.weights.ndim == 2:
133
- # Truncate to the max sequence length in the batch to avoid indexing errors
134
- indices = (indices, slice(None, x.shape[1]))
163
+ # One weight row per example: a shuffled multi-epoch stream reaches
164
+ # its rows through "example_ids". Truncate to the max sequence
165
+ # length in the batch to avoid indexing errors.
166
+ rows = (
167
+ torch.tensor(batch["example_ids"], device=self.device)
168
+ if "example_ids" in batch
169
+ else indices
170
+ )
171
+ indices = (rows, slice(None, x.shape[1]))
135
172
  elif "doc_ids" in batch:
136
173
  indices = torch.tensor(batch["doc_ids"], device=self.device)
137
174
  # doc_ids may be longer than the per-batch padded seq_len (unpacked