Post-training and Alignment

Catastrophic forgetting

Training a model on something new overwrites what it already knew, and the collapse is far more complete than most people expect.

On this page 8
  1. The short answer
  2. The analogy you have already lived
  3. Why this matters
  4. How it works
  5. What actually helps
  6. Where you have already seen this
  7. Remember this
  8. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

Teach a model something new and it overwrites what it already knew, often completely.

The analogy you have already lived

You spend three weeks cramming for a chemistry exam. You do well.

A month later somebody asks you a physics question from the term before. You cannot remember the formula. You knew it perfectly six weeks ago.

Nothing physically removed it. The revision for chemistry sat on top of the physics and pushed it out of reach.

Models do the same thing, and they do it far faster and far more completely than people do.

Why this matters

Fine-tuning is how models are specialised. A hospital fine-tunes on medical notes. A bank fine-tunes on its own documents.

The intention is "keep everything you know, and add this". What actually happens is closer to "replace what you know with this".

The model gets better at the new thing and quietly loses abilities nobody was testing for. Safety behaviour, other languages, arithmetic, following instructions — all can degrade without a single warning in the training logs.

How it works

Every weight in the model is shared by every ability. There is no shelf holding "French" and a separate shelf holding "arithmetic".

Training pushes the weights toward the new task. It moves the same numbers that were carrying the old ones.

   before                     after fine-tuning on task B

   task A ability  ####       task A ability  .
   task B ability  .          task B ability  #####

   the model did not "add" B. It moved to B.

The word "catastrophic" is not decoration. In the measurement below a model scores 87% on its first task. After training on a second, it scores zero.

What actually helps

Mix the old data back in. Keep some of the original examples in every batch. This is called replay, and it is the only reliable fix.

Train less. A smaller learning rate and fewer steps means less movement. Less damage, and also less learning of the new task.

Freeze most of the model. Train only a small add-on. Helps, but only if you freeze the parts that were carrying the old ability.

The measurement below shows replay working and the other two mostly not. That ordering surprises people, and it is worth trusting over intuition.

Where you have already seen this

  • Forgetting last term's subject after cramming for this term's.
  • Learning a new phone's keyboard and fumbling on your old one.
  • Recording over a cassette or a video tape.

Remember this

  • Fine-tuning on something new overwrites old abilities rather than adding to them.
  • The collapse can be total, and it does not show up in your training loss.
  • Mixing old data back in is the fix that reliably works.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Runs on a CPU in about fifteen seconds.

Watching a skill disappear

Two tasks with disjoint label spaces, so one model could do both, and a marker feature telling it which task an input belongs to. There is no contradiction here — only interference.

forgetting.py
import copy
import torch
import torch.nn as nn
import torch.nn.functional as F

# Two tasks, one shared 10-class head, with DISJOINT label spaces so a single
# model could do both. A marker feature says which task the input belongs to.
# TASK A: bin the first 20 features -> classes 0-4
# TASK B: bin the last  20 features -> classes 5-9
D, C = 41, 10


def task(which, n, gen):
    x = torch.randn(n, D, generator=gen)
    x[:, -1] = 0.0 if which == "A" else 1.0            # the task marker
    half = x[:, :20] if which == "A" else x[:, 20:40]
    b = (half.sum(1) / 3 + 2.5).clamp(0, 5 - 1e-3).long()
    return x, b + (0 if which == "A" else 5)


def net(seed=0):
    torch.manual_seed(seed)
    return nn.Sequential(nn.Linear(D, 128), nn.ReLU(),
                         nn.Linear(128, 128), nn.ReLU(), nn.Linear(128, C))


gen = torch.Generator().manual_seed(0)
ax, ay = task("A", 4000, gen)
bx, by = task("B", 4000, gen)
tax, tay = task("A", 2000, torch.Generator().manual_seed(5))
tbx, tby = task("B", 2000, torch.Generator().manual_seed(6))


def fit(model, x, y, steps, lr=3e-3, freeze_body=False):
    params = list(model[-1].parameters()) if freeze_body else list(model.parameters())
    if freeze_body:
        for p in model.parameters():
            p.requires_grad_(False)
        for p in model[-1].parameters():
            p.requires_grad_(True)
    opt = torch.optim.AdamW(params, lr=lr)
    for _ in range(steps):
        loss = F.cross_entropy(model(x), y)
        opt.zero_grad(); loss.backward(); opt.step()
    return model


def acc(model, x, y):
    with torch.no_grad():
        return (model(x).argmax(-1) == y).float().mean().item()


base = fit(net(), ax, ay, 600)
print(f"after training on task A only:  A {acc(base, tax, tay):.3f}   B {acc(base, tbx, tby):.3f}")

