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.
- {bergson-2.2.1 → bergson-2.2.2}/PKG-INFO +25 -31
- {bergson-2.2.1 → bergson-2.2.2}/README.md +26 -32
- {bergson-2.2.1 → bergson-2.2.2}/bergson/__init__.py +1 -1
- {bergson-2.2.1 → bergson-2.2.2}/bergson/config/config.py +10 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/cli.py +59 -49
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/data_stream.py +72 -35
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/fsdp.py +11 -5
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/metasmoothness.py +21 -8
- bergson-2.2.2/bergson/magic/shard_load.py +119 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/trainer.py +16 -4
- bergson-2.2.2/bergson/score/output_influence.py +154 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/score/score.py +117 -2
- {bergson-2.2.1 → bergson-2.2.2}/bergson/score/scorer.py +4 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/worker_utils.py +21 -8
- {bergson-2.2.1 → bergson-2.2.2}/bergson/validate.py +19 -31
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/PKG-INFO +25 -31
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/SOURCES.txt +5 -0
- {bergson-2.2.1 → bergson-2.2.2}/pyproject.toml +1 -1
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_ddp.py +2 -2
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_magic.py +55 -25
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_metasmoothness.py +3 -2
- bergson-2.2.2/tests/test_output_influence.py +191 -0
- bergson-2.2.2/tests/test_padding.py +50 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_per_query_magic.py +16 -18
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_per_token_lds.py +36 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_loss_padding.py +2 -2
- bergson-2.2.2/tests/test_shard_load.py +46 -0
- {bergson-2.2.1 → bergson-2.2.2}/LICENSE +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/__main__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/adam_preconditioner.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/approx_unrolling_math.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/pipeline.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/precompute_checkpoints.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/segment_aggregation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/approx_unrolling/train_cfg_io.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/build.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/builder.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/commands.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/trackstar.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/cli/trak.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collection.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/collector.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/dist_autocorrelation_gradient_collector.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/gradient_collectors.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/in_memory_collector.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/collector/projection_matrix.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/config/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/config/config_io.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/config/validation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/data.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/diagnose.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/distributed.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/format.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/gradients.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/apply_hessian.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/astra.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/autocorrelation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/eigenvectors.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/hessian_approximations.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/inversion.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/kfac.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/pipeline.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/preconditioner.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/shampoo.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/sharded_computation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/hessians/tkfac.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/huggingface.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/config.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/dtensor_patch.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/grad_accum.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/optim.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/rtl_tqdm.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/score_plot.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/magic/swap.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/moe.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/process_autocorrelation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/process_grads.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/query/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/query/attributor.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/query/faiss_index.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/query/query_index.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/facts.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/generate.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/recall/recall.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/score/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/score/candidates.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/score/score_writer.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/__init__.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/batch_size.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/csv_writer.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/gradcheck.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/load_from_optimizer.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/logger.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/logging.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/math.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/peft.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/step_state.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/trainer_export.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson/utils/utils.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/dependency_links.txt +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/entry_points.txt +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/requires.txt +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/bergson.egg-info/top_level.txt +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/setup.cfg +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_adam_state_loading.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_advantages.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_approx_unrolling_checkpoint_parse.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_attention_head_normalizers.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_attribute_tokens.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_attributor.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_bank_loss_cache.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_banked_model_loading.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_batch_size_invariance.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_build.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_builder.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_candidates.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_ckpt_avg_query_grads.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_cli.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_compute_lambda.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_config_runner.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_contrast.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_data.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_diagnose.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_batch_budget.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_cap.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_failure.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_distributed_magic.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_epoch_shuffle.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_faiss_ann_cpu_load.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_force_math_sdp.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_format.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_global_kfac.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_global_projection.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_grad_clipping.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_gradcheck.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_gradients.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_hessian.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_launch_devices.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_log_fn.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_logit_scale.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_moe.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_multi_query_validate.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_multinode.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_muon.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_normalizer.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_outer_product_gradients.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_pretokenized.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_inner_products.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_matrix.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_projection_settings.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_batching_invariance.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_query_set_config.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_recall.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_reduce.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_save_mode_final.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_score.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_score_plot.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_score_writer_dist.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_fisher_normalization.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_optimizer_variants.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_resume_polarity.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_source_trainer_integration.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_step_matrix.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_step_state.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_tokenize.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_trainer_callback.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_trak.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_truncation.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_validate_filter.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_validate_routing.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_validation_config.py +0 -0
- {bergson-2.2.1 → bergson-2.2.2}/tests/test_wandb_logging.py +0 -0
- {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.
|
|
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 |
|
|
92
|
-
|
|
93
|
-
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] |
|
|
94
|
-
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] |
|
|
95
|
-
|
|
|
96
|
-
|
|
|
97
|
-
|
|
|
98
|
-
| [
|
|
99
|
-
|
|
|
100
|
-
|
|
|
101
|
-
|
|
|
102
|
-
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] |
|
|
103
|
-
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] |
|
|
91
|
+
| Method | QLD | LDS | Per-token QLD |
|
|
92
|
+
|:---|:---:|:---:|:---:|
|
|
93
|
+
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] | 1.375 [1.348, 1.403] |
|
|
94
|
+
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] | 1.082 [1.047, 1.116] |
|
|
95
|
+
| Eigenvalue-corrected Shampoo | 0.071 [0.060, 0.082] | 0.517 [0.491, 0.539] | 0.922 [0.887, 0.957] |
|
|
96
|
+
| SOURCE (Adam, EK-FAC) | 0.071 [0.060, 0.084] | 0.473 [0.446, 0.498] | 0.908 [0.877, 0.942] |
|
|
97
|
+
| EK-FAC | 0.070 [0.058, 0.082] | 0.454 [0.426, 0.479] | 0.863 [0.833, 0.894] |
|
|
98
|
+
| BM25 | 0.062 [0.048, 0.076] | 0.220 [0.185, 0.252] | 0.677 [0.650, 0.704] |
|
|
99
|
+
| Semantic search ([Qwen3 8B](https://huggingface.co/Qwen/Qwen3-Embedding-8B)) | 0.049 [0.038, 0.061] | 0.132 [0.093, 0.169] | 0.008 [0.006, 0.010] |
|
|
100
|
+
| TrackStar (no optimizer correction, projection 64) | 0.045 [0.036, 0.055] | 0.270 [0.240, 0.295] | 0.489 [0.465, 0.515] |
|
|
101
|
+
| TRAK (8-model ensemble) | 0.032 [0.024, 0.040] | 0.138 [0.111, 0.165] | 0.119 [0.112, 0.125] |
|
|
102
|
+
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] | 0.217 [0.206, 0.229] |
|
|
103
|
+
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] | 0.028 [0.018, 0.039] |
|
|
104
104
|
|
|
105
|
-
Results for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
105
|
+
Results with 95% confidence intervals for GPT-2 finetuned on 4 epochs of the WikiText corpus. Held-out loss dropped from 3.545 to 3.111 over training.
|
|
106
106
|
|
|
107
107
|
## Functionality
|
|
108
108
|
|
|
@@ -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
|
-
#
|
|
135
|
+
# Getting Started
|
|
136
136
|
|
|
137
|
-
There are many example YAMLs in the `examples` directory.
|
|
137
|
+
There are many example YAMLs in the `examples` directory, including various paper experiment replications. Use `bergson <yaml_path>` to run them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
|
|
138
138
|
|
|
139
139
|
```bash
|
|
140
140
|
bergson examples/magic/gpt2_wikitext_tiny.yaml
|
|
141
141
|
```
|
|
142
142
|
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
| | Notebook | Description |
|
|
146
|
-
|---|----------|-------------|
|
|
147
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
148
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
149
|
-
|
|
150
|
-
Construct and query an on-disk index of randomly projected gradients:
|
|
143
|
+
You can use the same fields used to specify experiments in the YAMLs to run experiments directly in the Bergson CLI. For example, to construct and query an on-disk index of randomly projected gradients from the CLI:
|
|
151
144
|
|
|
152
145
|
```bash
|
|
153
146
|
bergson build runs/index --model EleutherAI/pythia-14m --dataset NeelNanda/pile-10k --truncation --token_batch_size 4096 --projection_dim 16
|
|
154
147
|
bergson query --index runs/index --unit_norm
|
|
155
148
|
```
|
|
156
149
|
|
|
157
|
-
|
|
150
|
+
Or check out a notebook for programmatic usage:
|
|
158
151
|
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
152
|
+
| | Notebook | Description |
|
|
153
|
+
|---|----------|-------------|
|
|
154
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
155
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
162
156
|
|
|
163
157
|
# Development
|
|
164
158
|
|
|
@@ -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 |
|
|
31
|
-
|
|
32
|
-
| MAGIC | 0.100 [0.090, 0.112] | 0.931 [0.925, 0.936] |
|
|
33
|
-
| EK-FAC + ASTRA | 0.074 [0.063, 0.087] | 0.643 [0.624, 0.660] |
|
|
34
|
-
|
|
|
35
|
-
|
|
|
36
|
-
|
|
|
37
|
-
| [
|
|
38
|
-
|
|
|
39
|
-
|
|
|
40
|
-
|
|
|
41
|
-
| Gradient cosine similarity | 0.021 [0.016, 0.027] | 0.156 [0.131, 0.181] |
|
|
42
|
-
| Activation similarity | 0.000 [-0.000, 0.001] | 0.110 [0.070, 0.149] |
|
|
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
|
-
#
|
|
74
|
+
# Getting Started
|
|
75
75
|
|
|
76
|
-
There are many example YAMLs in the `examples` directory.
|
|
76
|
+
There are many example YAMLs in the `examples` directory, including various paper experiment replications. Use `bergson <yaml_path>` to run them. For example, to MAGIC-attribute a GPT-2 WikiText fine-tune:
|
|
77
77
|
|
|
78
78
|
```bash
|
|
79
79
|
bergson examples/magic/gpt2_wikitext_tiny.yaml
|
|
80
80
|
```
|
|
81
81
|
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
| | Notebook | Description |
|
|
85
|
-
|---|----------|-------------|
|
|
86
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
87
|
-
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
88
|
-
|
|
89
|
-
Construct and query an on-disk index of randomly projected gradients:
|
|
82
|
+
You can use the same fields used to specify experiments in the YAMLs to run experiments directly in the Bergson CLI. For example, to construct and query an on-disk index of randomly projected gradients from the CLI:
|
|
90
83
|
|
|
91
84
|
```bash
|
|
92
85
|
bergson build runs/index --model EleutherAI/pythia-14m --dataset NeelNanda/pile-10k --truncation --token_batch_size 4096 --projection_dim 16
|
|
93
86
|
bergson query --index runs/index --unit_norm
|
|
94
87
|
```
|
|
95
88
|
|
|
96
|
-
|
|
89
|
+
Or check out a notebook for programmatic usage:
|
|
97
90
|
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
91
|
+
| | Notebook | Description |
|
|
92
|
+
|---|----------|-------------|
|
|
93
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/poison_detection.ipynb) | **Poison Detection** | Detect poisoned training examples with gradient attribution (T4, ~5 min) |
|
|
94
|
+
| [](https://colab.research.google.com/github/EleutherAI/bergson/blob/main/notebooks/style_ablation.ipynb) | **Style Ablation** | Suppress style to recover semantic matching (A100, ~20 min) |
|
|
101
95
|
|
|
102
96
|
# Development
|
|
103
97
|
|
|
@@ -885,6 +885,16 @@ class ScoreConfig(Serializable):
|
|
|
885
885
|
capability (e.g. in influence functions). False for unrolled
|
|
886
886
|
differentiation."""
|
|
887
887
|
|
|
888
|
+
token_influence: Literal["gradient", "output"] = "gradient"
|
|
889
|
+
"""What each row scores with ``attribute_tokens``. ``gradient`` scores row
|
|
890
|
+
``t`` by the per-token gradient at position ``t``, which is position ``t``'s
|
|
891
|
+
effect on the loss of every later token. ``output`` scores row ``t`` by the
|
|
892
|
+
loss on token ``t + 1`` alone, the output token influence of Grosse et al.
|
|
893
|
+
(2023). Both kinds of row sum to the per-example score, so without
|
|
894
|
+
``attribute_tokens``, ``output`` only changes how that score is computed.
|
|
895
|
+
``output`` costs one forward-mode pass per query column, so aggregate the
|
|
896
|
+
query when you can, and needs an unprojected query and dot-product scoring."""
|
|
897
|
+
|
|
888
898
|
candidates: CandidateConfig = field(default_factory=CandidateConfig)
|
|
889
899
|
"""Score only the rows an earlier run ranked highest. The store then has one
|
|
890
900
|
row per candidate; ``candidates.npy`` beside it holds their training rows."""
|
|
@@ -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
|
|
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
|
-
|
|
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,
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
379
|
-
|
|
380
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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,
|
|
568
|
-
|
|
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
|
-
|
|
588
|
-
|
|
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
|
-
|
|
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,
|
|
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,
|
|
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
|
-
|
|
735
|
-
|
|
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
|
-
|
|
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.
|
|
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,
|
|
65
|
+
) -> tuple[Dataset, int, Padding]:
|
|
28
66
|
"""Pad dataset to be divisible by batch_size by repeating the last example.
|
|
29
67
|
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
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,
|
|
75
|
+
return dataset, num_docs, Padding()
|
|
51
76
|
|
|
52
77
|
pad_count = batch_size - remainder
|
|
53
78
|
total = len(dataset)
|
|
54
|
-
|
|
55
|
-
|
|
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
|
-
|
|
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
|
-
|
|
90
|
+
padding = Padding(pad_count, num_docs=1, num_examples=num_examples)
|
|
66
91
|
else:
|
|
67
|
-
num_docs =
|
|
68
|
-
|
|
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,
|
|
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
|
-
#
|
|
134
|
-
|
|
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
|