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