Lesson 17 / 27

Catastrophic Forgetting and Replay

Training on a new task can erase an old one; mixing old examples back in helps.

New skill in, old skill out

When a network is trained only on new data, weights move to fit it and can overwrite what supported earlier skills: catastrophic forgetting. In LLMs it shows up as lost general ability, worse formatting elsewhere or forgotten languages after narrow tuning. Mitigations: a low learning rate and few epochs, LoRA (the base stays frozen), replay (mix a sample of general or earlier data into the training set), and an evaluation that includes the old skills, not only the new one.

What fine-tuning can quietly break

Specialising a model can erase old skills, add confident wrong facts and weaken safety behaviour.

Three ideas: forgetting, facts, safety.
Figure 5.1 — Forgetting, facts and safety.

Forgetting and replay in a small network, run

I ran this with Python, numpy 2.5.3 and scikit-learn 1.9.1, with fixed random seeds. It trains small classical models, not a language model: the mechanics (gradient descent, learning rate, overfitting, forgetting, low-rank updates) are the same ideas that apply to fine-tuning an LLM, but the numbers are not LLM results. A small neural network learns task A (accuracy 0.965 on A). After fine-tuning only on task B it scores 0.961 on B but 0.041 on A, so A is forgotten. Fine-tuning on B plus 500 replayed A examples keeps both (A 0.961, B 0.956). The two tasks here are a constructed regime-flag toy, not an LLM.

import numpy as np
from sklearn.neural_network import MLPClassifier

rng = np.random.default_rng(3)
w = np.array([1.5, -1.0, 1.0, 0.5, -1.5, 0, 0, 0])
def task(flag, n=2000):                              # same inputs, but the flag feature says which rule applies
    X = rng.normal(size=(n, 8)); X[:, 7] = flag
    s = X @ w; y = ((s if flag == 0 else -s) + rng.normal(size=n) * 0.3 > 0).astype(int)
    return X, y
XA, yA = task(0); XB, yB = task(1); XAt, yAt = task(0); XBt, yBt = task(1)

def run(replay):
    m = MLPClassifier(hidden_layer_sizes=(32,), learning_rate_init=0.01, random_state=0)
    for _ in range(40): m.partial_fit(XA, yA, classes=[0, 1])
    before = (m.score(XAt, yAt), m.score(XBt, yBt))
    X, y = (np.vstack([XB, XA[:500]]), np.concatenate([yB, yA[:500]])) if replay else (XB, yB)
    for _ in range(80): m.partial_fit(X, y)
    return before, (m.score(XAt, yAt), m.score(XBt, yBt))

b, a = run(replay=False)
print(f"trained on task A      : accuracy on A {b[0]:.3f} | on B {b[1]:.3f}")
print(f"then fine-tuned on B   : accuracy on A {a[0]:.3f} | on B {a[1]:.3f}   <- A forgotten")
b, a = run(replay=True)
print(f"fine-tuned on B + 500 A: accuracy on A {a[0]:.3f} | on B {a[1]:.3f}   <- both kept")

Output:

trained on task A      : accuracy on A 0.965 | on B 0.045
then fine-tuned on B   : accuracy on A 0.041 | on B 0.961   <- A forgotten
fine-tuned on B + 500 A: accuracy on A 0.961 | on B 0.956   <- both kept

Keep a regression suite

Keep a small set of old-skill prompts and rerun it after every tuning run to catch forgetting early.

Quick check: Which helps reduce catastrophic forgetting?

  • Removing the validation set
  • Training longer on only the new task
  • Replaying a sample of earlier data, a low learning rate or LoRA
  • Increasing the learning rate a lot
Answer

Replaying a sample of earlier data, a low learning rate or LoRA — Keep old behaviour in the loop and in the evaluation.