Attention Mechanisms and Transformers

Fine-Tuning BERT and GPT Models


Here is a training log that looks like success and is not:

Text
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.771

Training 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.

Warmup then decay, learning rate in millionths02017131073001234567peak 2e-5after 10 pctend at zeroThe first steps run at a near-zero rate while Adam's moment estimates are still noise.
Warmup exists because the freshly attached head sends a large, meaningless gradient into pretrained weights.

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.

Python
import torchimport torch.nn as nnfrom transformers import AutoModel, AutoTokenizerclass BertClassifier(nn.Module):    def __init__(self, model_name="bert-base-uncased", num_labels=2, dropout=0.1):        super().__init__()        self.bert = AutoModel.from_pretrained(model_name)        self.dropout = nn.Dropout(dropout)        self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels)    def forward(self, input_ids, attention_mask, token_type_ids=None):        out = self.bert(input_ids=input_ids,                        attention_mask=attention_mask,                        token_type_ids=token_type_ids)        cls = out.last_hidden_state[:, 0]          # the [CLS] position        return self.classifier(self.dropout(cls))

Count what you have added. For binary classification the head is 768×2+2=1,538768 \times 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

Python
from torch.utils.data import Dataset, DataLoaderclass TextDataset(Dataset):    def __init__(self, texts, labels, tokenizer, max_len=128):        self.texts, self.labels = texts, labels        self.tok, self.max_len = tokenizer, max_len    def __len__(self):        return len(self.texts)    def __getitem__(self, i):        enc = self.tok(            str(self.texts[i]),            truncation=True,            max_length=self.max_len,            padding="max_length",            return_tensors="pt",        )        return {            "input_ids":      enc["input_ids"].squeeze(0),            "attention_mask": enc["attention_mask"].squeeze(0),            "labels":         torch.tensor(self.labels[i], dtype=torch.long),        }

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:

Python
lengths = [len(tokenizer.encode(t)) for t in texts]import numpy as npprint(np.percentile(lengths, [50, 90, 95, 99]))# 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:

Python
from transformers import DataCollatorWithPaddingcollator = DataCollatorWithPadding(tokenizer)   # pads per batchloader = DataLoader(dataset, batch_size=16, shuffle=True, collate_fn=collator)

The training loop

Python
from transformers import get_linear_schedule_with_warmupimport torch.nn.functional as Fdef fine_tune(model, train_loader, val_loader, epochs=3, lr=2e-5,              warmup_ratio=0.1, device="cuda"):    model.to(device)    # No weight decay on biases or LayerNorm parameters — decaying them    # shifts the normalisation statistics and hurts.    decay, no_decay = [], []    for name, p in model.named_parameters():        if not p.requires_grad:            continue        (no_decay if any(k in name for k in ["bias", "LayerNorm.weight"])         else decay).append(p)    optimizer = torch.optim.AdamW(        [{"params": decay, "weight_decay": 0.01},         {"params": no_decay, "weight_decay": 0.0}],        lr=lr, eps=1e-8)    total_steps = len(train_loader) * epochs    scheduler = get_linear_schedule_with_warmup(        optimizer,        num_warmup_steps=int(total_steps * warmup_ratio),        num_training_steps=total_steps)    best_acc, best_state = 0.0, None    for epoch in range(epochs):        model.train()        for batch in train_loader:            batch = {k: v.to(device) for k, v in batch.items()}            logits = model(batch["input_ids"], batch["attention_mask"])            loss = F.cross_entropy(logits, batch["labels"])            optimizer.zero_grad(set_to_none=True)            loss.backward()            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)            optimizer.step()            scheduler.step()        acc = evaluate(model, val_loader, device)        print(f"epoch {epoch+1}  val_acc {acc:.4f}")        if acc > best_acc:            best_acc = acc            best_state = {k: v.detach().cpu().clone()                          for k, v in model.state_dict().items()}    model.load_state_dict(best_state)     # keep the best epoch, not the last    return model, best_acc

