Lesson 100: BERT Fine-Tuning Strategies

USAAIO Lesson 100, from Phase 3. It compares feature extraction with full fine-tuning, then covers layer-wise learning rates, two-stage domain adaptation, and catastrophic forgetting with elastic weight consolidation (EWC). A TinyBERT-like toy model - d=64, 4 layers, about 208K parameters - is used to trace all the strategies, and the EWC loss formula, the diagonal of the Fisher information, and the layer-wise learning-rate decay were verified analytically with numpy. The lesson runs to 27 slides.

Subject: Machine Learning · 53 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. BERT Fine-Tuning Four Strategies

Title

USAAIO · Lesson 100 · Phase 3

From frozen backbone to full-model updates: feature extraction, full fine-tuning, layer-wise learning rates, domain adaptation, and catastrophic forgetting prevention with EWC — all grounded in why each knob exists.

2. By the end of this lesson you can

Objectives

  1. Implement feature extraction by freezing BERT weights and training only a classification head, and state the data-size regime where it wins
  2. Implement full fine-tuning and explain why a low uniform LR (e.g. 2e-5) is critical
  3. Construct a layer-wise LR schedule (10× decay per group going down) in PyTorch using param_groups
  4. Describe two-stage domain adaptation (domain MLM → task fine-tune) and the BioBERT evidence for it
  5. Implement elastic weight consolidation (EWC) to resist catastrophic forgetting: compute the Fisher diagonal, form the penalty, and add it to the task loss

3. What survived from Eigenvalues & Eigenvectors?

Warm-up

Discussion prompt

Before we open Lesson 100: BERT Fine-Tuning Strategies: without looking back, what was the main idea of Eigenvalues & Eigenvectors, and what could you do by the end of it that you could not do before?

Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.

Answer:

the eigen-equation Av=λv, the characteristic polynomial, the spectral theorem for symmetric matrices, power iteration for the dominant eigenvector, and the connection to PCA — build power iteration and PCA from scratch and verify against sklearn.

4. Feature extraction — freeze and reuse

Section

Part 1 of 4

5. What feature extraction means

Concept

Feature extraction treats the pretrained BERT stack as a fixed representation function. Only the task head — a single linear layer mapping the [CLS] embedding to class logits — is updated. The rest of the 110 M parameters are frozen.

componentrequires_gradparams updated
BERT embeddingsFalse0
BERT transformer layers (×12)False0
Classification head (Linear)Trued_model × n_classes + n_classes

For BERT-base (d_model=768, 2-class): head has 768×2 + 2 = 1,538 trainable params out of 110 M. This is why it runs on tiny datasets without overfitting.

6. Fill in: requires_grad for What feature extraction means

Comparison

Comparison matrix

From What feature extraction means: refill the requires_grad column from what you know. The rest of the table is as it appeared.

componentrequires_gradparams updated
BERT embeddingsFalse0
BERT transformer layers (×12)False0
Classification head (Linear)Trued_model × n_classes + n_classes

7. What has to happen first: Feature extraction in PyTorch

Ranking

Put in order

Put the moves of Feature extraction in PyTorch into the order they have to happen.

  1. Freeze all BERT parameters with a requires_grad=False loop
  2. Pass only requires_grad=True params to the optimizer
  3. Train with a standard classification loop; the BERT layers never update

Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. Frozen params skip .grad allocation and optimizer updates; only the head receives gradients.

8. Feature extraction in PyTorch

Worked example

Freeze all BERT parameters with a requires_grad=False loop

Why: Frozen params skip .grad allocation and optimizer updates; only the head receives gradients. This is the canonical pattern for transfer learning with a frozen backbone.

import torch, torch.nn as nn

class BERTClassifier(nn.Module):
    def __init__(self, bert, d=768, n_cls=2):
        super().__init__()
        self.bert = bert
        self.head = nn.Linear(d, n_cls)

    def forward(self, input_ids, attention_mask):
        # [CLS] is position 0
        out = self.bert(input_ids, attention_mask)
        cls = out.last_hidden_state[:, 0]  # (B, 768)
        return self.head(cls)              # (B, 2)

# --- freeze BERT, unfreeze head ---
model = BERTClassifier(bert)
for name, p in model.named_parameters():
    if 'head' not in name:
        p.requires_grad = False

trainable = sum(p.numel() for p in model.parameters()
                if p.requires_grad)
print(f'Trainable params: {trainable}')   # 1538

Pass only requires_grad=True params to the optimizer

Why: Without this filter, AdamW allocates momentum buffers for all 110 M params even though only the head has gradients — wasting memory and slowing each step.

Train with a standard classification loop; the BERT layers never update