print(f"\n{'then train on task B with...':<34} {'A':>7} {'B':>7}")
recipes = {
    "task B only, lr 3e-3": dict(x=bx, y=by, lr=3e-3),
    "task B only, lr 3e-4": dict(x=bx, y=by, lr=3e-4),
    "task B + 5% replay of A": dict(x=torch.cat([bx, ax[:200]]),
                                    y=torch.cat([by, ay[:200]]), lr=3e-3),
    "task B + 50% replay of A": dict(x=torch.cat([bx, ax[:2000]]),
                                     y=torch.cat([by, ay[:2000]]), lr=3e-3),
    "task B, body frozen": dict(x=bx, y=by, lr=3e-3, freeze_body=True),
}
for name, kw in recipes.items():
    m = fit(copy.deepcopy(base), steps=600, **kw)
    print(f"{name:<34} {acc(m, tax, tay):>7.3f} {acc(m, tbx, tby):>7.3f}")
Output
after training on task A only:  A 0.871   B 0.000

then train on task B with...             A       B
task B only, lr 3e-3                 0.000   0.899
task B only, lr 3e-4                 0.013   0.340
task B + 5% replay of A              0.101   0.853
task B + 50% replay of A             0.780   0.887
task B, body frozen                  0.000   0.345

Written against PyTorch 2.5.1, CPU, all seeds fixed — reproducible on this build.

Five rows, four lessons

Row 1 is the whole phenomenon: 0.871 → 0.000. Not degraded. Gone. And nothing in the task-B loss curve gives any hint, because the task-B loss was falling beautifully the entire time.

Lowering the learning rate cost more than it saved. Task A recovered from 0.000 to 0.013 — statistical noise — while task B fell from 0.899 to 0.340. Training less does mean forgetting less, and it also means learning less. This is not a good trade at this ratio.

5% replay bought 0.101 on A for almost nothing on B. A small but real effect from 200 old examples mixed into 4,000 new ones.

50% replay bought 0.780 on A while keeping 0.887 on B. Both tasks, one model, no clever algorithm. Replay is the fix. The cost is that you must still have the old data, which in practice is exactly the problem — nobody ships the pretraining corpus with the model.

Freezing the body did not help at all. A stayed at 0.000. The frozen-body run only trained the final layer, and that final layer is the shared output head carrying task A's classes. Freezing the wrong part protects nothing. If you use adapters for this, keep a separate head per task, or accept that the shared head is a single point of forgetting.

What this looks like on a real LLM

The damage is rarely in the thing you are measuring. It shows up in abilities you stopped testing.

A practical protocol:

python
# before fine-tuning, capture a baseline on abilities you are NOT training
BASELINE_SUITE = ["gsm8k_subset", "safety_refusals", "multilingual_qa",
                  "instruction_format", "json_validity"]
# after fine-tuning, re-run all of them and diff

No output block — the suite is yours to build. The rule is what matters: an evaluation you did not run before fine-tuning cannot tell you what fine-tuning cost.

Techniques that work on real models, in rough order of how much they help:

  • Replay. Mix 5–30% of general instruction data into the domain data. This is what every serious continued-pretraining recipe does.
  • Low learning rate, few epochs. One to three epochs at 1e-5 for full fine-tuning. More epochs is where most damage happens.
  • Adapters. LoRA forgets measurably less than full fine-tuning, because the frozen base is literally unchanged and the adapter can be removed.
  • Model merging. Average the fine-tuned weights back toward the original.
  • Regularisation toward the original weights. An L2 penalty on the distance from the starting point, or the Fisher-weighted version described below.

Common mistakes

Only evaluating the target task. The single most common and most expensive mistake in this entire lesson.

Assuming a small dataset is safe. A few hundred narrow examples at a high learning rate can wreck instruction-following. Volume is not what determines damage; update magnitude is.

Fine-tuning away safety behaviour by accident. Published work shows safety alignment can be removed by fine-tuning on a handful of benign examples. If you ship a fine-tuned model, re-test its refusals.

Reusing a base model's evaluation numbers for your fine-tuned model. They no longer apply.

Fine-tuning on top of a fine-tune, repeatedly. Each round compounds the drift. Go back to the base model and retrain from a combined dataset instead.

Try it yourself

Change 50% replay to use ax[:400] (10%) and then ax[:800] (20%). Plot task-A accuracy against replay fraction. The curve is steep at first and then flattens — and where it flattens is the replay budget you actually need.

What to learn next

Researcher — Mathematics and papers.

The mechanism

Catastrophic interference was identified in connectionist networks by McCloskey and Cohen, 1989 and Ratcliff, 1990. The cause is representational overlap: gradient descent on task B moves shared parameters without any term in the loss preserving task A's function.

Formally, after training on A the parameters sit at $\theta_A^$, a minimiser of $\mathcal{L}_A$. Training on B follows $-\nabla\mathcal{L}_B$, which has no component constraining $\mathcal{L}_A$. Since $\theta_A^$ lies on a low-dimensional manifold of solutions for A, almost any direction leaves it.

The corollary that matters practically: there generally exists a $\theta$ good at both — the 50%-replay row demonstrates one — but sequential gradient descent will not find it, because nothing in the sequential objective points there.

