A Perfectly Valid Certificate for the Wrong Thing

My position on multilingual unlearning, with the maths behind it. Certify the fact and not the string, close the forget set across languages, and report the worst language instead of the average.

Published
Reading time
8 min
Section
RESEARCH
On this page
  1. The position, in three sentences
  2. The certificate is honest about the wrong object
  3. Closing the forget set
  4. Even with the right set, steps do not transfer
  5. A price you can read off one number
  6. What I think we should report
  7. Where I could be wrong
  8. The short version
  9. Footnotes#

Imagine a compliance officer holding a freshly printed certificate that says our model has forgotten Ramesh’s phone number. The certificate is mathematically correct. Every symbol has been checked twice by people who enjoy checking symbols. The model, meanwhile, is in the next room reciting the number to anyone who asks in Tamil.

Both things are true at the same time, and that is the entire problem. I work on unlearning and I work on multilingual models, and for a while I assumed those were two separate worries I could carry in two separate bags. They are one bag. This post is my position on why, followed by the maths that made me take that position, written as plainly as I can manage.

The position, in three sentences

  1. The object we should certify is a fact, and a training string is only one of the many ways the model was told that fact.
  2. A forget set that lives in one language is not a forget set. It has to be closed under translation before any unlearning method touches it.
  3. We should report the worst language, because an average over languages is how the smallest ones get quietly sacrificed.

Everything below is an attempt to make those three sentences earn their keep.

The certificate is honest about the wrong object

Recall the standard definition. We have a learning algorithm A\mathcal{A}, a dataset DD, and an example zz we want gone. A removal mechanism M\mathcal{M} is (ε,δ)(\varepsilon,\delta)-certified if its output is statistically close to A(D∖{z})\mathcal{A}(D\setminus\{z\}), which is what we would have gotten by retraining without zz.

Read that again slowly. The reference is retraining without zz. If the fact that zz expresses is also expressed by other training examples, then the reference model knows the fact too. The certificate can be perfect and the fact can survive, because the thing we compared against was never ignorant.

Let me make this concrete with a toy model. Say a fact ff is supported by mm training strings z1,…,zmz_1,\dots,z_m, written in different languages. Each string contributes some teaching weight wi≥0w_i \ge 0. A reasonable and very crude way to say how well the model knows the fact is a saturating curve,

k(W)  =  1−e−W,W=∑i=1mwi.k(W) \;=\; 1 - e^{-W}, \qquad W = \sum_{i=1}^{m} w_i .

More evidence means more knowledge, with diminishing returns, like revising for an exam. Now let the supporting strings be

LanguageWeight wwShare of total
English1.037%
Hindi1.037%
Tamil0.519%
Maithili0.27%

The total is W=2.7W = 2.7, so the model knows the fact at k=1−e−2.7≈0.93k = 1 - e^{-2.7} \approx 0.93. We receive a request to forget the English string. Retraining without it leaves W=1.7W = 1.7, which gives k≈0.82k \approx 0.82. A certified removal mechanism is allowed to land right there, and we can print the certificate with a straight face. The fact went from 93 percent known to 82 percent known, and everyone is invited to the party.

This is not a flaw in the definition. The definition does exactly what it says. It is a flaw in what we choose to call zz.

Closing the forget set

The fix is to forget the whole support of the fact, Zf={z1,…,zm}Z_f = \{z_1,\dots,z_m\}, and certify against retraining on D∖ZfD\setminus Z_f. The difficulty is that nobody hands us ZfZ_f. We have to find it, and finding it is a retrieval problem. In practice that means embedding the forget string with a multilingual sentence encoder and pulling out everything nearby, across languages, above some similarity threshold.

Retrieval is never perfect, and here is where the fairness problem walks in wearing a lab coat. Let rℓr_\ell be the recall of our retriever in language ℓ\ell, meaning the fraction of that language’s supporting weight we manage to find. Whatever we miss stays in the data. The weight left behind is