Why: The CLS embedding is effectively a fixed feature vector extracted by BERT. Only the linear boundary between classes is learned — identical to training a logistic regression on top of fixed embeddings.

epochtrain losstest acc %
10.693152.5
50.620362.5
100.567172.5
150.528477.5
200.501280.0

9. What each one costs: Feature extraction in PyTorch

Trade off

Comparison matrix

From Feature extraction in PyTorch: every row here is a choice with a cost. Fill the train loss column, then say which row you would actually pick and what you give up for it.

epochtrain losstest acc %
10.693152.5
50.620362.5
100.567172.5
150.528477.5
200.501280.0

10. Full fine-tuning and the LR danger zone

Section

Part 2 of 4

11. Full fine-tuning: every parameter moves

Concept

Full fine-tuning unlocks all 110 M parameters simultaneously. The optimizer updates every weight: embeddings, all transformer layers, and the head. This typically achieves higher task accuracy than feature extraction when you have >10 k labeled examples.

strategytrainablewhen to use
Feature extraction~0.001% (head only)<1 k labeled examples
Full fine-tuning100%>10 k labeled examples
Layer-wise LR100% (different rates)Medium data; best accuracy

The critical hyper-parameter is learning rate. BERT was pretrained at ~1e-4; fine-tuning at the same rate destroys the pretrained representations in a few steps. The Devlin et al. paper recommends 2e-5 to 5e-5 — low enough for gentle gradient steps.

12. Something is wrong here: using a high uniform learning rate for fine-tuning

Anomaly

Predict first

A student writes this, and it looks reasonable:

Student sets lr=1e-3 (the "standard" Adam default) and trains all BERT layers for 5 epochs.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: This is the default that works well for training from scratch.

Fine-tuning requires a LR ≈ 50–500× smaller than scratch training to protect pretrained weights.

Why: This is the default that works well for training from scratch.

13. Trap: using a high uniform learning rate for fine-tuning

Trap

The trap

Student sets lr=1e-3 (the "standard" Adam default) and trains all BERT layers for 5 epochs.

optimizer = AdamW(model.parameters(), lr=1e-3)

Why: This is the default that works well for training from scratch.

Outcome: validation loss spikes after epoch 1, then diverges. Pretrained representations are wiped out — the model is effectively re-initialized with random gradient noise.

The fix

Fine-tuning requires a LR ≈ 50–500× smaller than scratch training to protect pretrained weights.

optimizer = AdamW(model.parameters(), lr=2e-5)

Why: 2e-5 produces gentle gradient steps that nudge each weight toward the task without erasing the language knowledge encoded in the pretrained weights.

Rule: for full fine-tuning of BERT-scale models, use 2e-5 to 5e-5. For very small datasets, lean toward 2e-5 or lower.

14. Break it on purpose: using a high uniform learning rate for…

Break the constraint

Discussion prompt

The rule this trap just fixed:

Fine-tuning requires a LR ≈ 50–500× smaller than scratch training to protect pretrained weights.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

This is the default that works well for training from scratch.

15. Layer-wise LR — protect what matters most

Section

Part 3 of 4

16. Why different layers need different learning rates

Concept

In a pretrained transformer, lower layers capture general syntax and morphology; upper layers capture task-relevant semantics. A uniform LR treats them identically — but the lower layers need smaller updates to preserve generalizable features.

\[ \eta_k = \eta_{\text{base}} \cdot \alpha^{L-k} \quad k=0,\ldots,L \text{ (layer index from bottom)} \]

Typical: alpha=0.1 (10× decay per layer), eta_base=2e-5 for the head. The bottom embedding layer ends up at 2e-5 × 0.1^4 = 2e-9 — barely moves.

17. Break it if you can: Why different layers need different learning rates

Counterexample

Discussion prompt

Typical: alpha=0.1 (10× decay per layer), eta_base=2e-5 for the head. The bottom embedding layer ends up at 2e-5 × 0.1^4 = 2e-9 — barely moves.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

18. Predict the next row: Layer-wise LR with param_groups

Pattern

Predict first

The table runs: embeddings | 2.00e-09 | 1 / 10,000 · layer[0-2] | 2.00e-08 | 1 / 1,000 · layer[3-5] | 2.00e-07 | 1 / 100 · layer[6-8] | 2.00e-06 | 1 / 10 · layer[9-11] | 2.00e-05 | 1 / 1

In Layer-wise LR with param_groups, given the rows so far: what is the next one — the row where group is head?

Correct: head | 2.00e-05 | 1 / 1

grouplrratio to head
embeddings2.00e-091 / 10,000
layer[0-2]2.00e-081 / 1,000
layer[3-5]2.00e-071 / 100
layer[6-8]2.00e-061 / 10
layer[9-11]2.00e-051 / 1
head2.00e-051 / 1