What actually preserves knowledge

Replay / rehearsal. Mix samples from the old distribution into the new batches. This is not a heuristic: it makes the objective $\mathcal{L}_A + \mathcal{L}_B$ again, which is the joint problem whose solution you want. Everything else is an approximation of it.

Elastic Weight Consolidation (Kirkpatrick et al., 2017) approximates the old task's loss with a quadratic penalty weighted by the diagonal Fisher information:

$$ \mathcal{L}(\theta) = \mathcal{L}_B(\theta) + \frac{\lambda}{2}\sum_i F_i\,(\theta_i - \theta^*_{A,i})^2 $$

$F_i$ estimates how much task A's log-likelihood depends on parameter $i$, so parameters that mattered to A are held rigid and the rest are free. It works, and it needs a Fisher estimate over the old data — which means you still need the old data, or a saved Fisher diagonal.

Related: Synaptic Intelligence (Zenke et al., 2017) accumulates importance online; Learning without Forgetting (Li and Hoiem, 2017) distils from the pre-fine-tuning model's outputs on the new data, which uniquely requires no old data at all.

Measured on real LLMs

  • Luo et al., 2024 (An Empirical Study of Catastrophic Forgetting in LLMs During Continual Fine-tuning) find forgetting of domain knowledge, reasoning and reading comprehension worsens monotonically with model size from 1B to 7B, and that decoder-only models forget more than encoder–decoder ones.
  • Qi et al., 2024 (Fine-tuning Aligned Language Models Compromises Safety, Even When Users Do Not Intend To!, ICLR 2024) is the alarming result: safety alignment of GPT-3.5 Turbo was removed by fine-tuning on 10 adversarially designed examples for under $0.20 through the public API; and fine-tuning on entirely benign datasets also measurably degraded safety.
  • Biderman et al., 2024 (LoRA Learns Less and Forgets Less) quantify the adapter trade-off: LoRA better preserves out-of-domain performance and learns less in-domain, with the effect explained by the low rank of its updates.
  • Ibrahim et al., 2024 (Simple and Scalable Strategies to Continually Pre-train LLMs) show that learning-rate re-warming plus re-decaying, combined with replaying a small fraction (as little as 5%) of the previous distribution, matches full retraining on the union of datasets at a fraction of the cost.

That last result is the most actionable in this list: LR re-warming plus ~5% replay is the current default recipe for continued pretraining.

Task vectors and merging as a mitigation

Defining the task vector $\tau = \theta_{\text{ft}} - \theta_{\text{base}}$ (Ilharco et al., 2023), a scaled interpolation

$$ \theta = \theta_{\text{base}} + \lambda\tau, \qquad \lambda \in (0, 1) $$

trades new-task performance against retained general ability, and $\lambda \approx 0.5$–$0.8$ frequently keeps most of the fine-tuning gain while recovering much of what was lost. Wortsman et al., 2022 (Robust fine-tuning of zero-shot models, WiSE-FT) established this for CLIP, reporting improved robustness and accuracy from the interpolation. See merging model weights.

The theoretical floor

Knoblauch et al., 2020 (Optimal Continual Learning has Perfect Memory and is NP-hard) prove that optimal continual learning requires perfect memory of all previous tasks, and that the resulting set-intersection problem is NP-hard. This is why every practical method is a trade-off rather than a fix, and why replay — which is literally partial memory — remains the strongest approach.

Papers

  • McCloskey and Cohen, Catastrophic Interference in Connectionist Networks, 1989 — Psychology of Learning and Motivation 24.
  • French, Catastrophic forgetting in connectionist networks, 1999 — Trends in Cognitive Sciences 3(4).
  • Kirkpatrick et al., Overcoming catastrophic forgetting in neural networks (EWC), PNAS 2017 — arxiv.org/abs/1612.00796
  • Li and Hoiem, Learning without Forgetting, TPAMI 2017 — arxiv.org/abs/1606.09282
  • Knoblauch et al., Optimal Continual Learning has Perfect Memory and is NP-hard, ICML 2020 — arxiv.org/abs/2006.05188
  • Wortsman et al., Robust fine-tuning of zero-shot models (WiSE-FT), CVPR 2022 — arxiv.org/abs/2109.01903
  • Ilharco et al., Editing Models with Task Arithmetic, ICLR 2023 — arxiv.org/abs/2212.04089
  • Qi et al., Fine-tuning Aligned Language Models Compromises Safety, ICLR 2024 — arxiv.org/abs/2310.03693
  • Luo et al., An Empirical Study of Catastrophic Forgetting in LLMs During Continual Fine-tuning, 2024 — arxiv.org/abs/2308.08747
  • Ibrahim et al., Simple and Scalable Strategies to Continually Pre-train LLMs, TMLR 2024 — arxiv.org/abs/2403.08763
  • Biderman et al., LoRA Learns Less and Forgets Less, TMLR 2024 — arxiv.org/abs/2405.09673

What to learn next