Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Fine-Tuning BERT and GPT Models
Here is a training log that looks like success and is not:
epoch 1 train_loss 0.412 train_acc 0.834 val_acc 0.821epoch 2 train_loss 0.153 train_acc 0.951 val_acc 0.812epoch 3 train_loss 0.041 train_acc 0.994 val_acc 0.786epoch 4 train_loss 0.008 train_acc 0.998 val_acc 0.771Training loss fell by a factor of fifty. Validation accuracy peaked at epoch 1 and lost five points after that. The best model was the one after a single pass over 800 examples, and everything since has been memorisation.
This is the characteristic shape of fine-tuning gone wrong, and it happens because fine-tuning is a fundamentally different activity from training. You are not teaching a model language — that already happened, at a cost of thousands of GPU-days. You are nudging 110 million already-good parameters towards a specific task using a few thousand examples. The techniques that work when training from scratch, chiefly "train longer with a healthy learning rate", actively destroy what you started with.
Fine-tuning BERT for classification
The head
BERT-base outputs a 768-dimensional vector per token. For sequence classification you take the [CLS] position — the leading special token that, through twelve layers of bidirectional attention, has become a learned summary of the whole input — and put a linear layer on it.
1import torch2import torch.nn as nn3from transformers import AutoModel, AutoTokenizer45class BertClassifier(nn.Module):6 def __init__(self, model_name="bert-base-uncased", num_labels=2, dropout=0.1):7 super().__init__()8 self.bert = AutoModel.from_pretrained(model_name)9 self.dropout = nn.Dropout(dropout)10 self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels)1112 def forward(self, input_ids, attention_mask, token_type_ids=None):13 out = self.bert(input_ids=input_ids,14 attention_mask=attention_mask,15 token_type_ids=token_type_ids)16 cls = out.last_hidden_state[:, 0] # the [CLS] position17 return self.classifier(self.dropout(cls))Count what you have added. For binary classification the head is 768×2+2=1,538 parameters against BERT's 109,482,240 — 0.0014% of the model. Nearly everything the classifier will know came from pre-training; you are learning a readout, not a capability.
One subtlety about pooler_output, which Hugging Face also returns. That is the [CLS] vector pushed through a tanh-activated dense layer trained on next-sentence prediction — an objective later work found contributes little. Using last_hidden_state[:, 0] directly is generally equal or better and is what most fine-tuning code does.
Preparing data
1from torch.utils.data import Dataset, DataLoader23class TextDataset(Dataset):4 def __init__(self, texts, labels, tokenizer, max_len=128):5 self.texts, self.labels = texts, labels6 self.tok, self.max_len = tokenizer, max_len78 def __len__(self):9 return len(self.texts)1011 def __getitem__(self, i):12 enc = self.tok(13 str(self.texts[i]),14 truncation=True,15 max_length=self.max_len,16 padding="max_length",17 return_tensors="pt",18 )19 return {20 "input_ids": enc["input_ids"].squeeze(0),21 "attention_mask": enc["attention_mask"].squeeze(0),22 "labels": torch.tensor(self.labels[i], dtype=torch.long),23 }Choosing max_len is a real decision, not a default to accept. Attention cost grows with the square of sequence length, so 128 tokens costs a quarter of what 256 does. Measure your data first:
1lengths = [len(tokenizer.encode(t)) for t in texts]2import numpy as np3print(np.percentile(lengths, [50, 90, 95, 99]))4# e.g. [ 41. 98. 127. 214.]With that distribution, max_len=128 keeps 95% of documents intact and truncates the rest. Setting it to 512 to be safe would quadruple your training time to protect 5% of examples. Setting it to 64 would silently cut a quarter of your data in half — and truncation is from the end, so the conclusion of every review disappears, which for sentiment is exactly the wrong half.
Prefer dynamic padding over padding="max_length" when you can. Padding each batch to its own longest sequence rather than a global maximum typically cuts training time by 30–50% on data with variable lengths:
from transformers import DataCollatorWithPaddingcollator = DataCollatorWithPadding(tokenizer) # pads per batchloader = DataLoader(dataset, batch_size=16, shuffle=True, collate_fn=collator)The training loop
1from transformers import get_linear_schedule_with_warmup2import torch.nn.functional as F34def fine_tune(model, train_loader, val_loader, epochs=3, lr=2e-5,5 warmup_ratio=0.1, device="cuda"):6 model.to(device)78 # No weight decay on biases or LayerNorm parameters — decaying them9 # shifts the normalisation statistics and hurts.10 decay, no_decay = [], []11 for name, p in model.named_parameters():12 if not p.requires_grad:13 continue14 (no_decay if any(k in name for k in ["bias", "LayerNorm.weight"])15 else decay).append(p)1617 optimizer = torch.optim.AdamW(18 [{"params": decay, "weight_decay": 0.01},19 {"params": no_decay, "weight_decay": 0.0}],20 lr=lr, eps=1e-8)2122 total_steps = len(train_loader) * epochs23 scheduler = get_linear_schedule_with_warmup(24 optimizer,25 num_warmup_steps=int(total_steps * warmup_ratio),26 num_training_steps=total_steps)2728 best_acc, best_state = 0.0, None2930 for epoch in range(epochs):31 model.train()32 for batch in train_loader:33 batch = {k: v.to(device) for k, v in batch.items()}34 logits = model(batch["input_ids"], batch["attention_mask"])35 loss = F.cross_entropy(logits, batch["labels"])3637 optimizer.zero_grad(set_to_none=True)38 loss.backward()39 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)40 optimizer.step()41 scheduler.step()4243 acc = evaluate(model, val_loader, device)44 print(f"epoch {epoch+1} val_acc {acc:.4f}")45 if acc > best_acc:46 best_acc = acc47 best_state = {k: v.detach().cpu().clone()48 for k, v in model.state_dict().items()}4950 model.load_state_dict(best_state) # keep the best epoch, not the last51 return model, best_accThat final load_state_dict is the direct answer to the opening log. Validation accuracy peaked at epoch 1; without checkpointing you would ship epoch 4 and be five points worse for no reason.
Fine-tuning GPT
For classification
A causal model has no [CLS] token, and even if you prepended one it would be useless: with causal masking, position 0 can attend only to itself, so it sees nothing. The information flows the other way. The last non-padding token is the only position that has seen the whole input.
1class GPTClassifier(nn.Module):2 def __init__(self, model_name="gpt2", num_labels=2):3 super().__init__()4 self.gpt = AutoModel.from_pretrained(model_name)5 self.classifier = nn.Linear(self.gpt.config.n_embd, num_labels)67 def forward(self, input_ids, attention_mask):8 h = self.gpt(input_ids=input_ids,9 attention_mask=attention_mask).last_hidden_state1011 # index of the final real token in each row12 last_idx = attention_mask.sum(dim=1) - 1 # (B,)13 pooled = h[torch.arange(h.size(0), device=h.device), last_idx]14 return self.classifier(pooled)Take the first token instead and you get a classifier reading a representation built from one word. Take a mean over all positions and it works, but worse — early positions have seen almost nothing, so you are averaging in mostly-uninformed vectors.
One practical wrinkle: GPT-2 has no padding token. You must assign one, and you must pad on the left if you want the last position to be uniform across the batch — or, as above, index the last real token explicitly and pad right.
tokenizer = AutoTokenizer.from_pretrained("gpt2")tokenizer.pad_token = tokenizer.eos_token # reuse EOS as PADmodel.config.pad_token_id = tokenizer.pad_token_idFor generation
Fine-tuning a generative model on your own text uses the same objective it was pre-trained with — predict the next token — on your data.
1from transformers import AutoModelForCausalLM23model = AutoModelForCausalLM.from_pretrained("gpt2")45def collate_lm(batch, tokenizer, max_len=256):6 enc = tokenizer([b["text"] for b in batch], truncation=True,7 max_length=max_len, padding=True, return_tensors="pt")8 labels = enc["input_ids"].clone()9 labels[enc["attention_mask"] == 0] = -100 # ignore padding in the loss10 enc["labels"] = labels11 return enc1213# The model shifts internally: it compares logits[:, :-1] with labels[:, 1:]14out = model(**batch)15loss = out.lossTwo things trip people here. The -100 sentinel is what PyTorch's cross-entropy ignores by default; without it a large share of your loss is the model being graded on predicting padding. And the shift is done inside Hugging Face's causal LM classes — if you write your own loss you must shift, and if you shift on top of theirs you are off by two.
For instruction-style data, where the prompt is given and only the response should be learned, mask the prompt tokens in the labels as well:
prompt_len = len(tokenizer(prompt)["input_ids"])labels[:prompt_len] = -100 # loss only on the responseHyperparameters
| Parameter | Range that works | Why |
|---|---|---|
| Learning rate | 2×10−5 to 5×10−5 | Roughly 50× smaller than training from scratch. At 10−3 the pre-trained weights are destroyed in the first hundred steps |
| Batch size | 16 or 32 | Larger needs a proportionally larger learning rate; smaller is noisy but often fine |
| Epochs | 2 to 4 | Beyond this you are almost always memorising. If more helps, your dataset is large enough that you should question whether it is really fine-tuning |
| Warmup | 6–10% of total steps | The randomly initialised head produces large early gradients that would otherwise wreck the encoder |
| Weight decay | 0.01, excluding bias and LayerNorm | Decaying LayerNorm weights towards zero distorts the normalisation the model relies on |
| Max grad norm | 1.0 | Caps the damage from a single bad batch |
| Dropout | 0.1 (the pre-training value) | Raise to 0.2–0.3 only if you are clearly overfitting |
Why warmup specifically
At step 0 the classification head is random. Its output is uninformative, so the loss is high and the gradient flowing back into BERT is large and points in an essentially arbitrary direction. Applying a full-size update at that moment moves 110 million carefully-trained parameters towards noise.
Warmup makes the first updates tiny. Concretely, with 4,000 training examples, batch 16 and 3 epochs:
At step 1 the learning rate is 751×2×10−5=2.7×10−7 — essentially zero. By step 75 it has climbed to the full 2×10−5, and by then the head has learned enough to produce a sensible gradient direction. From step 75 it decays linearly to zero at step 750.
Warmup is not about the encoder; it is about protecting the encoder from the random head sitting on top of it.
A related and underused technique is layer-wise learning rate decay. Lower layers of a pre-trained model encode general syntax that your task almost certainly does not need to change; upper layers encode task-relevant abstractions that do. Scale the learning rate by a factor per layer going down — with a factor of 0.95 across 12 layers, layer 0 gets 0.9512=0.540 of the top rate.
1def layerwise_lr_groups(model, base_lr=2e-5, decay=0.95):2 groups = [{"params": model.classifier.parameters(), "lr": base_lr}]3 layers = model.bert.encoder.layer4 for i, layer in enumerate(reversed(layers)): # top layer first5 groups.append({"params": layer.parameters(),6 "lr": base_lr * (decay ** (i + 1))})7 groups.append({"params": model.bert.embeddings.parameters(),8 "lr": base_lr * (decay ** (len(layers) + 1))})9 return groupsWhat goes wrong, and how to tell
Overfitting on a small dataset
The opening log is the signature: training accuracy marching to 0.99 while validation peaks early and declines. With 800 examples and 110 million parameters, memorisation is not a risk — it is the default outcome.
| Remedy | Effect | Cost |
|---|---|---|
| Fewer epochs (2 instead of 4) | Often the whole fix | None |
| Early stopping on validation | Same, done automatically | Needs a real validation split |
| Smaller model (DistilBERT) | 66 M parameters, ~97% of BERT's quality | Slightly lower ceiling |
| Freeze the bottom k layers | Fewer trainable parameters, faster, less to overfit with | Lower ceiling if the domain is far from pre-training |
| Higher dropout (0.2–0.3) | Modest regularisation | Slower convergence |
| Parameter-efficient methods (below) | Trains 0.3% of parameters; strongly resistant to overfitting | An extra dependency |
1# Freeze embeddings and the bottom 6 encoder layers2for p in model.bert.embeddings.parameters():3 p.requires_grad = False4for layer in model.bert.encoder.layer[:6]:5 for p in layer.parameters():6 p.requires_grad = False78trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)9print(f"trainable: {trainable:,}") # ~43M instead of ~110MInstability
Loss oscillates, or one seed reaches 0.94 and another 0.71 with identical settings. This is well documented for BERT fine-tuning on small datasets.
| Cause | Fix |
|---|---|
| Learning rate too high | Drop from 5×10−5 to 2×10−5 |
| No warmup | Add 10% warmup |
| Unclipped gradients | clip_grad_norm_(..., 1.0) |
| Wrong Adam epsilon | Use 10−8; PyTorch's default of 10−8 is fine, but some configurations carry 10−6 over from elsewhere |
| Genuine seed sensitivity | Run 3–5 seeds and report the mean and spread. A single number from a single seed is not a result |
Both training and validation are poor
This is underfitting, and it needs the opposite treatment. Check in this order:
- Are the labels correct? Print twenty examples with their labels and read them. Off-by-one label mappings and shuffled columns are more common than anyone admits.
- Can it overfit 100 examples? Train on a tiny subset with no regularisation. If the loss will not go near zero, you have a wiring bug, not a learning problem.
- Is the sequence length adequate? Compare the token-length percentiles against
max_len. - Is the learning rate too low? 10−6 barely moves anything in three epochs.
- Is the domain too far from pre-training? General BERT on clinical notes underperforms a domain-adapted checkpoint by a wide margin.
Parameter-efficient fine-tuning
Full fine-tuning updates every parameter, which means every task gets its own complete copy of the model. Ten tasks, ten times 440 MB. And the optimiser state for AdamW is two additional float32 tensors per parameter:
| Item (BERT-base, 110 M parameters) | Memory |
|---|---|
| Weights (fp32) | 440 MB |
| Gradients | 440 MB |
| Adam first moment | 440 MB |
| Adam second moment | 440 MB |
| Total before activations | 1.76 GB |
LoRA (low-rank adaptation) attacks this. The observation is that the update a fine-tune applies to a weight matrix is empirically low-rank — it does not need all 768×768 degrees of freedom. So freeze W0 and learn a factored correction:
With rank r=8, the trainable part of one 768×768 matrix is
against 589,824 — 2.08%. Apply it to the query and value projections of all 12 layers and you train
which is 0.27% of BERT-base. The optimiser state shrinks from 880 MB to about 2.4 MB. At inference you can fold BA back into W0, so there is no added latency — unlike adapters, which insert extra layers into the forward pass.
1from peft import LoraConfig, get_peft_model, TaskType23config = LoraConfig(4 task_type=TaskType.SEQ_CLS,5 r=8, # rank6 lora_alpha=16, # scaling: the update is multiplied by alpha/r7 lora_dropout=0.1,8 target_modules=["query", "value"],9)1011model = get_peft_model(base_model, config)12model.print_trainable_parameters()13# trainable params: 296,450 || all params: 109,780,228 || trainable%: 0.2700| Method | Trainable | Inference cost | Notes |
|---|---|---|---|
| Full fine-tuning | 100% | Baseline | Highest ceiling; one full copy per task |
| LoRA | 0.1–1% | None once merged | The default choice; matches full fine-tuning on most tasks |
| Adapters | 0.5–5% | Small but real — extra layers stay in the graph | Predates LoRA; largely superseded by it |
| Prefix / prompt tuning | <0.1% | Consumes context length | Effective mainly on very large models |
| BitFit | ~0.09% (biases only) | None | Surprisingly competitive on small datasets; a strong sanity baseline |
A note on lora_alpha, which is routinely misunderstood: the update is scaled by α/r. With α=16 and r=8 that factor is 2. If you raise r to 32 and leave α at 16, the effective scale drops to 0.5 and the adaptation becomes much weaker — people then conclude "higher rank did not help". Scale α with r.
What this means when you build something
Run this sequence in order. Each step is cheap and each one rules out a class of failure.
| Step | Action | What it proves |
|---|---|---|
| 1 | Print 20 examples with labels and read them | The data pipeline is correct |
| 2 | Record the majority-class baseline accuracy | What "good" has to beat. On 80/20 data, 0.80 accuracy is worthless |
| 3 | Overfit 100 examples to near-zero loss | Model, loss and optimiser are wired correctly |
| 4 | Fine-tune with the defaults: 2×10−5, batch 16, 3 epochs, 10% warmup | A real baseline number |
| 5 | Repeat with 3 seeds | Whether your improvements are real or noise |
| 6 | Only then tune, and change one thing at a time | Attribution |
Step 2 is skipped constantly and it is the one that most often changes the conclusion. A fraud classifier reporting 97% accuracy on data that is 97% legitimate has learned to say "legitimate". Report macro-F1 or per-class recall on imbalanced data, never accuracy alone.
On method selection: use full fine-tuning when you have a single task, plenty of GPU memory and more than a few thousand examples. Use LoRA when you have several tasks to serve from one base model, when memory is tight, or when your dataset is small enough that overfitting is the main risk — training 0.3% of the parameters is itself a powerful regulariser. Use a frozen encoder with a trained head only as a fast baseline; it is almost always a point or two behind.
And keep the number of trainable parameters in view relative to your dataset size. Two thousand examples against 110 million trainable parameters is a ratio of about 1 to 55,000. That is not a training problem to be solved with a better schedule — it is a reason to train fewer parameters.