Predictive data attribution
Joint work with
Hamza Golubovic, Han Tong
Arian Maleki, Andrew Ilyas
Applications of predictive data attribution
Which data removals would improve performance on a target task?
Which training examples would change a particular prediction if removed?
How would excluding a set of records affect held-out behavior?
Example: attributing LLM predictions
Remove a specified subset of training text
Next-token loss after a held-out passage
How that loss would change if the model had trained without this text
Parameterizing predictive data attribution
Each training example receives a continuous weight
Model output function
Data weights
Trained parameters
Final measurement
Meta-gradients
Gradients through the whole training procedure
MAGIC
MAGIC uses a first-order Taylor expansion
REPLAY computes the meta-gradients
MAGIC is state of the art for predictive data attribution
ResNet-9 / CIFAR-10 · random 1% deletions · mean Spearman correlation
Challenge
Computation scales linearly with the number of queries
Can we attribute the whole test dataset with fewer replays?
Attribution as matrix estimation
Estimate the influence matrix with B<k measurements
Two distinct objectives
Recover the influence scores
Predict the effects of data interventions
The chain rule factors the influence matrix
The Jacobian of the training procedure
Query gradients at the trained parameters
One replay can measure a combination of rows
One replay measures one query’s row
Combining query gradients gives a measured direction u
Orthogonal projection
Let U=span{u1,…,uB}, with orthogonal ub
We know U but still need ⟨Yq,ub⟩ for every query
Forward mode completes the projection
- Forward mode computes Jub⊤
- Multiplying by V⊤ gives the inner products
An online viewpoint of PCA
Maximizing captured energy minimizes reconstruction error
MAGE
We use the query-gradient Gram matrix as a proxy
Probe its leading eigenvectors in order
From LDS to a projection objective
A normalized PCA objective for attribution
The squared-correlation surrogate maximizes captured energy
of the row-normalized matrix
SPELL
Approximate PCA on normalized rows of the influence matrix
Start with MAGE measurements
Estimate row norms from the current projection
Choose the next direction from the normalized residual Gram matrix
MAGE closely tracks the reconstruction oracle
k=500, B=100 · mean ± SE · rank-B SVD requires the full matrix
SPELL leads on predictive data attribution
k=500, B=100 · 1% deletions · mean ± SE · oracle SVD · full MAGIC uses B=500
SPELL gains accuracy on queries with small meta-gradients
1% deletions · CIFAR-10: k=1000,B=500 · TinyStories: k=600,B=300