Why: The relationship between the columns, not the individual numbers, is what generates the next row. AdamW (like all PyTorch optimizers) accepts a list of {"params": ..., "lr": ...} dicts — one per group.

19. Layer-wise LR with param_groups

Worked example

Partition parameters into groups by depth, assign a decreasing LR to each

Why: AdamW (like all PyTorch optimizers) accepts a list of {"params": ..., "lr": ...} dicts — one per group. Each group maintains separate step-size and momentum state.

base_lr = 2e-5
alpha   = 0.1   # 10x decay per group going down

# bert.encoder.layer is a ModuleList of 12 layers (index 0=bottom)
param_groups = [
    {"params": bert.embeddings.parameters(),
     "lr": base_lr * alpha**4},     # 2e-09
    {"params": bert.encoder.layer[0:3].parameters(),
     "lr": base_lr * alpha**3},     # 2e-08
    {"params": bert.encoder.layer[3:6].parameters(),
     "lr": base_lr * alpha**2},     # 2e-07
    {"params": bert.encoder.layer[6:9].parameters(),
     "lr": base_lr * alpha**1},     # 2e-06
    {"params": bert.encoder.layer[9:12].parameters(),
     "lr": base_lr * alpha**0},     # 2e-05
    {"params": head.parameters(),
     "lr": base_lr},                # 2e-05
]
optimizer = AdamW(param_groups)

Verify the schedule printed before training starts

Why: A common bug: using the same param object in multiple groups — PyTorch warns but still trains, causing silent double-counting. Print and cross-check.

grouplrratio to head
embeddings2.00e-091 / 10,000
layer[0-2]2.00e-081 / 1,000
layer[3-5]2.00e-071 / 100
layer[6-8]2.00e-061 / 10
layer[9-11]2.00e-051 / 1
head2.00e-051 / 1

20. Inspect it line by line: Layer-wise LR with param_groups

Error analysis

Annotate

Walk the callouts on Layer-wise LR with param_groups. Each one is a place this is easy to get subtly wrong.

  • AdamW (like all PyTorch optimizers) accepts a list of {"params": ..., "lr": ...} dicts — one per group. Each group maintains separate step-size and momentum state.
  • A common bug: using the same param object in multiple groups — PyTorch warns but still trains, causing silent double-counting. Print and cross-check.

21. Domain adaptation then task fine-tune

Section

Part 4a of 4

22. Two-stage domain adaptation

Concept

Domain adaptation runs a second round of MLM pretraining on domain-specific text (e.g. PubMed abstracts) before task fine-tuning. This shifts BERT's vocabulary statistics and contextual representations toward the target domain — no task labels required.

Stage 1: Domain MLM
Unlabeled domain corpus. Mask 15% tokens, predict them. LR ≈ 1e-5. Outcome: domain-shifted weights.
Stage 2: Task fine-tune
Labeled task data. Swap in classification head. LR ≈ 2e-5. Starts from domain-adapted weights.

BioBERT (Lee et al. 2020) ran Stage 1 on 4.5B biomedical words, then fine-tuned on NER/QA tasks. Gain over standard BERT: +3.6 F1 on biomedical NER — entirely from Stage 1 MLM, zero extra labels.

23. By analogy: Two-stage domain adaptation

Analogy

Discussion prompt

Explain Two-stage domain adaptation by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

BioBERT (Lee et al. 2020) ran Stage 1 on 4.5B biomedical words, then fine-tuned on NER/QA tasks. Gain over standard BERT: +3.6 F1 on biomedical NER — entirely from Stage 1 MLM, zero extra labels.