Wleft  =  ∑ℓ(1−rℓ) Wℓ,kleft  =  1−e−Wleft.W_{\text{left}} \;=\; \sum_{\ell} (1 - r_\ell)\, W_\ell , \qquad k_{\text{left}} \;=\; 1 - e^{-W_{\text{left}}} .

Retrievers are trained mostly on the languages that already have the most data, so rℓr_\ell falls as the language gets smaller. Put plausible numbers on the same toy example.

LanguageWeight WℓW_\ellRecall rℓr_\ellWeight left behind
English1.00.990.01
Hindi1.00.950.05
Tamil0.50.800.10
Maithili0.20.500.10

We left behind Wleft=0.26W_{\text{left}} = 0.26, so the model still knows the fact at k≈0.23k \approx 0.23. That is a big improvement over 0.82, and it is still not zero. More interesting is where the leftover lives. Tamil and Maithili supplied only 26 percent of the evidence and 77 percent of what remains. The languages with the least data are the ones that keep the secret, which is a strange way to reward them for being under-resourced.1

Even with the right set, steps do not transfer

Suppose we found ZfZ_f perfectly. We still have to change the weights, and now a second geometric problem appears. The gradient of the log probability of the fact in language ℓ\ell is gℓg_\ell. A step Δ\Delta changes that log probability by about gℓ⊤Δg_\ell^{\top}\Delta. I wrote about the first order version of this in an earlier post, where an English-only step helps Hindi exactly as much as the two gradients lean the same way.

Here I want to add the thing that post left out, which is the price. Forcing every language to forget by the same amount cc means solving

G⊤Δ=−c 1,G=[ g1    ⋯    gL ].G^{\top}\Delta = -c\,\mathbf{1}, \qquad G = [\,g_1 \;\; \cdots \;\; g_L\,].

There are infinitely many solutions, and last time I picked the shortest one in ordinary distance. Ordinary distance in weight space is a poor measure of damage, because some directions barely matter to the model and some directions break it. A better ruler is the Fisher information matrix FF of the data we want to keep, since 12Δ⊤FΔ\tfrac12\Delta^{\top}F\Delta approximates how far the model’s outputs on that data move, in KL divergence. So we ask for the cheapest step in that ruler,

min⁡Δ  12 Δ⊤FΔsubject toG⊤Δ=−c 1.\min_{\Delta}\; \tfrac12\,\Delta^{\top}F\Delta \quad\text{subject to}\quad G^{\top}\Delta = -c\,\mathbf{1}.

Introduce a Lagrange multiplier λ\lambda. Stationarity gives FΔ=GλF\Delta = G\lambda, so Δ=F−1Gλ\Delta = F^{-1}G\lambda. Plugging into the constraint gives G⊤F−1G λ=−c 1G^{\top}F^{-1}G\,\lambda = -c\,\mathbf{1}. Name the small L×LL\times L matrix M=G⊤F−1GM = G^{\top}F^{-1}G, and then

Δ⋆=−c F−1G M−11,cost=12 Δ⋆⊤FΔ⋆=c22  1⊤M−11.\Delta^{\star} = -c\,F^{-1}G\,M^{-1}\mathbf{1}, \qquad \text{cost} = \tfrac12\,\Delta^{\star\top}F\Delta^{\star} = \tfrac{c^{2}}{2}\;\mathbf{1}^{\top}M^{-1}\mathbf{1}.

That last expression is the price of forgetting in every language at once. Everything interesting sits inside MM, which is just the table of how much each pair of languages agrees, measured in the Fisher ruler.

A price you can read off one number

Rescale each language so its own gradient has unit length in the Fisher ruler, so that MM has ones on the diagonal. Suppose every pair of languages agrees with the same correlation ρ\rho. Then M=(1−ρ)I+ρ 11⊤M = (1-\rho)I + \rho\,\mathbf{1}\mathbf{1}^{\top}, and the inverse is easy enough to do on the back of a bus ticket,

