Predictive data attribution

Predictive data attribution
MAGIC and sketched metadifferentiation
Eric
01 / 25

Joint work with

Hamza Golubovic, Han Tong

Arian Maleki, Andrew Ilyas

02 / 25

Applications of predictive data attribution

Data curation

Which data removals would improve performance on a target task?

Failure analysis

Which training examples would change a particular prediction if removed?

Data removal

How would excluding a set of records affect held-out behavior?

03 / 25

Example: attributing LLM predictions

Intervention

Remove a specified subset of training text

Measurement

Next-token loss after a held-out passage

Target

How that loss would change if the model had trained without this text

04 / 25

Parameterizing predictive data attribution

L(θ;w)=1n∑i=1nwi ℓ(xi;θ)\mathcal L(\theta;\mathbf w)=\frac{1}{n}\sum_{i=1}^{n} w_i\,\ell(x_i;\theta)
Original training set
w=1n\mathbf w=\mathbf1_n
Remove a subset S
w=1n−1S\mathbf w=\mathbf1_n-\mathbf 1_S

Each training example receives a continuous weight

05 / 25

Model output function

Data weights

w\mathbf w
→A\xrightarrow{\mathcal A}

Trained parameters

θ(w)\theta(\mathbf w)
→φq\xrightarrow{\varphi^{q}}

Final measurement

fq(w)f^{q}(\mathbf w)
fq(w)=(φq∘A)(w)f^{q}(\mathbf w)=(\varphi^{q}\circ\mathcal A)(\mathbf w)
06 / 25

Meta-gradients

∇wfq(w)∣w=1n=∇w[φq(A(w))]∣w=1n\left.\nabla_{\mathbf w}f^{q}(\mathbf w)\right|_{\mathbf w=\mathbf1_n}=\left.\nabla_{\mathbf w}\big[\varphi^{q}(\mathcal A(\mathbf w))\big]\right|_{\mathbf w=\mathbf1_n}

Gradients through the whole training procedure

07 / 25

MAGIC

MAGIC uses a first-order Taylor expansion

fq(w)≈fq(1n)+⟨∇wfq(1n), w−1n⟩f^{q}(\mathbf w)\approx f^{q}(\mathbf1_n)+\big\langle\nabla_{\mathbf w}f^{q}(\mathbf1_n),\,\mathbf w-\mathbf1_n\big\rangle

REPLAY computes the meta-gradients

08 / 25

MAGIC is state of the art for predictive data attribution

00.250.50.751EK-FAC0.23TRAK0.35MAGIC0.96

ResNet-9 / CIFAR-10 · random 1% deletions · mean Spearman correlation

09 / 25

Challenge

Computation scales linearly with the number of queries

k queries⟹k replaysk\text{ queries}\quad\Longrightarrow\quad k\text{ replays}

Can we attribute the whole test dataset with fewer replays?

10 / 25

Attribution as matrix estimation

Yqi=∂fq∂wi∣w=1n,Y∈Rk×nY_{qi}=\left.\frac{\partial f^{q}}{\partial w_i}\right|_{\mathbf w=\mathbf1_n},\qquad Y\in\mathbb R^{k\times n}
Row YqY_q: query qqColumn ii: training example ii

Estimate the influence matrix with B<kB<k measurements

11 / 25

Two distinct objectives

Reconstruction
min⁡Y^  ∥Y−Y^∥F2\min_{\widehat Y}\;\|Y-\widehat Y\|_F^2

Recover the influence scores

LDS
max⁡Y^  1k∑q=1kLDSq\max_{\widehat Y}\;\frac1k\sum_{q=1}^{k}\mathrm{LDS}_q

Predict the effects of data interventions

12 / 25

The chain rule factors the influence matrix

J=∂θ(w)∂w∣w=1nJ=\left.\frac{\partial\theta(\mathbf w)}{\partial\mathbf w}\right|_{\mathbf w=\mathbf1_n}

The Jacobian of the training procedure

V=[v1,…,vk],vq=∇θφq(θ(1n))V=[v_1,\ldots,v_k],\qquad v_q=\nabla_{\theta}\varphi^{q}\big(\theta(\mathbf1_n)\big)

Query gradients at the trained parameters

Y=V⊤JY=V^{\top}J
13 / 25

One replay can measure a combination of rows

One replay measures one query’s row

REPLAY⁡(vq)=vq⊤J=Yq\operatorname{REPLAY}(v_q)=v_q^{\top}J=Y_q

Combining query gradients gives a measured direction uu

u=REPLAY⁡(Vz)=(Vz)⊤J=z⊤Yu=\operatorname{REPLAY}(Vz)=(Vz)^{\top}J=z^{\top}Y
14 / 25

Orthogonal projection

Let U=span⁡{u1,…,uB}U=\operatorname{span}\{u_1,\ldots,u_B\}, with orthogonal ubu_b

Y^q:=YqΠU=∑b=1B⟨Yq,ub⟩∥ub∥22 ub\widehat Y_q:=Y_q\Pi_U=\sum_{b=1}^{B}\frac{\langle Y_q,u_b\rangle}{\|u_b\|_2^2}\,u_b