24. When domain adaptation helps (and when it doesn't)

Concept

scenariodomain adapt helps?reason
Biomedical QA (BioASQ)Yes — +3.6 F1Rare domain vocabulary, BERT OOV-heavy
Legal contract NERYes — significantSpecialized clause structure not in Wikipedia
News sentimentMarginalGeneral BERT already covers news prose well
Short social-media textMinimalDomain MLM data hard to curate; little gain

Key signal: if your domain text contains many subword splits of domain-specific terms (e.g. ##idine, ##ase from BERT's WordPiece on drug names), domain adaptation is likely to help.

25. Fill in: domain adapt helps? for When domain adaptation helps (and when it…

Comparison

Comparison matrix

From When domain adaptation helps (and when it doesn't): refill the domain adapt helps? column from what you know. The rest of the table is as it appeared.

scenariodomain adapt helps?reason
Biomedical QA (BioASQ)Yes — +3.6 F1Rare domain vocabulary, BERT OOV-heavy
Legal contract NERYes — significantSpecialized clause structure not in Wikipedia
News sentimentMarginalGeneral BERT already covers news prose well
Short social-media textMinimalDomain MLM data hard to curate; little gain

26. Catastrophic forgetting and EWC

Section

Part 4b of 4

27. What catastrophic forgetting is

Concept

Catastrophic forgetting: after fine-tuning on Task B, accuracy on previously learned Task A collapses — the new task's gradients overwrite the weight configurations that solved Task A.

training orderTask A accTask B acc
After Task A only87.5%N/A
After naive Task B fine-tune42.5%87.5%
After EWC Task B fine-tune80.0%82.5%

Naive fine-tuning is not broken for Task B — it excels there. The damage is invisible to a single-task evaluation, which is why catastrophic forgetting is easy to miss in practice.

28. Break it if you can: What catastrophic forgetting is

Counterexample

Discussion prompt

Naive fine-tuning is not broken for Task B — it excels there. The damage is invisible to a single-task evaluation, which is why catastrophic forgetting is easy to miss in practice.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

29. Elastic Weight Consolidation (EWC)

Concept

EWC (Kirkpatrick et al. 2017) adds a quadratic penalty to the Task B loss that resists movement of parameters that were important to Task A. Importance is measured by the diagonal of the Fisher information matrix.

\[ \mathcal{L}_{\text{EWC}} = \mathcal{L}_{B}(\theta) + \frac{\lambda}{2} \sum_i F_i \bigl(\theta_i - \theta^*_{A,i}\bigr)^2 \]

\[ F_i = \mathbb{E}_{x\sim\mathcal{D}_A}\!\left[\left(\frac{\partial \log p(y|x,\theta^*_A)}{\partial \theta_i}\right)^2\right] \]

Intuition: F_i is large for parameters whose gradient variance is high on Task A data — they encode critical Task A structure. The penalty makes the optimizer pay a high cost for moving those parameters far from theta*_A.

30. What has to happen first: EWC in PyTorch: Fisher diagonal + penalty

Ranking

Put in order

Put the moves of EWC in PyTorch: Fisher diagonal + penalty into the order they have to happen.

  1. After Task A training converges, save theta_A and compute the diagonal Fisher by looping over Task A data
  2. During Task B training, add the EWC penalty to the cross-entropy loss before backward()
  3. Verify the EWC penalty is non-zero at the start of Task B training

Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. The diagonal Fisher is the expected squared gradient of the log-likelihood.

31. EWC in PyTorch: Fisher diagonal + penalty

Worked example

After Task A training converges, save theta_A and compute the diagonal Fisher by looping over Task A data

Why: The diagonal Fisher is the expected squared gradient of the log-likelihood. We approximate the expectation by averaging over the Task A training set — this is the empirical Fisher.

# --- Step 1: save Task A optimal params ---
theta_A = {n: p.clone().detach()
            for n, p in model.named_parameters()}

# --- Step 2: diagonal Fisher on Task A data ---
fisher = {n: torch.zeros_like(p)
          for n, p in model.named_parameters()}
model.eval()
for x, y in task_a_loader:
    model.zero_grad()
    log_probs = torch.log_softmax(model(x), dim=1)
    # gradient of log p(y|x) w.r.t. params
    log_probs[range(len(y)), y].mean().backward()
    for n, p in model.named_parameters():
        if p.grad is not None:
            fisher[n] += p.grad.data.pow(2)
for n in fisher:
    fisher[n] /= len(task_a_loader.dataset)

During Task B training, add the EWC penalty to the cross-entropy loss before backward()

Why: The penalty gradient pushes theta back toward theta_A for high-Fisher params. Lambda controls the trade-off: too low → forgetting; too high → Task B learns slowly.

Verify the EWC penalty is non-zero at the start of Task B training

Why: At epoch 0, theta ≈ theta_A so the penalty should be near 0. After a few Task B steps it grows as weights move. If the penalty is 0 throughout, the Fisher was not computed correctly.

epochtask_B lossewc_penaltytotal loss
10.69210.00120.6933
100.51028.45319.0633
200.423115.781216.2043
300.387419.402119.7895

32. Inspect it line by line: EWC in PyTorch: Fisher diagonal + penalty

Error analysis

Annotate

Walk the callouts on EWC in PyTorch: Fisher diagonal + penalty. Each one is a place this is easy to get subtly wrong.

  • The diagonal Fisher is the expected squared gradient of the log-likelihood. We approximate the expectation by averaging over the Task A training set — this is the empirical Fisher.
  • The penalty gradient pushes theta back toward theta_A for high-Fisher params. Lambda controls the trade-off: too low → forgetting; too high → Task B learns slowly.
  • At epoch 0, theta ≈ theta_A so the penalty should be near 0. After a few Task B steps it grows as weights move. If the penalty is 0 throughout, the Fisher was not computed correctly.

33. EWC single-parameter walkthrough

Concept

Concrete example: one parameter with F_i = 1.0, theta*_A_i = 0.80, lambda = 1000.

\[ \text{penalty}_i = \frac{1000}{2} \times 1.0 \times (\theta_i - 0.80)^2 \]

theta_i (current)displacementpenalty
0.800.000.0
0.820.020.2
0.900.105.0
1.000.2020.0

The penalty is quadratic — doubling the displacement quadruples the penalty. This creates a soft elastic constraint: small task-driven moves are cheap; large moves that erase Task A structure are expensive.

34. What each one costs: EWC single-parameter walkthrough

Trade off

Comparison matrix

From EWC single-parameter walkthrough: every row here is a choice with a cost. Fill the displacement column, then say which row you would actually pick and what you give up for it.

theta_i (current)displacementpenalty
0.800.000.0
0.820.020.2
0.900.105.0
1.000.2020.0

35. Without one step: Fine-tuning strategy selection recipe

Constraint

Discussion prompt

Run Fine-tuning strategy selection recipe with this step confiscated:

Set the learning rate. Full fine-tuning: 2e-5 to 5e-5. Layer-wise: head at 2e-5, 10× decay per group downward.

Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.

Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.

Answer:

  1. Count your labeled examples. <1 k → feature extraction; >10 k → full fine-tuning or layer-wise LR.
  2. Check domain gap. Many OOV/rare subwords? Add Stage 1 domain MLM before task fine-tuning.
  3. Set the learning rate. Full fine-tuning: 2e-5 to 5e-5. Layer-wise: head at 2e-5, 10× decay per group downward.
  4. Need multiple tasks? After Task A converges: (a) save theta_A, (b) compute diagonal Fisher, (c) add EWC penalty to Task B loss. Lambda in [100, 10000] —…
  5. Validate retention. Evaluate on Task A test set after every Task B epoch. A flat or rising Task A curve confirms EWC is working; a collapsing curve means…

36. Fine-tuning strategy selection recipe

Pattern

  1. Count your labeled examples. <1 k → feature extraction; >10 k → full fine-tuning or layer-wise LR.
  2. Check domain gap. Many OOV/rare subwords? Add Stage 1 domain MLM before task fine-tuning.
  3. Set the learning rate. Full fine-tuning: 2e-5 to 5e-5. Layer-wise: head at 2e-5, 10× decay per group downward.
  4. Need multiple tasks? After Task A converges: (a) save theta_A, (b) compute diagonal Fisher, (c) add EWC penalty to Task B loss. Lambda in [100, 10000] — tune on a held-out Task A split.
  5. Validate retention. Evaluate on Task A test set after every Task B epoch. A flat or rising Task A curve confirms EWC is working; a collapsing curve means raise lambda.

37. Where does it stop working: Fine-tuning strategy selection recipe

Edge cases

Discussion prompt

Fine-tuning strategy selection recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".

Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.

Answer:

  1. Count your labeled examples. <1 k → feature extraction; >10 k → full fine-tuning or layer-wise LR.
  2. Check domain gap. Many OOV/rare subwords? Add Stage 1 domain MLM before task fine-tuning.
  3. Set the learning rate. Full fine-tuning: 2e-5 to 5e-5. Layer-wise: head at 2e-5, 10× decay per group downward.
  4. Need multiple tasks? After Task A converges: (a) save theta_A, (b) compute diagonal Fisher, (c) add EWC penalty to Task B loss. Lambda in [100, 10000] —…
  5. Validate retention. Evaluate on Task A test set after every Task B epoch. A flat or rising Task A curve confirms EWC is working; a collapsing curve means…

38. Rule out three: Check 1 — feature extraction

Elimination

Eliminate the wrong options

You freeze all BERT-base layers and add a 2-class classification head (d_model=768). How many parameters are trainable?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. 768 × 2 + 2 = 1,538
  • B. 768 × 2 = 1,536
  • C. 768 + 2 = 770
  • D. 110,000,000

Survives elimination: A

Why: A Linear(768, 2) layer has weight matrix 768×2 = 1,536 params plus a bias of size 2, giving 1,538 total. Everything else in BERT-base (≈110 M params) is frozen.

39. Check 1 — feature extraction

Check

Work out the answer before clicking.

Check your understanding

You freeze all BERT-base layers and add a 2-class classification head (d_model=768). How many parameters are trainable?

  • A. 768 × 2 + 2 = 1,538 (correct)
  • B. 768 × 2 = 1,536
  • C. 768 + 2 = 770
  • D. 110,000,000

Answer: A

Why: A Linear(768, 2) layer has weight matrix 768×2 = 1,536 params plus a bias of size 2, giving 1,538 total. Everything else in BERT-base (≈110 M params) is frozen.

Why B tempts people
Forgets the bias term. nn.Linear includes a bias by default: d_model × n_classes weights + n_classes biases.
Why C tempts people
Confuses the output dimension (2) with the bias and treats the weight as 768 scalars — mixing up the shape of the weight matrix.
Why D tempts people
Counts all of BERT-base's parameters; the question specifies only the head is trainable with everything else frozen.

40. Answer it before you see the options: Check 2 — layer-wise learning rates

Prediction

Predict first

You use base_lr=2e-5 and alpha=0.1. With 4 transformer layers (indices 0–3), what LR does layer[0] (the bottom layer) get?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: 2e-08

Why: The formula eta_k = base_lr × alpha^(L-k) with L=3 (top layer index), k=0 gives 2e-5 × 0.1^3 = 2e-5 × 0.001 = 2e-8. Bottom layer receives 2e-08.

41. Check 2 — layer-wise learning rates

Check

Trace the formula before clicking.

Check your understanding

You use base_lr=2e-5 and alpha=0.1. With 4 transformer layers (indices 0–3), what LR does layer[0] (the bottom layer) get?

  • A. 2e-05
  • B. 2e-06
  • C. 2e-08 (correct)
  • D. 2e-09

Answer: C

Why: The formula eta_k = base_lr × alpha^(L-k) with L=3 (top layer index), k=0 gives 2e-5 × 0.1^3 = 2e-5 × 0.001 = 2e-8. Bottom layer receives 2e-08.

Why A tempts people
This is the head/top-layer LR — no decay applied. Layer[0] is the bottom layer and receives the maximum decay.
Why B tempts people
This is what layer[2] gets (alpha^1 = 0.1 decay from base): 2e-5 × 0.1^1 = 2e-6. Confuses the depth index.
Why D tempts people
This would be correct with a 5-group schedule including embeddings at depth 4 (2e-5 × 0.1^4 = 2e-9), but with 4 transformer layers the bottom transformer layer is depth 3, not 4.

42. Rule out three: Check 3 — EWC penalty

Elimination

Eliminate the wrong options

For a single parameter: F_i = 2.0, theta*_A_i = 0.5, theta_i = 0.7, lambda = 1000. What is the EWC penalty contribution from this parameter?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. 40.0
  • B. 20.0
  • C. 80.0
  • D. 0.04

Survives elimination: A

Why: penalty = (lambda/2) × F_i × (theta_i - theta*_A_i)^2 = (1000/2) × 2.0 × (0.7 - 0.5)^2 = 500 × 2.0 × 0.04 = 40.0.

43. Check 3 — EWC penalty

Check

Compute the exact value before clicking.

Check your understanding

For a single parameter: F_i = 2.0, theta*_A_i = 0.5, theta_i = 0.7, lambda = 1000. What is the EWC penalty contribution from this parameter?

  • A. 40.0 (correct)
  • B. 20.0
  • C. 80.0
  • D. 0.04

Answer: A

Why: penalty = (lambda/2) × F_i × (theta_i - theta*_A_i)^2 = (1000/2) × 2.0 × (0.7 - 0.5)^2 = 500 × 2.0 × 0.04 = 40.0.

Why B tempts people
Forgets to multiply by F_i = 2.0: (1000/2) × 1.0 × 0.04 = 20.0. Treats Fisher weight as 1 instead of 2.
Why C tempts people
Uses lambda directly instead of lambda/2: 1000 × 2.0 × 0.04 = 80.0. Omits the 1/2 normalization factor from the EWC formula.
Why D tempts people
Computes only the squared displacement (0.7-0.5)^2 = 0.04 without multiplying by F_i or lambda. Skips the two scaling factors entirely.

44. Your Turn — Project overview

Concept

You will implement all four fine-tuning strategies on the same toy BERT-like model and measure the outcome of each. By the end you will have a single script that demonstrates the full spectrum from frozen backbone to EWC-regularized continual learning.

Milestone 1
Build TinyBERT and freeze it for feature extraction. Confirm trainable param count.
Milestone 2
Full fine-tune with uniform LR=2e-5. Compare accuracy vs feature extraction.
Milestone 3
Add layer-wise LR schedule. Verify each group's LR before training.
Milestone 4
Implement EWC. Show Task A retention before and after Task B fine-tune.

45. Which is which: Your Turn — Project overview

Matching

Match the pairs

From Your Turn — Project overview — match each one to what it actually does. The descriptions have been shuffled.

  • c1. Milestone 1
  • c2. Milestone 2
  • c3. Milestone 3
  • c4. Milestone 4
  • b1. Build TinyBERT and freeze it for feature extraction. Confirm trainable param count.
  • b2. Full fine-tune with uniform LR=2e-5. Compare accuracy vs feature extraction.
  • b3. Add layer-wise LR schedule. Verify each group's LR before training.
  • b4. Implement EWC. Show Task A retention before and after Task B fine-tune.

Why: Milestone 1, Milestone 2, Milestone 3, Milestone 4 are easy to tell apart while they are sitting next to their descriptions and much harder afterwards, which is what this checks.

46. Why is this step legal: Build TinyBERT (d=64, 4 layers, vocab=100…

Explain it to yourself

Discussion prompt

In Milestone 1 — TinyBERT feature extraction this move is made:

Build TinyBERT (d=64, 4 layers, vocab=100, seq_len=16) and count total vs trainable params

Why is that legal? Name the rule or definition it rests on before you read on.

Hint: If you can only say "because that is what you do", the rule is the thing to go and find.

Answer:

A small but structurally correct BERT-like model lets you verify every concept without needing pretrained weights or GPU. All param counts should match the formulas you've derived.

47. Milestone 1 — TinyBERT feature extraction

Worked example

Build TinyBERT (d=64, 4 layers, vocab=100, seq_len=16) and count total vs trainable params

Why: A small but structurally correct BERT-like model lets you verify every concept without needing pretrained weights or GPU. All param counts should match the formulas you've derived.

import torch, torch.nn as nn
torch.manual_seed(42)

class TinyBERTLayer(nn.Module):
    def __init__(self, d=64, h=2):
        super().__init__()
        self.attn = nn.MultiheadAttention(d, h, batch_first=True)
        self.ln1  = nn.LayerNorm(d)
        self.ff   = nn.Sequential(
            nn.Linear(d, d*4), nn.GELU(), nn.Linear(d*4, d))
        self.ln2  = nn.LayerNorm(d)
    def forward(self, x):
        a, _ = self.attn(x, x, x)
        x = self.ln1(x + a)
        return self.ln2(x + self.ff(x))

class TinyBERT(nn.Module):
    def __init__(self, vocab=100, d=64, L=4, n_cls=2):
        super().__init__()
        self.embed  = nn.Embedding(vocab, d)
        self.pos    = nn.Embedding(32, d)
        self.layers = nn.ModuleList([TinyBERTLayer(d) for _ in range(L)])
        self.head   = nn.Linear(d, n_cls)
    def forward(self, x):
        p = torch.arange(x.size(1)).unsqueeze(0)
        h = self.embed(x) + self.pos(p)
        for layer in self.layers: h = layer(h)
        return self.head(h[:, 0])

model = TinyBERT()
for n, p in model.named_parameters():
    if 'head' not in n: p.requires_grad = False
total     = sum(p.numel() for p in model.parameters())
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f'total={total}, trainable={trainable}')

Expected output: total≈208514, trainable=130 (Linear(64,2) weight+bias)

Why: Head is 64×2 + 2 = 130. All 208,384 BERT-body params are frozen. If trainable equals total, the freeze loop is missing or the condition is wrong.

componentparamstrainable
embed + pos8,4480
4 × TinyBERTLayer199,9360
head Linear(64,2)130130
total208,514130

48. Which is which, by trainable

Discrimination

Sort into buckets

Sort these by trainable, from memory, without looking back at Milestone 1 — TinyBERT feature extraction. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

0
embed + pos; 4 × TinyBERTLayer
130
head Linear(64,2); total
g1
trainable is "0" for embed + pos, 4 × TinyBERTLayer — that is what the table on "Milestone 1 — TinyBERT feature…" records, and it is the single property separating this group from the rest.
g2
trainable is "130" for head Linear(64,2), total — that is what the table on "Milestone 1 — TinyBERT feature…" records, and it is the single property separating this group from the rest.

49. Predict the next row: Milestone 4 — EWC full program

Pattern

Predict first

The table runs: Naive fine-tune (no EWC) | ~42% | ~87% · EWC (lambda=1000) | ~80% | ~82%

In Milestone 4 — EWC full program, given the rows so far: what is the next one — the row where method is Expected: EWC retains Task A?

Correct: Expected: EWC retains Task A | >70% | >75%

methodTask A acc (retained)Task B acc (new)
Naive fine-tune (no EWC)~42%~87%
EWC (lambda=1000)~80%~82%
Expected: EWC retains Task A>70%>75%

Why: The relationship between the columns, not the individual numbers, is what generates the next row. This is the complete EWC loop. The Fisher must be computed in eval mode with gradients enabled; the EWC penalty must be added every Task B step before calling backward().

50. Milestone 4 — EWC full program

Worked example

Train on Task A, save theta_A, compute diagonal Fisher, then fine-tune on Task B with EWC penalty

Why: This is the complete EWC loop. The Fisher must be computed in eval mode with gradients enabled; the EWC penalty must be added every Task B step before calling backward().

# Task A and B datasets (binary classification, synthetic)
torch.manual_seed(0)
X  = torch.randint(0, 100, (400, 16))
yA = (X.float().mean(1) > 50).long()   # mean-token rule
yB = (X[:, 0] > 50).long()             # first-token rule

model = TinyBERT(); loss_fn = nn.CrossEntropyLoss()
opt   = torch.optim.AdamW(model.parameters(), lr=2e-4)

# --- Train Task A ---
for _ in range(40):
    l = loss_fn(model(X[:320]), yA[:320])
    opt.zero_grad(); l.backward(); opt.step()

theta_A = {n: p.clone().detach() for n, p in model.named_parameters()}

# --- Compute diagonal Fisher on Task A ---
fisher = {n: torch.zeros_like(p) for n, p in model.named_parameters()}
model.eval()
for i in range(320):
    model.zero_grad()
    lp = torch.log_softmax(model(X[i:i+1]), dim=1)
    lp[0, yA[i]].backward()
    for n, p in model.named_parameters():
        if p.grad is not None:
            fisher[n] += p.grad.pow(2)
for n in fisher: fisher[n] /= 320

# --- Fine-tune on Task B with EWC ---
model.train(); lam = 1000
opt2 = torch.optim.AdamW(model.parameters(), lr=2e-4)
for _ in range(40):
    task_loss = loss_fn(model(X[:320]), yB[:320])
    ewc = sum((fisher[n] * (p - theta_A[n]).pow(2)).sum()
               for n, p in model.named_parameters())
    (task_loss + lam/2 * ewc).backward()
    opt2.step(); opt2.zero_grad()

# --- Evaluate retention ---
with torch.no_grad():
    accA = (model(X[320:]).argmax(1)==yA[320:]).float().mean()
    accB = (model(X[320:]).argmax(1)==yB[320:]).float().mean()
print(f'Task A acc={accA:.2%}  Task B acc={accB:.2%}')

Compare with naive fine-tune (no EWC): Task A acc should collapse; with EWC it should stay above 70%

Why: If Task A and Task B have orthogonal decision rules (mean vs first token), naive gradient descent will overwrite the Task A weights. EWC resists this via the Fisher-weighted penalty.

methodTask A acc (retained)Task B acc (new)
Naive fine-tune (no EWC)~42%~87%
EWC (lambda=1000)~80%~82%
Expected: EWC retains Task A>70%>75%

51. Fill in: Task B acc (new) for Milestone 4 — EWC full program

Comparison

Comparison matrix

From Milestone 4 — EWC full program: refill the Task B acc (new) column from what you know. The rest of the table is as it appeared.

methodTask A acc (retained)Task B acc (new)
Naive fine-tune (no EWC)~42%~87%
EWC (lambda=1000)~80%~82%
Expected: EWC retains Task A>70%>75%

52. Connect it up: Lesson 100: BERT Fine-Tuning Strategies

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — Feature extraction — freeze and reuse · Full fine-tuning and the LR danger zone · Layer-wise LR — protect what matters most · Domain adaptation then task fine-tune · Catastrophic forgetting and EWC. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

53. Lesson 100 recap

Recap

You now have four fine-tuning strategies in your toolkit, each with a clear data-size and domain regime where it wins.

strategykey ideawhen to use
Feature extractionFreeze BERT; train head only (130 params for tiny)<1 k labels
Full fine-tuningAll params update; lr 2e-5 to 5e-5>10 k labels
Layer-wise LR10× LR decay per group downward (2e-9 at bottom)Medium data, best accuracy
Domain adaptationStage 1 MLM on domain corpus, Stage 2 task fine-tuneSpecialized vocabulary (bio, legal)
EWCFisher-weighted L2 penalty vs theta*_A prevents forgettingContinual / multi-task learning

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 100 — BERT Fine-Tuning Strategies — Barron · USAAIO Round 2 Preparation, 2026
  2. Devlin et al. 'BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding' (NAACL 2019) — arXiv:1810.04805
  3. Kirkpatrick et al. 'Overcoming catastrophic forgetting in neural networks' (PNAS 2017) — EWC — arXiv:1612.00796
  4. Lee et al. 'BioBERT: a pre-trained biomedical language representation model for biomedical text mining' (Bioinformatics 2020) — arXiv:1901.08746
  5. TinyBERT param count, layer-wise LR schedule, and EWC penalty verified with numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108