1⊤M−11  =  L1+(L−1)ρ,cost  =  c22⋅L1+(L−1)ρ.\mathbf{1}^{\top}M^{-1}\mathbf{1} \;=\; \frac{L}{1 + (L-1)\rho}, \qquad \text{cost} \;=\; \frac{c^{2}}{2}\cdot\frac{L}{1+(L-1)\rho}.

Look at the two ends. If ρ=1\rho = 1, all languages share one gradient, the cost is c2/2c^2/2 whatever LL is, and forgetting in ten languages costs the same as forgetting in one. If ρ=0\rho = 0, the languages share nothing, the cost is Lc2/2Lc^2/2, and every language is a separate bill. Real models sit between those ends, and the formula says how far.

Take L=4L = 4 and ρ=0.2\rho = 0.2. The honest all-language step costs c22⋅41.6=1.25 c2\tfrac{c^2}{2}\cdot\tfrac{4}{1.6} = 1.25\,c^2. The English-only step costs 0.5 c20.5\,c^2, but it only moves each other language by ρc=0.2c\rho c = 0.2c. We asked for cc and delivered a fifth of it. The all-language step costs 2.5 times as much, and that ratio is the real price of treating the other three languages as people.

fisher_forget.py python
import numpy as np

def fisher_forget_step(grads, fisher, c=1.0, ridge=1e-6):
    """Cheapest step, in the Fisher ruler, that lowers the fact's log prob by c in every language.

    grads  is an (L, d) array with one gradient of log p(fact) per language.
    fisher is a (d, d) Fisher matrix of the data we want to keep.
    Returns the step and its cost.
    """
    G = grads.T                                          # (d, L)
    FinvG = np.linalg.solve(fisher + ridge * np.eye(fisher.shape[0]), G)
    M = G.T @ FinvG + ridge * np.eye(G.shape[1])         # (L, L) agreement between languages
    w = np.linalg.solve(M, np.ones(G.shape[1]))
    step = -c * (FinvG @ w)
    cost = 0.5 * c**2 * np.ones(G.shape[1]) @ w          # (c^2 / 2) * 1' M^-1 1
    return step, cost

In a real network FF is far too large to store, so you would use a block diagonal or Kronecker approximation, and then the cost formula is only as good as that approximation.

What I think we should report

Putting the two halves together, a forgetting result has two numbers worth printing and most papers print neither.

  • Residual knowledge in the worst language, max⁡ℓkℓ\max_\ell k_\ell, measured with prompts written by native speakers and not machine translated from English. A machine translated prompt tests the translator.
  • Retrieval recall per language for whatever process built the forget set, because it caps how much any method can forget.

I would also stop averaging. An average over fifty languages in which forty-nine are clean and one holds the whole secret reads as a 98 percent success, which is a very good score for a model that is still leaking.

Where I could be wrong

Four things I would not bet my degree on.

  1. The saturating curve k(W)=1−e−Wk(W)=1-e^{-W} is a cartoon. Real knowledge is not a sum of independent contributions, and some facts are learned from one great example and not from many mediocre ones.
  2. The equal correlation assumption flatters the maths. Real language pairs cluster, so Hindi and Urdu agree strongly with each other and weakly with Tamil, and the price depends on the whole matrix and not on one ρ\rho.
  3. The first order step ignores curvature, so large cc needs several small steps with GG recomputed each time.
  4. Fisher distance on the retain data measures damage to what we kept. It says nothing about damage to languages we never sampled, which is the same blind spot again in a different place.

None of these change the position. They change how confident I am about the numbers.

The short version

A certificate tells you that you removed the string you pointed at. It cannot tell you that you pointed at the right thing. Pointing at the right thing means finding the fact in every language it was written in, and finding it worst in the languages with the least data, which are also the languages where someone is least likely to check.

My father would put it more briefly. Before you say you have thrown out the old newspapers, check the storeroom in the other house.

Footnotes#

  1. The recall numbers here are invented to be plausible and easy to add up. The direction of the effect, with recall falling as resources fall, is the part I would defend. The sizes are for the arithmetic. ↩

Where to next?

Writing / Research