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
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.
Objectives
param_groupsWarm-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.
Section
Part 1 of 4
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.
| component | requires_grad | params updated |
|---|---|---|
| BERT embeddings | False | 0 |
| BERT transformer layers (×12) | False | 0 |
| Classification head (Linear) | True | d_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.
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.
| component | requires_grad | params updated |
|---|---|---|
| BERT embeddings | False | 0 |
| BERT transformer layers (×12) | False | 0 |
| Classification head (Linear) | True | d_model × n_classes + n_classes |
Ranking
Put in order
Put the moves of Feature extraction in PyTorch into the order they have to happen.
requires_grad=False looprequires_grad=True params to the optimizerWhy: 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.
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}') # 1538Pass 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.
| epoch | train loss | test acc % |
|---|---|---|
| 1 | 0.6931 | 52.5 |
| 5 | 0.6203 | 62.5 |
| 10 | 0.5671 | 72.5 |
| 15 | 0.5284 | 77.5 |
| 20 | 0.5012 | 80.0 |
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.
| epoch | train loss | test acc % |
|---|---|---|
| 1 | 0.6931 | 52.5 |
| 5 | 0.6203 | 62.5 |
| 10 | 0.5671 | 72.5 |
| 15 | 0.5284 | 77.5 |
| 20 | 0.5012 | 80.0 |
Section
Part 2 of 4
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.
| strategy | trainable | when to use |
|---|---|---|
| Feature extraction | ~0.001% (head only) | <1 k labeled examples |
| Full fine-tuning | 100% | >10 k labeled examples |
| Layer-wise LR | 100% (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.
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.
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.
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.
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.
Section
Part 3 of 4
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.
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.
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
| group | lr | ratio to head |
|---|---|---|
| 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 |
| head | 2.00e-05 | 1 / 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.
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.
| group | lr | ratio to head |
|---|---|---|
| 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 |
| head | 2.00e-05 | 1 / 1 |
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.param object in multiple groups — PyTorch warns but still trains, causing silent double-counting. Print and cross-check.Section
Part 4a of 4
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.
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.
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.
Concept
| scenario | domain adapt helps? | reason |
|---|---|---|
| Biomedical QA (BioASQ) | Yes — +3.6 F1 | Rare domain vocabulary, BERT OOV-heavy |
| Legal contract NER | Yes — significant | Specialized clause structure not in Wikipedia |
| News sentiment | Marginal | General BERT already covers news prose well |
| Short social-media text | Minimal | Domain 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.
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.
| scenario | domain adapt helps? | reason |
|---|---|---|
| Biomedical QA (BioASQ) | Yes — +3.6 F1 | Rare domain vocabulary, BERT OOV-heavy |
| Legal contract NER | Yes — significant | Specialized clause structure not in Wikipedia |
| News sentiment | Marginal | General BERT already covers news prose well |
| Short social-media text | Minimal | Domain MLM data hard to curate; little gain |
Section
Part 4b of 4
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 order | Task A acc | Task B acc |
|---|---|---|
| After Task A only | 87.5% | N/A |
| After naive Task B fine-tune | 42.5% | 87.5% |
| After EWC Task B fine-tune | 80.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.
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.
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.
Ranking
Put in order
Put the moves of EWC in PyTorch: Fisher diagonal + penalty into the order they have to happen.
theta_A and compute the diagonal Fisher by looping over Task A databackward()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.
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.
| epoch | task_B loss | ewc_penalty | total loss |
|---|---|---|---|
| 1 | 0.6921 | 0.0012 | 0.6933 |
| 10 | 0.5102 | 8.4531 | 9.0633 |
| 20 | 0.4231 | 15.7812 | 16.2043 |
| 30 | 0.3874 | 19.4021 | 19.7895 |
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.
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) | displacement | penalty |
|---|---|---|
| 0.80 | 0.00 | 0.0 |
| 0.82 | 0.02 | 0.2 |
| 0.90 | 0.10 | 5.0 |
| 1.00 | 0.20 | 20.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.
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) | displacement | penalty |
|---|---|---|
| 0.80 | 0.00 | 0.0 |
| 0.82 | 0.02 | 0.2 |
| 0.90 | 0.10 | 5.0 |
| 1.00 | 0.20 | 20.0 |
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:
theta_A, (b) compute diagonal Fisher, (c) add EWC penalty to Task B loss. Lambda in [100, 10000] —…Pattern
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.lambda.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:
theta_A, (b) compute diagonal Fisher, (c) add EWC penalty to Task B loss. Lambda in [100, 10000] —…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.
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.
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?
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.
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.
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?
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.
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.
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.
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?
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.
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.
Matching
Match the pairs
From Your Turn — Project overview — match each one to what it actually does. The descriptions have been shuffled.
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.
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.
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.
| component | params | trainable |
|---|---|---|
| embed + pos | 8,448 | 0 |
| 4 × TinyBERTLayer | 199,936 | 0 |
| head Linear(64,2) | 130 | 130 |
| total | 208,514 | 130 |
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.
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%
| method | Task 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().
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.
| method | Task 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% |
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.
| method | Task 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% |
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.
Recap
You now have four fine-tuning strategies in your toolkit, each with a clear data-size and domain regime where it wins.
| strategy | key idea | when to use |
|---|---|---|
| Feature extraction | Freeze BERT; train head only (130 params for tiny) | <1 k labels |
| Full fine-tuning | All params update; lr 2e-5 to 5e-5 | >10 k labels |
| Layer-wise LR | 10× LR decay per group downward (2e-9 at bottom) | Medium data, best accuracy |
| Domain adaptation | Stage 1 MLM on domain corpus, Stage 2 task fine-tune | Specialized vocabulary (bio, legal) |
| EWC | Fisher-weighted L2 penalty vs theta*_A prevents forgetting | Continual / multi-task learning |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.