That 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.

Python
class GPTClassifier(nn.Module):    def __init__(self, model_name="gpt2", num_labels=2):        super().__init__()        self.gpt = AutoModel.from_pretrained(model_name)        self.classifier = nn.Linear(self.gpt.config.n_embd, num_labels)    def forward(self, input_ids, attention_mask):        h = self.gpt(input_ids=input_ids,                     attention_mask=attention_mask).last_hidden_state        # index of the final real token in each row        last_idx = attention_mask.sum(dim=1) - 1                   # (B,)        pooled = h[torch.arange(h.size(0), device=h.device), last_idx]        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.

Python
tokenizer = AutoTokenizer.from_pretrained("gpt2")tokenizer.pad_token = tokenizer.eos_token       # reuse EOS as PADmodel.config.pad_token_id = tokenizer.pad_token_id

For 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.

Python
from transformers import AutoModelForCausalLMmodel = AutoModelForCausalLM.from_pretrained("gpt2")def collate_lm(batch, tokenizer, max_len=256):    enc = tokenizer([b["text"] for b in batch], truncation=True,                    max_length=max_len, padding=True, return_tensors="pt")    labels = enc["input_ids"].clone()    labels[enc["attention_mask"] == 0] = -100      # ignore padding in the loss    enc["labels"] = labels    return enc# The model shifts internally: it compares logits[:, :-1] with labels[:, 1:]out = model(**batch)loss = out.loss

Two 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:

Python
prompt_len = len(tokenizer(prompt)["input_ids"])labels[:prompt_len] = -100          # loss only on the response

Hyperparameters

ParameterRange that worksWhy
Learning rate2×10−52\times10^{-5} to 5×10−55\times10^{-5}Roughly 50× smaller than training from scratch. At 10−310^{-3} the pre-trained weights are destroyed in the first hundred steps
Batch size16 or 32Larger needs a proportionally larger learning rate; smaller is noisy but often fine
Epochs2 to 4Beyond this you are almost always memorising. If more helps, your dataset is large enough that you should question whether it is really fine-tuning
Warmup6–10% of total stepsThe randomly initialised head produces large early gradients that would otherwise wreck the encoder
Weight decay0.01, excluding bias and LayerNormDecaying LayerNorm weights towards zero distorts the normalisation the model relies on
Max grad norm1.0Caps the damage from a single bad batch
Dropout0.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:

steps per epoch=400016=250,total steps=750,warmup=0.1×750=75\text{steps per epoch} = \frac{4000}{16} = 250, \qquad \text{total steps} = 750, \qquad \text{warmup} = 0.1 \times 750 = 75

