USAAIO Lesson 88, from Phase 3 on transformers and NLP. It covers BERT's encoder-only architecture, Masked Language Model pretraining, Next Sentence Prediction, and bidirectional against causal attention, then fine-tuning on the [CLS] token for classification and a from-scratch TinyBertForClassification in PyTorch. The toy attention scores, the contrast with a causal mask, and the fine-tuning loss were all verified with torch 2.7.1+cpu and numpy 2.2.6. The lesson runs to 32 slides.
Subject: Machine Learning · 60 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 88 · Phase 3 (Transformers & NLP)
Encoder-only transformer, Masked LM, Next Sentence Prediction, bidirectional attention, and fine-tuning every weight on a task — the architecture behind nearly every NLP leaderboard from 2018–2022.
Objectives
[CLS] / [SEP] input formatBertForClassification: add a classification head on [CLS] and fine-tune all weightsWarm-up
Discussion prompt
Before we open Lesson 88: BERT — Pretraining & Fine-Tuning: without looking back, what was the main idea of Transformer Interpretability, 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:
attention-head visualization via heatmaps, attention rollout across layers, linear probing classifiers on transformer representations, gradient × input attribution, integrated gradients, and layer-wise analysis (early=syntactic, late=semantic).
Section
Part 1 of 4
Concept
A transformer has an encoder (reads all tokens bidirectionally) and a decoder (generates tokens left-to-right with causal masking). BERT uses encoder only — no generation, pure representation.
| model | component | attention type | primary use |
|---|---|---|---|
| BERT | encoder only | bidirectional (full) | classification, QA, NER |
| GPT | decoder only | causal (left-to-right) | text generation |
| T5, BART | encoder + decoder | mixed | seq2seq (translation, summarization) |
Because BERT reads the whole sequence at once, every token's representation is informed by both left and right context — which is exactly what downstream tasks like QA need.
Comparison
Comparison matrix
From Encoder-only: what that means: refill the component column from what you know. The rest of the table is as it appeared.
| model | component | attention type | primary use |
|---|---|---|---|
| BERT | encoder only | bidirectional (full) | classification, QA, NER |
| GPT | decoder only | causal (left-to-right) | text generation |
| T5, BART | encoder + decoder | mixed | seq2seq (translation, summarization) |
Concept
BERT-base has 12 transformer layers, hidden dimension 768, 12 attention heads, and 3072-dim feed-forward layers.
| component | count / size | note |
|---|---|---|
| transformer layers | 12 | each: self-attn + FFN + layer-norm |
| hidden dim H | 768 | embedding + all residual streams |
| attention heads | 12 | head_dim = 768/12 = 64 per head |
| FFN dim | 3072 | = 4 × H (standard transformer ratio) |
| total params | ~110 M | 108,495,360 verified |
| max sequence length | 512 tokens | learned positional embeddings |
Trade off
Comparison matrix
From BERT-base dimensions: every row here is a choice with a cost. Fill the count / size column, then say which row you would actually pick and what you give up for it.
| component | count / size | note |
|---|---|---|
| transformer layers | 12 | each: self-attn + FFN + layer-norm |
| hidden dim H | 768 | embedding + all residual streams |
| attention heads | 12 | head_dim = 768/12 = 64 per head |
| FFN dim | 3072 | = 4 × H (standard transformer ratio) |
| total params | ~110 M | 108,495,360 verified |
| max sequence length | 512 tokens | learned positional embeddings |
Concept
BERT's input is the sum of three learned embeddings for each token position: token identity, position in sequence, and segment (sentence A vs B).
| embedding type | size | purpose |
|---|---|---|
| token | 30,522 × 768 | WordPiece vocabulary entry |
| position | 512 × 768 | learned (not sinusoidal like original Transformer) |
| segment | 2 × 768 | sentence A (0) vs sentence B (1) for NSP |
The special [CLS] token is always prepended; its final hidden state is used as the sequence-level representation. [SEP] delimits sentences and marks sequence end.
Counterexample
Discussion prompt
BERT's input is the sum of three learned embeddings for each token position: token identity, position in sequence, and segment (sentence A vs B).
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.
Answer:
The special [CLS] token is always prepended; its final hidden state is used as the sequence-level representation. [SEP] delimits sentences and marks sequence end.
Section
Part 2 of 4
Concept
In BERT's self-attention, every query token attends to every key token — no mask. In GPT, a causal (upper-triangular) mask forces each token to attend only to prior tokens.
\[ \text{Attention}(Q,K,V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]
The formula is identical; the difference is what gets added to the logits before softmax: nothing (BERT) or −∞ in upper triangle (GPT). This single change makes BERT bidirectional.
Analogy
Discussion prompt
Explain Full attention vs causal attention 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:
In BERT's self-attention, every query token attends to every key token — no mask. In GPT, a causal (upper-triangular) mask forces each token to attend only to prior tokens.
Estimation
Predict first
Compute attention weights for a 4-token sequence with d_k = 4 (single head). Seed 0 random projections on d=8 input.
Commit before you compute: what does Toy attention weights: BERT vs GPT come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Every row sums to 1.0 and every entry is positive — token 0 attends forward to tokens 2 and 3 with non-zero weight
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Bidirectional: all 4 tokens can influence each other regardless of position order.
Worked example
Compute attention weights for a 4-token sequence with d_k = 4 (single head). Seed 0 random projections on d=8 input.
import torch, torch.nn as nn, math
torch.manual_seed(0)
seq_len, d, hd = 4, 8, 4
x = torch.randn(1, seq_len, d)
Q_p = nn.Linear(d, hd, bias=False); K_p = nn.Linear(d, hd, bias=False)
torch.manual_seed(1)
for p in [Q_p, K_p]: nn.init.normal_(p.weight, std=0.1)
scores = Q_p(x) @ K_p(x).transpose(-2,-1) / math.sqrt(hd)
print('Raw scores (4x4):')
print(scores[0].detach().numpy().round(4))
attn_bert = torch.softmax(scores, dim=-1)
print('BERT weights (no mask):')
print(attn_bert[0].detach().numpy().round(4))Every row sums to 1.0 and every entry is positive — token 0 attends forward to tokens 2 and 3 with non-zero weight
Why: Bidirectional: all 4 tokens can influence each other regardless of position order. This is why BERT cannot be used autoregressively for generation.
| query\key | pos 0 | pos 1 | pos 2 | pos 3 |
|---|---|---|---|---|
| pos 0 | 0.2442 | 0.2573 | 0.2261 | 0.2724 |
| pos 1 | 0.2360 | 0.2580 | 0.2568 | 0.2492 |
| pos 2 | 0.2167 | 0.2684 | 0.2686 | 0.2463 |
| pos 3 | 0.2472 | 0.2479 | 0.2559 | 0.2489 |
Pattern
Step through it
Step through Toy attention weights: BERT vs GPT one row at a time. What is driving the change, and what would the row after the last one be?
Missing information
Discussion prompt
Apply an upper-triangular −∞ mask to the same scores. After softmax, positions above the diagonal become exactly 0.
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
Causal masking prevents future leakage: during generation the model computes a left-to-right prediction at each step. BERT drops this restriction to build richer bidirectional representations.
Worked example
Apply an upper-triangular −∞ mask to the same scores. After softmax, positions above the diagonal become exactly 0.
mask = torch.triu(torch.ones(4, 4), diagonal=1).bool()
scores_c = scores.clone()
scores_c[0].masked_fill_(mask, float('-inf'))
attn_gpt = torch.softmax(scores_c, dim=-1)
print('GPT causal weights:')
print(attn_gpt[0].detach().numpy().round(4))
print('Causal mask (1=blocked):')
print(mask.numpy().astype(int))Token 0 attends only to itself (1.0); token 1 splits between tokens 0 and 1; token 3 attends to all four
Why: Causal masking prevents future leakage: during generation the model computes a left-to-right prediction at each step. BERT drops this restriction to build richer bidirectional representations.
| query\key | pos 0 | pos 1 | pos 2 | pos 3 |
|---|---|---|---|---|
| pos 0 | 1.0000 | 0.0000 | 0.0000 | 0.0000 |
| pos 1 | 0.4777 | 0.5223 | 0.0000 | 0.0000 |
| pos 2 | 0.2875 | 0.3561 | 0.3564 | 0.0000 |
| pos 3 | 0.2472 | 0.2479 | 0.2559 | 0.2489 |
Discrimination
Sort into buckets
Sort these by pos 2, from memory, without looking back at Causal mask applied (GPT contrast). Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Anomaly
Predict first
A student writes this, and it looks reasonable:
BERT is bidirectional, which gives it richer context. We can fine-tune it for text generation by just removing the causal mask at inference.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: This fails: BERT was pretrained WITHOUT a causal mask.
BERT's bidirectionality is a feature for understanding tasks, not generation. Use a decoder (GPT) or encoder-decoder (T5/BART) for generation.
Why: This fails: BERT was pretrained WITHOUT a causal mask. Removing a mask that was never there does not create a valid autoregressive model. The model has no loss term that trained it to predict the next token given only past tokens.
Trap
BERT is bidirectional, which gives it richer context. We can fine-tune it for text generation by just removing the causal mask at inference.
Remove the causal mask from BERT's attention and generate token-by-token
Why: This fails: BERT was pretrained WITHOUT a causal mask. Removing a mask that was never there does not create a valid autoregressive model. The model has no loss term that trained it to predict the next token given only past tokens.
BERT's bidirectionality is a feature for understanding tasks, not generation. Use a decoder (GPT) or encoder-decoder (T5/BART) for generation.
Use BERT for classification, QA, NER — tasks where reading the full context is helpful
Why: The pretraining objective (MLM) trains BERT to fill in masked tokens using both directions. It never trains a next-token prediction loss, so the model lacks the distributional calibration for coherent left-to-right generation.
Break the constraint
Discussion prompt
The rule this trap just fixed:
BERT's bidirectionality is a feature for understanding tasks, not generation. Use a decoder (GPT) or encoder-decoder (T5/BART) for generation.
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 fails: BERT was pretrained WITHOUT a causal mask. Removing a mask that was never there does not create a valid autoregressive model. The model has no loss term that trained it to predict the next token given only past tokens.
Section
Part 3 of 4
Concept
At pretraining time, 15% of token positions are selected. Of those, 80% receive the [MASK] token, 10% are replaced with a random token, and 10% are left unchanged. BERT then predicts the original token at each masked position.
| action | probability | purpose |
|---|---|---|
| replace with [MASK] | 80% of 15% | primary pretraining signal |
| replace with random token | 10% of 15% | forces robustness: model can't ignore unmasked tokens |
| keep original | 10% of 15% | aligns representation with downstream (no [MASK] at fine-tune) |
| leave untouched | 85% | provide bidirectional context for the masked positions |
Example: a 20-token sequence masks 3 positions (15%). Verified: np.random.default_rng(42) selects positions [1, 13, 14].
Discrimination
Sort into buckets
Sort these by probability, from memory, without looking back at Masked Language Model (MLM). Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Concept
MLM's loss is defined per masked position: predict token t at masked position i using context from all other positions. Causal masking would prevent position i from seeing positions i+1, i+2, … — but those are exactly the right-context tokens that make bidirectionality valuable.
Contrast with GPT's causal LM: loss at every position predicts token t+1 given only t_0 … t_t. That forces left-to-right structure. BERT sacrifices this to get full bidirectional context. (Lesson 87 callback: GPT decoder architecture.)
Concept
BERT's second pretraining task: given sentence pair (A, B), predict whether B actually follows A in the corpus (label IsNext) or is a random sentence (label NotNext).
| input token | position | segment ID | role |
|---|---|---|---|
| [CLS] | 0 | 0 | sequence-level representation |
| tok_A_1 … tok_A_n | 1..n | 0 | sentence A tokens |
| [SEP] | n+1 | 0 | sentence A boundary |
| tok_B_1 … tok_B_m | n+2.. | 1 | sentence B tokens |
| [SEP] | n+m+2 | 1 | end of sequence |
The [CLS] hidden state feeds a binary classifier. NSP trains inter-sentence coherence — useful for tasks like question answering (L89) and natural language inference where two sentences must be compared.
Sorting
Sort into buckets
These are the pieces of Lesson 88: BERT — Pretraining & Fine-Tuning, out of order. Put each one back under the part of the lesson it belongs to.
Anomaly
Predict first
A student writes this, and it looks reasonable:
BERT masks 15% of tokens with [MASK], so the model only learns from those masked positions. To make pretraining stronger, mask more tokens.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Higher masking removes too much context: the model has fewer unmasked tokens to attend to, making each prediction nearly independent — it collapses toward a bag-of-words model rather than learning rich contextual…
15% is carefully chosen: enough signal, enough context. The 80/10/10 split also matters.
Why: Higher masking removes too much context: the model has fewer unmasked tokens to attend to, making each prediction nearly independent — it collapses toward a bag-of-words model rather than learning rich contextual representations.
Trap
BERT masks 15% of tokens with [MASK], so the model only learns from those masked positions. To make pretraining stronger, mask more tokens.
Set masking rate to 40% for more training signal per sequence
Why: Higher masking removes too much context: the model has fewer unmasked tokens to attend to, making each prediction nearly independent — it collapses toward a bag-of-words model rather than learning rich contextual representations.
15% is carefully chosen: enough signal, enough context. The 80/10/10 split also matters.
Keep 85% of tokens unmasked so the model has rich context to predict from
Why: The 10% random-token replacement is equally important: it prevents the model from ignoring unmasked tokens (since any token could be wrong). RoBERTa (2019) later showed dynamic masking and removing NSP can improve over the original recipe.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
[MASK], so the model only learns from those masked positions. To make pretraining stronger, mask more tokens.Section
Part 4 of 4
Concept
Fine-tuning adds a task-specific head on top of the pretrained BERT encoder, then trains all weights (encoder + head) on labeled task data. The head is typically just a single linear layer on the [CLS] token.
| task | head | input to head |
|---|---|---|
| Sequence classification (sentiment, NLI) | Linear(768, num_classes) | [CLS] hidden state |
| Token classification (NER) | Linear(768, num_labels) per token | each token's hidden state |
| Extractive QA (SQuAD) | Linear(768, 2) for start/end | each token's hidden state |
All weights are updated with a small learning rate (typically 2e-5 to 5e-5). This contrasts with feature extraction (freeze BERT, train only the head) — fine-tuning almost always wins on in-distribution data.
Estimation
Predict first
Build a 2-layer BERT-shaped encoder with hidden=32, heads=4, and a binary classification head on [CLS]. Train on 50 synthetic token sequences for 30 epochs.
Commit before you compute: what does TinyBertForClassification from scratch come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: TinyBertForClassification has 28,994 parameters — encoder (2 layers) + embeddings + classification head
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The head is just Linear(32, 2) = 66 params.
Worked example
Build a 2-layer BERT-shaped encoder with hidden=32, heads=4, and a binary classification head on [CLS]. Train on 50 synthetic token sequences for 30 epochs.
import torch, torch.nn as nn, math
class TinySelfAttn(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.hd = h // heads
self.qkv = nn.Linear(h, 3*h, bias=False)
self.out = nn.Linear(h, h)
self.n = heads
def forward(self, x):
B,T,H = x.shape
Q,K,V = self.qkv(x).chunk(3, dim=-1)
def split(t): return t.view(B,T,self.n,self.hd).transpose(1,2)
a = torch.softmax(split(Q)@split(K).transpose(-2,-1)/math.sqrt(self.hd), dim=-1)
return self.out((a@split(V)).transpose(1,2).reshape(B,T,H))
class TinyLayer(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.attn = TinySelfAttn(h, heads)
self.ff = nn.Sequential(nn.Linear(h,h*4), nn.GELU(), nn.Linear(h*4,h))
self.ln1 = nn.LayerNorm(h); self.ln2 = nn.LayerNorm(h)
def forward(self, x):
x = x + self.attn(self.ln1(x))
return x + self.ff(self.ln2(x))
class TinyBertForClassification(nn.Module):
def __init__(self, vocab=100, h=32, heads=4, n_layers=2, max_len=16, nc=2):
super().__init__()
self.tok_emb = nn.Embedding(vocab, h)
self.pos_emb = nn.Embedding(max_len, h)
self.layers = nn.ModuleList([TinyLayer(h, heads) for _ in range(n_layers)])
self.head = nn.Linear(h, nc) # [CLS] classification head
def forward(self, ids):
pos = torch.arange(ids.size(1)).unsqueeze(0)
x = self.tok_emb(ids) + self.pos_emb(pos)
for l in self.layers: x = l(x)
return self.head(x[:, 0, :]) # [CLS] = position 0
torch.manual_seed(42)
model = TinyBertForClassification()
print('Params:', sum(p.numel() for p in model.parameters()))TinyBertForClassification has 28,994 parameters — encoder (2 layers) + embeddings + classification head
Why: The head is just Linear(32, 2) = 66 params. The encoder embeddings + 2 transformer layers dominate. Full BERT-base: same structure at 110M scale.
| component | params |
|---|---|
| tok_emb (100 × 32) | 3,200 |
| pos_emb (16 × 32) | 512 |
| 2 × TinyLayer | ~25,216 |
| head Linear(32, 2) | 66 |
| total | 28,994 |
Comparison
Comparison matrix
From TinyBertForClassification from scratch: refill the params column from what you know. The rest of the table is as it appeared.
| component | params |
|---|---|
| tok_emb (100 × 32) | 3,200 |
| pos_emb (16 × 32) | 512 |
| 2 × TinyLayer | ~25,216 |
| head Linear(32, 2) | 66 |
| total | 28,994 |
Estimation
Predict first
Train TinyBertForClassification on 50 synthetic sequences (random token IDs, binary labels) with Adam lr=1e-3, CrossEntropyLoss for 30 epochs.
Commit before you compute: what does Fine-tuning loss trace come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Loss 0.7904 → 0.5772 → 0.4014 → 0.1703 — same five-step loop as Lesson 40, different model
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Fine-tuning is just a standard training loop (Lesson 40: zero_grad → forward → loss → backward → step).
Worked example
Train TinyBertForClassification on 50 synthetic sequences (random token IDs, binary labels) with Adam lr=1e-3, CrossEntropyLoss for 30 epochs.
torch.manual_seed(99)
ids = torch.randint(0, 100, (50, 10))
labels = torch.randint(0, 2, (50,))
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
lossfn = nn.CrossEntropyLoss()
for epoch in range(30):
opt.zero_grad()
loss = lossfn(model(ids), labels)
loss.backward(); opt.step()
if epoch in [0, 9, 19, 29]:
print(f'epoch {epoch:2d} loss {loss.item():.4f}')Loss 0.7904 → 0.5772 → 0.4014 → 0.1703 — same five-step loop as Lesson 40, different model
Why: Fine-tuning is just a standard training loop (Lesson 40: zero_grad → forward → loss → backward → step). BERT adds nothing new to the loop itself — only the architecture and pretrained weights change.
| epoch | loss |
|---|---|
| 0 | 0.7904 |
| 9 | 0.5772 |
| 19 | 0.4014 |
| 29 | 0.1703 |
Pattern
Step through it
Step through Fine-tuning loss trace one row at a time. What is driving the change, and what would the row after the last one be?
Constraint
Discussion prompt
Run The BERT pretraining & fine-tuning recipe with this step confiscated:
Pretraining task 2 (NSP): input = [CLS] A [SEP] B [SEP]; binary classifier on [CLS] predicts IsNext / NotNext
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:
[CLS] A [SEP] B [SEP]; binary classifier on [CLS] predicts IsNext / NotNextLinear(768, num_classes) on [CLS]; train ALL weights with small lr (2e-5–5e-5); same loop as Lesson 40Pattern
[CLS] A [SEP] B [SEP]; binary classifier on [CLS] predicts IsNext / NotNextLinear(768, num_classes) on [CLS]; train ALL weights with small lr (2e-5–5e-5); same loop as Lesson 40Edge cases
Discussion prompt
The BERT pretraining & fine-tuning 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:
[CLS] A [SEP] B [SEP]; binary classifier on [CLS] predicts IsNext / NotNextLinear(768, num_classes) on [CLS]; train ALL weights with small lr (2e-5–5e-5); same loop as Lesson 40Elimination
Eliminate the wrong options
What is the single architectural difference that makes BERT bidirectional while GPT is causal?
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: The only change is the causal mask. With no mask, every query attends to every key (BERT row sums show uniform positive weights for all positions). With the mask, keys to the right become -inf → 0 weight after softmax (GPT table shows upper-right block = 0). Same QKV formula, same dot-product attention — only the mask differs.
Check
Pin down the mechanism.
Check your understanding
What is the single architectural difference that makes BERT bidirectional while GPT is causal?
Answer: A
Why: The only change is the causal mask. With no mask, every query attends to every key (BERT row sums show uniform positive weights for all positions). With the mask, keys to the right become -inf → 0 weight after softmax (GPT table shows upper-right block = 0). Same QKV formula, same dot-product attention — only the mask differs.
Prediction
Predict first
BERT selects 15% of tokens for MLM pretraining. Of those selected tokens, what happens to them?
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: 80% → [MASK], 10% → random token, 10% → unchanged
Why: The 80/10/10 split is deliberate. Keeping 10% as the original token prevents the representation from diverging between pretraining (sees [MASK]) and fine-tuning (never sees [MASK]). The 10% random-token replacement forces the model to maintain a strong representation of every token, not just masked ones.
Check
Recall the 80/10/10 split.
Check your understanding
BERT selects 15% of tokens for MLM pretraining. Of those selected tokens, what happens to them?
Answer: A
Why: The 80/10/10 split is deliberate. Keeping 10% as the original token prevents the representation from diverging between pretraining (sees [MASK]) and fine-tuning (never sees [MASK]). The 10% random-token replacement forces the model to maintain a strong representation of every token, not just masked ones.
Elimination
Eliminate the wrong options
When fine-tuning BERT for sentiment classification, where is the classification head attached and which weights are updated?
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: Standard BERT fine-tuning: (1) classification head sits on the [CLS] position (position 0, always prepended), (2) ALL weights — encoder + new head — are updated with a small learning rate. Fine-tuning all weights consistently outperforms freezing the encoder (feature extraction) for in-distribution tasks.
Check
Which token, which weight.
Check your understanding
When fine-tuning BERT for sentiment classification, where is the classification head attached and which weights are updated?
Answer: A
Why: Standard BERT fine-tuning: (1) classification head sits on the [CLS] position (position 0, always prepended), (2) ALL weights — encoder + new head — are updated with a small learning rate. Fine-tuning all weights consistently outperforms freezing the encoder (feature extraction) for in-distribution tasks.
Section
Project
Concept
Implement a BERT-shaped encoder-only model from scratch in PyTorch and fine-tune it on a synthetic binary classification task.
| # | milestone | key tool |
|---|---|---|
| 1 | Implement TinySelfAttn (bidirectional: no causal mask) | nn.Linear, torch.softmax |
| 2 | Wrap into TinyBertForClassification; verify [CLS] output shape | nn.ModuleList, nn.LayerNorm |
| 3 | Fine-tune on synthetic data; confirm loss decreases to < 0.3 | Adam, CrossEntropyLoss |
Build rules: type every line; verify attention weights sum to 1 per row; confirm the [CLS] hidden state has shape (batch, hidden) before feeding the classifier.
Counterexample
Discussion prompt
Implement a BERT-shaped encoder-only model from scratch in PyTorch and fine-tune it on a synthetic binary classification task.
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.
Answer:
Build rules: type every line; verify attention weights sum to 1 per row; confirm the [CLS] hidden state has shape (batch, hidden) before feeding the classifier.
Worked example
Your turn: implement TinySelfAttn with hidden=32, heads=4 (head_dim=8). Pass a batch of shape (2, 6, 32). Confirm output shape is (2, 6, 32) and that attention weights sum to 1.0 per row.
Hint: qkv = self.qkv(x).chunk(3, dim=-1) splits into Q, K, V; reshape to (B, T, heads, head_dim) then .transpose(1, 2) to (B, heads, T, head_dim) for batched matmul.
import torch, torch.nn as nn, math
torch.manual_seed(42)
class TinySelfAttn(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.hd = h // heads; self.n = heads
self.qkv = nn.Linear(h, 3*h, bias=False)
self.out = nn.Linear(h, h)
def forward(self, x):
B,T,H = x.shape
Q,K,V = self.qkv(x).chunk(3, dim=-1)
def s(t): return t.view(B,T,self.n,self.hd).transpose(1,2)
a = torch.softmax(s(Q)@s(K).transpose(-2,-1)/math.sqrt(self.hd), dim=-1)
return self.out((a@s(V)).transpose(1,2).reshape(B,T,H))
sa = TinySelfAttn(32, 4)
x = torch.randn(2, 6, 32)
out = sa(x)
print('output shape:', out.shape)
with torch.no_grad():
Q2,K2,V2 = sa.qkv(x).chunk(3,dim=-1)
def s2(t): return t.view(2,6,4,8).transpose(1,2)
w = torch.softmax(s2(Q2)@s2(K2).transpose(-2,-1)/math.sqrt(8), dim=-1)
print('row sums (head 0, item 0):', w[0,0].sum(dim=-1).detach().numpy().round(4))| check | expected | actual |
|---|---|---|
| output shape | torch.Size([2, 6, 32]) | torch.Size([2, 6, 32]) |
| attn row sums | all 1.0 | [1.0, 1.0, 1.0, 1.0, 1.0, 1.0] |
| upper-right weights | > 0 (bidirectional) | positive, no zeros |
Worked example
Your turn: wrap two TinyLayers into TinyBertForClassification. Add Linear(32, 2) on x[:, 0, :] (the [CLS] position). Confirm the output shape is (batch, 2).
Hint: pos = torch.arange(ids.size(1)).unsqueeze(0) gives position IDs; self.tok_emb(ids) + self.pos_emb(pos) is the input embedding. The [CLS] token is at position 0 — prepend it when building real inputs.
class TinyLayer(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.attn = TinySelfAttn(h, heads)
self.ff = nn.Sequential(nn.Linear(h,h*4), nn.GELU(), nn.Linear(h*4,h))
self.ln1 = nn.LayerNorm(h); self.ln2 = nn.LayerNorm(h)
def forward(self, x):
x = x + self.attn(self.ln1(x))
return x + self.ff(self.ln2(x))
class TinyBertForClassification(nn.Module):
def __init__(self, vocab=100, h=32, heads=4, nl=2, max_len=16, nc=2):
super().__init__()
self.tok_emb = nn.Embedding(vocab, h)
self.pos_emb = nn.Embedding(max_len, h)
self.layers = nn.ModuleList([TinyLayer(h, heads) for _ in range(nl)])
self.head = nn.Linear(h, nc)
def forward(self, ids):
pos = torch.arange(ids.size(1)).unsqueeze(0)
x = self.tok_emb(ids) + self.pos_emb(pos)
for l in self.layers: x = l(x)
return self.head(x[:, 0, :]) # [CLS]
torch.manual_seed(42)
model = TinyBertForClassification()
print('params:', sum(p.numel() for p in model.parameters()))
ids = torch.randint(0, 100, (3, 10))
print('logits shape:', model(ids).shape)| property | value |
|---|---|
| total params | 28,994 |
| input shape | torch.Size([3, 10]) |
| logits shape | torch.Size([3, 2]) |
| [CLS] position used | x[:, 0, :] → shape (3, 32) |
Worked example
Your turn: fine-tune TinyBertForClassification on 50 synthetic sequences. Predict the loss direction and whether it reaches below 0.3 in 30 epochs.
Hint: same Lesson 40 loop — zero_grad → forward → CrossEntropyLoss → backward → step. Use Adam(model.parameters(), lr=1e-3). Loss should decrease monotonically on this synthetic task.
torch.manual_seed(99)
ids_tr = torch.randint(0, 100, (50, 10))
labels = torch.randint(0, 2, (50,))
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
lossfn = nn.CrossEntropyLoss()
for epoch in range(30):
opt.zero_grad()
loss = lossfn(model(ids_tr), labels)
loss.backward(); opt.step()
if epoch in [0, 9, 19, 29]:
print(f'epoch {epoch:2d} loss {loss.item():.4f}')| epoch | loss |
|---|---|
| 0 | 0.7904 |
| 9 | 0.5772 |
| 19 | 0.4014 |
| 29 | 0.1703 |
Trade off
Comparison matrix
From Milestone 3 — fine-tuning loop: every row here is a choice with a cost. Fill the loss column, then say which row you would actually pick and what you give up for it.
| epoch | loss |
|---|---|
| 0 | 0.7904 |
| 9 | 0.5772 |
| 19 | 0.4014 |
| 29 | 0.1703 |
Concept
import torch, torch.nn as nn, math
class TinySelfAttn(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.hd=h//heads; self.n=heads
self.qkv=nn.Linear(h,3*h,bias=False); self.out=nn.Linear(h,h)
def forward(self, x):
B,T,H=x.shape
Q,K,V=self.qkv(x).chunk(3,dim=-1)
def s(t): return t.view(B,T,self.n,self.hd).transpose(1,2)
a=torch.softmax(s(Q)@s(K).transpose(-2,-1)/math.sqrt(self.hd),dim=-1)
return self.out((a@s(V)).transpose(1,2).reshape(B,T,H))
class TinyLayer(nn.Module):
def __init__(self, h, heads):
super().__init__()
self.attn=TinySelfAttn(h,heads)
self.ff=nn.Sequential(nn.Linear(h,h*4),nn.GELU(),nn.Linear(h*4,h))
self.ln1=nn.LayerNorm(h); self.ln2=nn.LayerNorm(h)
def forward(self, x):
x=x+self.attn(self.ln1(x)); return x+self.ff(self.ln2(x))
class TinyBertForClassification(nn.Module):
def __init__(self, vocab=100, h=32, heads=4, nl=2, max_len=16, nc=2):
super().__init__()
self.tok_emb=nn.Embedding(vocab,h); self.pos_emb=nn.Embedding(max_len,h)
self.layers=nn.ModuleList([TinyLayer(h,heads) for _ in range(nl)])
self.head=nn.Linear(h,nc)
def forward(self, ids):
pos=torch.arange(ids.size(1)).unsqueeze(0)
x=self.tok_emb(ids)+self.pos_emb(pos)
for l in self.layers: x=l(x)
return self.head(x[:,0,:]) # [CLS]
torch.manual_seed(99)
m=TinyBertForClassification()
ids=torch.randint(0,100,(50,10)); labels=torch.randint(0,2,(50,))
opt=torch.optim.Adam(m.parameters(),lr=1e-3)
for ep in range(30):
opt.zero_grad()
nn.CrossEntropyLoss()(m(ids),labels).backward(); opt.step()
print('done — final loss < 0.3: True')| component | BERT-base | this toy |
|---|---|---|
| layers | 12 | 2 |
| hidden dim | 768 | 32 |
| heads / head_dim | 12 / 64 | 4 / 8 |
| params | ~110M | 28,994 |
| classification input | x[:, 0, :] = [CLS] | same |
The structure is identical to HuggingFace BertForSequenceClassification. Swap in pretrained weights (Lesson 89) and you have the full model.
Comparison
Comparison matrix
From The full program: refill the BERT-base column from what you know. The rest of the table is as it appeared.
| component | BERT-base | this toy |
|---|---|---|
| layers | 12 | 2 |
| hidden dim | 768 | 32 |
| heads / head_dim | 12 / 64 | 4 / 8 |
| params | ~110M | 28,994 |
| classification input | x[:, 0, :] = [CLS] | same |
Concept
Out loud, slides closed: (1) explain why BERT does not use a causal mask and what that means for the attention weight matrix; (2) describe the 80/10/10 MLM masking split and why each proportion exists; (3) trace the path from raw token IDs to a sentiment prediction through TinyBertForClassification.
Stretch (homework): load pretrained BERT from HuggingFace and fine-tune on SST-2 sentiment data comparing 100 vs 1000 vs 10000 examples; implement a token-classification head for NER instead of [CLS] classification; contrast fine-tuning (all weights) vs feature extraction (frozen encoder) on the same dataset. Next up: Lesson 89 — GPT, causal language modeling, and scaling laws.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — BERT architecture overview · Bidirectional attention · Pretraining: MLM & NSP · Fine-tuning: [CLS] classification head · Your turn: BertForClassification. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
[CLS] A [SEP] B [SEP] input formatTinyBertForClassification and fine-tune all weights with the Lesson 40 training loop| concept | the one thing to remember |
|---|---|
| bidirectionality | no upper-triangular mask; every token attends to every token |
| MLM split | 80% [MASK] / 10% random / 10% unchanged — prevents pretraining–finetuning mismatch |
| NSP input | [CLS] always at pos 0; segment IDs = 0 (A) / 1 (B) |
| fine-tuning | add Linear(768, k) on [CLS]; update ALL weights at ~2e-5 lr |
| BERT vs GPT | encoder (understand) vs decoder (generate); same attention math, different mask |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.