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.
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.