At step 1 the learning rate is 175×2×10−5=2.7×10−7\frac{1}{75} \times 2\times10^{-5} = 2.7\times10^{-7} — essentially zero. By step 75 it has climbed to the full 2×10−52\times10^{-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.5400.95^{12} = 0.540 of the top rate.

Python
def layerwise_lr_groups(model, base_lr=2e-5, decay=0.95):    groups = [{"params": model.classifier.parameters(), "lr": base_lr}]    layers = model.bert.encoder.layer    for i, layer in enumerate(reversed(layers)):        # top layer first        groups.append({"params": layer.parameters(),                       "lr": base_lr * (decay ** (i + 1))})    groups.append({"params": model.bert.embeddings.parameters(),                   "lr": base_lr * (decay ** (len(layers) + 1))})    return groups

What 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.

RemedyEffectCost
Fewer epochs (2 instead of 4)Often the whole fixNone
Early stopping on validationSame, done automaticallyNeeds a real validation split
Smaller model (DistilBERT)66 M parameters, ~97% of BERT's qualitySlightly lower ceiling
Freeze the bottom kk layersFewer trainable parameters, faster, less to overfit withLower ceiling if the domain is far from pre-training
Higher dropout (0.2–0.3)Modest regularisationSlower convergence
Parameter-efficient methods (below)Trains 0.3% of parameters; strongly resistant to overfittingAn extra dependency
Python
# Freeze embeddings and the bottom 6 encoder layersfor p in model.bert.embeddings.parameters():    p.requires_grad = Falsefor layer in model.bert.encoder.layer[:6]:    for p in layer.parameters():        p.requires_grad = Falsetrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)print(f"trainable: {trainable:,}")     # ~43M instead of ~110M

Instability

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.

CauseFix
Learning rate too highDrop from 5×10−55\times10^{-5} to 2×10−52\times10^{-5}
No warmupAdd 10% warmup
Unclipped gradientsclip_grad_norm_(..., 1.0)
Wrong Adam epsilonUse 10−810^{-8}; PyTorch's default of 10−810^{-8} is fine, but some configurations carry 10−610^{-6} over from elsewhere
Genuine seed sensitivityRun 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:

  1. 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.
  2. 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.
  3. Is the sequence length adequate? Compare the token-length percentiles against max_len.
  4. Is the learning rate too low? 10−610^{-6} barely moves anything in three epochs.
  5. 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
Gradients440 MB
Adam first moment440 MB
Adam second moment440 MB
Total before activations1.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×768768 \times 768 degrees of freedom. So freeze W0W_0 and learn a factored correction:

W=W0+BA,B∈R768×r, A∈Rr×768W = W_0 + BA, \qquad B \in \mathbb{R}^{768 \times r},\ A \in \mathbb{R}^{r \times 768}

With rank r=8r = 8, the trainable part of one 768×768768\times768 matrix is

768×8+8×768=6,144+6,144=12,288768 \times 8 + 8 \times 768 = 6{,}144 + 6{,}144 = 12{,}288

against 589,824 — 2.08%. Apply it to the query and value projections of all 12 layers and you train

12×2×12,288=294,912 parameters12 \times 2 \times 12{,}288 = 294{,}912 \text{ parameters}

which is 0.27% of BERT-base. The optimiser state shrinks from 880 MB to about 2.4 MB. At inference you can fold BABA back into W0W_0, so there is no added latency — unlike adapters, which insert extra layers into the forward pass.

Python
from peft import LoraConfig, get_peft_model, TaskTypeconfig = LoraConfig(    task_type=TaskType.SEQ_CLS,    r=8,                       # rank    lora_alpha=16,             # scaling: the update is multiplied by alpha/r    lora_dropout=0.1,    target_modules=["query", "value"],)model = get_peft_model(base_model, config)model.print_trainable_parameters()# trainable params: 296,450 || all params: 109,780,228 || trainable%: 0.2700
MethodTrainableInference costNotes
Full fine-tuning100%BaselineHighest ceiling; one full copy per task
LoRA0.1–1%None once mergedThe default choice; matches full fine-tuning on most tasks
Adapters0.5–5%Small but real — extra layers stay in the graphPredates LoRA; largely superseded by it
Prefix / prompt tuning<0.1%Consumes context lengthEffective mainly on very large models
BitFit~0.09% (biases only)NoneSurprisingly competitive on small datasets; a strong sanity baseline

A note on lora_alpha, which is routinely misunderstood: the update is scaled by α/r\alpha/r. With α=16\alpha = 16 and r=8r = 8 that factor is 2. If you raise rr to 32 and leave α\alpha at 16, the effective scale drops to 0.5 and the adaptation becomes much weaker — people then conclude "higher rank did not help". Scale α\alpha with rr.

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.

StepActionWhat it proves
1Print 20 examples with labels and read themThe data pipeline is correct
2Record the majority-class baseline accuracyWhat "good" has to beat. On 80/20 data, 0.80 accuracy is worthless
3Overfit 100 examples to near-zero lossModel, loss and optimiser are wired correctly
4Fine-tune with the defaults: 2×10−52\times10^{-5}, batch 16, 3 epochs, 10% warmupA real baseline number
5Repeat with 3 seedsWhether your improvements are real or noise
6Only then tune, and change one thing at a timeAttribution

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.