Distilling Diffusion Score Discrepancy for Efficient Training Data Attribution
Organizations: University of Illinois at Urbana-Champaign · Sony AI · University of Texas at Austin · Sony Group Corporation
Abstract
Training data attribution for diffusion models aims to identify the training samples that influence a generated instance, but existing methods either require costly per-sample gradient computation or query-specific model optimization. Moreover, most methods attribute changes in a proxy loss rather than changes in the actual model's generative behavior. We address these limitations by formulating attribution directly with a local score discrepancy measure, which applies to any diffusion variant (including DDPM, EDM, and flow matching), and by showing that such measure can be estimated without retraining, as a preconditioned gradient similarity. We instantiate this estimator as Training-data Influence via score Discrepancy (TID), which uses Kronecker-factored curvature to avoid random projections and per-sample gradient storage. We then distill TID into TIDE, a forward-only student trained online to reproduce the teacher's rankings from the diffusion model's internal activations. Under counterfactual evaluation on CIFAR-10, ArtBench-10, and MS-COCO, TID matches or outperforms state-of-the-art approaches, while TIDE retains most of TID's accuracy at four to five orders of magnitude lower per-query cost, attributing generated samples in milliseconds and faster than the generation itself.
Figures & tables
| Approach | Categ. | SSIM | SSCD | LPIPS | CLIP | ||
|---|---|---|---|---|---|---|---|
| CIFAR-10 | CLIP | MA | 0.598 0.011 | 0.585 0.008 | 0.610 0.013 | 0.661 0.010 | 0.613 0.007 |
| DINO | MA | 0.604 0.018 | 0.616 0.017 | 0.638 0.021 | 0.665 0.027 | 0.631 0.019 | |
| D-TRAK | G | 0.691 0.037 | 0.649 0.033 | 0.703 0.026 | 0.640 0.042 | 0.671 0.033 | |
| DAS | G | 0.760 0.031 | 0.753 0.036 | 0.789 0.040 | 0.746 0.035 | 0.762 0.035 | |
| MUCS | U | 0.864 0.025 | 0.854 0.023 | 0.908 0.015 | 0.857 0.028 | 0.870 0.021 | |
| TID (ours) | G | 0.836 0.031 | 0.815 0.031 | 0.872 0.023 | 0.818 0.027 | 0.835 0.027 |
| RQ/Variant | SSIM | SSCD | LPIPS | CLIP | ||
| TID (ours) | 0.836 | 0.815 | 0.872 | 0.818 | 0.835 | — |
| RQ1: Attribution target | ||||||
| 0.815 | 0.797 | 0.855 | 0.790 | 0.814 | 2.5% | |
| 0.777 | 0.737 | 0.785 | 0.731 | 0.757 | 9.3% | |
| RQ2: Corruption alignment | ||||||
| Independent draws | 0.759 | 0.743 | 0.798 | 0.732 | 0.758 | 9.2% |
Appendix figures & tables12 assets
Supplementary material from the paper’s appendix.
Appendix
| Method | Target | Curvature | Kernel | Normalization |
|---|---|---|---|---|
| D-TRAK | GN ridge | random projection | none | |
| DAS | GN ridge | random projection | gradient | |
| K-FAC influence | MC-GGN | K-FAC | none | |
| TID (ours) | empirical Fisher | K-FAC | factors | |
| Ablation variants of TID | ||||
| RQ1a: squared output | empirical Fisher | K-FAC | factors | |
| CIFAR-10 | ArtBench-10 | MS-COCO | ||
| Data | Training samples | 50,000 | 49,917 | 118,287 |
| Resolution | ||||
| Conditioning | none | class label | CLIP text embedding | |
| Backbone | Architecture | DiT | DiT | DiT |
| Transformer blocks | 12 | 12 | 12 | |
| Model width | 768 | 768 | 768 |
| Setting | ||
| TID | Timestep schedule | Karras |
| Corruption draw | shared | |
| Relative damping | 0.1 | |
| Prenormalization | per sample | |
| Inference draws | 32 | |
| Inference draw alignment | aligned |
| Method | CIFAR-10 | ArtBench-10 | MS-COCO |
|---|---|---|---|
| Random | 16,636 (33.3%) | 16,589 (33.2%) | 39,332 (33.3%) |
| CLIP | 13,025 (26.1%) | 10,349 (20.7%) | 34,384 (29.1%) |
| DINO | 13,728 (27.5%) | 11,857 (23.8%) | 35,952 (30.4%) |
| D-TRAK | 15,586 (31.2%) | 13,874 (27.8%) | 35,154 (29.7%) |
| DAS | 13,758 (27.5%) | 14,262 (28.6%) | 38,006 (32.1%) |
| MUCS | 15,448 (30.9%) | 14,093 (28.2%) | 34,096 (28.8%) |
| TIDE (ours) | 0.784 | 0.858 | 0.910 | 0.938 |
|---|
| Approach | Category | SSIM | SSCD | LPIPS | CLIP | ||
|---|---|---|---|---|---|---|---|
| CIFAR-10 | CLIP | MA | |||||
| DINO | MA | ||||||
| D-TRAK | G | ||||||
| DAS | G | ||||||
| MUCS | U | ||||||
| TID (ours) | G |
| Method | One-off compute | Cache | Per-query compute | Score flops |
|---|---|---|---|---|
| CLIP / DINO | ||||
| D-TRAK | ||||
| DAS | ||||
| TID | ||||
| TIDE | ||||
| MUCS | — | — |
| CIFAR-10 | ArtBench-10 | MS-COCO | ||
| Method | One-off computation | wall-clock / cache size | ||
| CLIP | encode training set | 98 s / 0.10 GB | 112 s / 0.10 GB | 277 s / 0.24 GB |
| DINO | encode training set | 18 s / 0.15 GB | 18 s / 0.15 GB | 37 s / 0.36 GB |
| DAS | error + projected-gradient store | 18.1 h / 3.3 GB | 12.3 h / 3.3 GB | 38.1 h / 7.8 GB |
| TIDE | distillation + embedding collection | 3.8 h / 0.15 GB | 4.2 h / 0.15 GB | 10.0 h / 0.36 GB |
| TID ( ) | factor estimation | 0.2 h / 2.4 GB | 0.2 h / 2.4 GB | 0.5 h / 2.5 GB |