We know UU but still need ⟨Yq,ub⟩\langle Y_q,u_b\rangle for every query

15 / 25

Forward mode completes the projection

[⟨Y1,ub⟩⋮⟨Yk,ub⟩]=Yub⊤=V⊤(Jub⊤)\begin{bmatrix}\langle Y_1,u_b\rangle\\[-.15em]\vdots\\[-.15em]\langle Y_k,u_b\rangle\end{bmatrix}=Yu_b^{\top}=V^{\top}(Ju_b^{\top})
  • Forward mode computes Jub⊤Ju_b^{\top}
  • Multiplying by V⊤V^{\top} gives the inner products
16 / 25

An online viewpoint of PCA

∥Y−YΠU∥F2=∥Y∥F2−∥YΠU∥F2\|Y-Y\Pi_U\|_F^2=\|Y\|_F^2-\|Y\Pi_U\|_F^2

Maximizing captured energy minimizes reconstruction error

z=vmax⁡ ⁣(Y(I−ΠU)Y⊤)z=v_{\max}\!\left(Y(I-\Pi_U)Y^{\top}\right)
17 / 25

MAGE

YY⊤=V⊤JJ⊤V  ≈  V⊤VYY^{\top}=V^{\top}JJ^{\top}V\;\approx\;V^{\top}V

We use the query-gradient Gram matrix as a proxy

Probe its leading eigenvectors in order

18 / 25

From LDS to a projection objective

Start with LDS
LDSq=Spearman-ρS ⁣(−⟨Y^q,1S⟩, fq(1n−1S)−fq(1n))\mathrm{LDS}_q=\text{Spearman-}\rho_S\!\left(-\langle\widehat Y_q,\mathbf1_S\rangle,\ f^{q}(\mathbf1_n-\mathbf1_S)-f^{q}(\mathbf1_n)\right)
Use a Pearson surrogate
≈Pearson-ρS ⁣(−⟨Y^q,1S⟩, fq(1n−1S)−fq(1n))\approx\text{Pearson-}\rho_S\!\left(-\langle\widehat Y_q,\mathbf1_S\rangle,\ f^{q}(\mathbf1_n-\mathbf1_S)-f^{q}(\mathbf1_n)\right)
Approximate with MAGIC
≈Pearson-ρS ⁣(−⟨Y^q,1S⟩, −⟨Yq,1S⟩)\approx\text{Pearson-}\rho_S\!\left(-\langle\widehat Y_q,\mathbf1_S\rangle,\ -\langle Y_q,\mathbf1_S\rangle\right)
Uniform subset sampling
=Pearson-ρ(Y^q,Yq)=\text{Pearson-}\rho(\widehat Y_q,Y_q)
Neglect row centering
≈⟨YqΠU,Yq⟩∥YqΠU∥2 ∥Yq∥2\approx\frac{\langle Y_q\Pi_U,Y_q\rangle}{\|Y_q\Pi_U\|_2\,\|Y_q\|_2}
Use orthogonality
=∥YqΠU∥2∥Yq∥2=\frac{\|Y_q\Pi_U\|_2}{\|Y_q\|_2}
19 / 25

A normalized PCA objective for attribution

∑q=1k∥YqΠU∥22∥Yq∥22=∥Y~ΠU∥F2\sum_{q=1}^{k}\frac{\|Y_q\Pi_U\|_2^2}{\|Y_q\|_2^2}=\|\widetilde Y\Pi_U\|_F^2
Y~q=Yq∥Yq∥2\widetilde Y_q=\frac{Y_q}{\|Y_q\|_2}

The squared-correlation surrogate maximizes captured energy
of the row-normalized matrix

20 / 25

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

21 / 25

MAGE closely tracks the reconstruction oracle

Relative Frobenius error · lower is better
CIFAR-1000.250.50.751MAGE0.1614SPELL0.7845Random0.3177PCA0.1869First B0.8542Rank-B SVD0.1545TinyStories00.250.50.751MAGE0.5257SPELL0.7310Random0.6665PCA0.5429First B0.8777Rank-B SVD0.5075

k=500, B=100k=500,\ B=100 · mean ± SE · rank-B SVD requires the full matrix

22 / 25

SPELL leads on predictive data attribution

Linear datamodeling score · higher is better
CIFAR-1000.250.50.751MAGE0.3425SPELL0.4538Random0.3572PCA0.2337First B0.3405Spherical SVD0.5189MAGIC0.9087TinyStories00.250.50.751MAGE0.4341SPELL0.5451Random0.4543PCA0.3428First B0.4120Spherical SVD0.6130MAGIC0.9952

k=500, B=100k=500,\ B=100 · 1% deletions · mean ± SE · oracle SVD · full MAGIC uses B=500B=500

23 / 25

SPELL gains accuracy on queries with small meta-gradients

LDS by row-norm quartile of the influence matrix
CIFAR-1000.250.50.751SmallestQ2Q3LargestTinyStories00.250.50.751SmallestQ2Q3Largest
MAGESPELLMAGIC

1% deletions · CIFAR-10: k=1000,B=500k=1000, B=500 · TinyStories: k=600,B=300k=600, B=300

24 / 25

Thank you

Thank you