Transfer Learning and Pretraining

Text Transfer Learning — BERT, DistilBERT, and RoBERTa


A team fine-tunes BERT to classify 3,000 product reviews as positive or negative. They use the optimiser settings they use for everything else: AdamW at lr=1e-3. Training runs, the loss drops from 0.71 to 0.69 and then sits there like a stone. Final accuracy on a balanced test set: 50.2%. A coin does better than that on a good day.

They change one number. lr=2e-5 — fifty times smaller. Same model, same data, same three epochs.

92.4%.

That single number is the difference between a model that has learned the task and a model whose pretrained weights were scrambled beyond recovery in the first two hundred steps. Text transfer learning is unusually unforgiving about this, and the reason is worth understanding: a pretrained language encoder is a far more delicate object than a randomly initialised network, and treating it like one is how most first attempts fail.

Learning rates for fine-tuning, and the band that works1e-35e-41e-45e-53e-52e-5012345loss sitsat 0.69the usual answerAt 1e-3 the first steps overwrite the pretrained weights, and 50.2% is a coin flip on a balanced set.
Fine-tuning nudges weights that are already good — the default rate you use for a fresh model destroys them in fifty steps.

What BERT learned, and why it transfers so well

Before 2018, building an NLP system for a new task meant designing a task-specific architecture and collecting tens of thousands of labelled examples. BERT changed the economics completely: pretrain one encoder on billions of words of unlabelled text, then fine-tune it on a few thousand labelled examples per task.

Bidirectionality, and the trick that made it possible

Earlier language models read left to right, predicting each word from those before it. That is a real limitation. In "the bank of the river", the disambiguating evidence for bank arrives three words later; a left-to-right model has already committed.

The obvious fix — let every position see every other position — is circular, because predicting the next word while being allowed to look at it is trivial. BERT's solution is masked language modelling: corrupt 15% of the tokens, then predict the originals using the full bidirectional context.

The corruption scheme has a detail people gloss over. Of the 15% selected tokens, 80% are replaced with [MASK], 10% with a random token, and 10% are left unchanged. The reason is a train/deploy mismatch: [MASK] never appears during fine-tuning or inference, so a model trained only on masked positions learns representations that are only good at masked positions. The random and unchanged cases force it to build a useful representation for every token, because it cannot tell which ones it will be scored on.

BERT also trained on next sentence prediction — given two segments, are they consecutive? — intended to help with tasks involving sentence pairs. Later work found this objective contributes little and sometimes hurts.

PropertyBERT-baseBERT-large
Transformer layers1224
Hidden size7681024
Attention heads1216
Parameters110M340M
Max sequence length512 tokens512 tokens
Vocabulary30,522 WordPiece30,522 WordPiece

Pretraining used BooksCorpus (800 million words) plus English Wikipedia (2.5 billion words) — about 3.3 billion words, roughly 16 GB of text, with no human annotation at all. That is the crucial economic difference from image pretraining: ImageNet needed 1.28 million human labels, while masked language modelling supervises itself, because the masked word is already in the text.

Language pretraining scales because the labels are free. That single property is why text models reached hundreds of billions of parameters while labelled image datasets stayed in the millions.

DistilBERT: 97% of the quality at 60% of the cost

DistilBERT is BERT-base compressed by knowledge distillation: a smaller student is trained to reproduce the full model's output distribution rather than just the hard labels. The soft probabilities carry more information than a one-hot target — if BERT assigns 0.7 to "positive", 0.25 to "neutral" and 0.05 to "negative", the student learns that neutral is the near miss, which a hard label cannot express.

The student keeps every second layer of the teacher (6 instead of 12), is initialised from those layers rather than randomly, and trains on three combined losses: distillation loss against the teacher's soft outputs, ordinary masked-language-modelling loss, and a cosine embedding loss aligning the student's hidden states with the teacher's.

Result: 66M parameters instead of 110M, roughly 60% faster inference, retaining about 97% of BERT-base's score on standard benchmarks. For most production classification work that trade is straightforwardly correct.

RoBERTa: the same architecture, trained properly

RoBERTa changed nothing about BERT's architecture. It changed the training recipe, and gained several points. Four changes did the work:

  • Ten times more data — about 160 GB versus BERT's 16 GB.
  • Dynamic masking. BERT masked its corpus once during preprocessing, so a given sentence had the same masks in every epoch. RoBERTa re-masks on the fly, so the model never sees the same corruption twice.
  • Next sentence prediction removed. Dropping it improved results, and the freed capacity went into longer contiguous sequences.
  • Much larger batches — 8,000 sequences instead of 256, with the learning rate scaled accordingly.

RoBERTa also swapped WordPiece for a 50,265-token byte-level BPE vocabulary, which never produces unknown tokens because it can always fall back to bytes.

BERT-baseDistilBERTRoBERTa-base
Parameters110M66M125M
Layers12612
Pretraining data16 GB16 GB (via teacher)160 GB
MaskingStaticStaticDynamic
Next sentence predictionYesNoNo
Vocabulary30,522 WordPiece30,522 WordPiece50,265 byte-BPE
Relative inference speed1.0×1.6×0.95×
Typical accuracyBaseline−1 to −3 points+2 to +4 points
Choose it whenYou want the safe defaultLatency or cost mattersAccuracy matters most

Text classification: sentiment on 3,000 reviews

For classification, BERT prepends a special [CLS] token to every sequence. Its final-layer representation is trained to summarise the whole input, and a single linear layer on top of it produces the class logits. That linear layer is the only randomly initialised part of the model.

Python
from transformers import AutoTokenizer, AutoModelForSequenceClassificationfrom datasets import load_datasetimport numpy as np, evaluatecheckpoint = "distilbert-base-uncased"tokenizer = AutoTokenizer.from_pretrained(checkpoint)model = AutoModelForSequenceClassification.from_pretrained(    checkpoint, num_labels=2,    id2label={0: "negative", 1: "positive"},    label2id={"negative": 0, "positive": 1},)def tokenize(batch):    return tokenizer(batch["text"], truncation=True, max_length=256)ds = load_dataset("stanfordnlp/imdb").map(tokenize, batched=True)

Loading the checkpoint prints a warning that some weights are newly initialised. That warning is correct and expected — it is the classification head. If it ever says the encoder weights were newly initialised, your checkpoint name is wrong and you are about to fine-tune a random transformer.

Note max_length=256 rather than the maximum 512. Self-attention costs O(n2)O(n^2) in sequence length: at 512 tokens each attention head computes 5122=262,144512^2 = 262{,}144 pairwise scores per layer, while at 256 it computes 65,53665{,}536 — four times fewer. Truncating to the shortest length that still contains the signal is the cheapest speedup available. Check your actual token-length distribution before choosing; if 95% of your reviews fit in 192 tokens, using 512 wastes most of your compute on padding.

Dynamic padding makes this concrete: pad each batch to its own longest sequence rather than to a global maximum.

Python
from transformers import DataCollatorWithPadding, TrainingArguments, Trainercollator = DataCollatorWithPadding(tokenizer=tokenizer)accuracy = evaluate.load("accuracy")def compute_metrics(eval_pred):    logits, labels = eval_pred    return accuracy.compute(predictions=np.argmax(logits, axis=-1),                            references=labels)args = TrainingArguments(    output_dir="sentiment",    learning_rate=2e-5,              # NOT 1e-3    per_device_train_batch_size=16,    num_train_epochs=4,    weight_decay=0.01,    warmup_steps=0.1,                # a float below 1 is a fraction of all steps    lr_scheduler_type="linear",    eval_strategy="epoch",    save_strategy="epoch",    load_best_model_at_end=True,    metric_for_best_model="accuracy",    fp16=True,)trainer = Trainer(model=model, args=args,                  train_dataset=ds["train"], eval_dataset=ds["test"],                  data_collator=collator, compute_metrics=compute_metrics)trainer.train()

Then prediction, which is where a different mistake hides:

Python
import torchtext = "The battery lasts two days but the camera is disappointing."inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=256)model.eval()with torch.no_grad():    logits = model(**inputs).logitsprobs = torch.softmax(logits, dim=-1)[0]label = model.config.id2label[int(probs.argmax())]print(f"{label}  ({probs.max():.3f})")

Two things that go wrong here. Forgetting model.eval() leaves dropout active, so the same input gives different answers on different calls. And using a different tokenizer at inference than at training — even a different casing variant of the same model — silently produces token IDs the model has never associated with those words. Always load the tokenizer from the same checkpoint directory as the model.

Token classification: named entity recognition

Classification assigns one label to a whole sequence. Token classification assigns a label to every token — person, organisation, location, or none. And here lives the single most common bug in NLP fine-tuning.

The subword alignment trap

Your annotations are per word. BERT operates on subword pieces. With bert-base-cased, the word "Zuckerberg" tokenises as ["Z", "##uck", "##er", "##berg"], and "Anthropic" as ["An", "##throp", "##ic"]. If you have 9 words and 14 tokens, and you simply hand over your 9 labels, everything after the first multi-piece word is misaligned. Training proceeds without error and the F1 score comes out around 40%.

The fix uses the tokenizer's word_ids(), which maps each token back to the word it came from.

Python
def align_labels(examples, tokenizer, label_all_subwords=False):    tokenized = tokenizer(examples["tokens"], truncation=True,                          is_split_into_words=True, max_length=256)    all_labels = []    for i, word_labels in enumerate(examples["ner_tags"]):        word_ids = tokenized.word_ids(batch_index=i)        previous, labels = None, []        for word_id in word_ids:            if word_id is None:                labels.append(-100)               # [CLS], [SEP], padding            elif word_id != previous:                labels.append(word_labels[word_id])   # first piece of a word            else:                # Later pieces: ignore them, or repeat the label.                labels.append(word_labels[word_id] if label_all_subwords else -100)            previous = word_id        all_labels.append(labels)    tokenized["labels"] = all_labels    return tokenized

-100 is the value PyTorch's cross-entropy ignores, so those positions contribute nothing to the loss. Labelling only the first subword piece of each word is the standard choice, and it matters for evaluation too — you score one prediction per word, not per piece.

Entity labels use the BIO scheme: B-PER begins a person entity, I-PER continues it, O is outside any entity. The B/I distinction exists to separate adjacent entities of the same type — without it, "Alice Bob" reads as one two-word person rather than two.

Python
from transformers import AutoModelForTokenClassification, DataCollatorForTokenClassificationlabels = ["O", "B-PER", "I-PER", "B-ORG", "I-ORG", "B-LOC", "I-LOC"]tokenizer = AutoTokenizer.from_pretrained("bert-base-cased")   # same checkpoint as the modelmodel = AutoModelForTokenClassification.from_pretrained(    "bert-base-cased", num_labels=len(labels),    id2label=dict(enumerate(labels)),    label2id={l: i for i, l in enumerate(labels)},)collator = DataCollatorForTokenClassification(tokenizer=tokenizer)

Use a cased checkpoint for NER. Capitalisation is one of the strongest signals that a token is a proper noun, and bert-base-uncased throws it away before the model ever sees it. This single choice is typically worth two to four F1 points.

Evaluate with entity-level F1, not token accuracy. In a typical NER corpus around 85% of tokens are O, so a model that predicts O for everything scores 85% token accuracy and 0% entity F1. Report precision, recall and F1 per entity type using a sequence-labelling metric that requires the whole span to match.

Extractive question answering

Given a question and a passage, find the span of the passage that answers it. The model is not generating text — it is choosing two positions.

The input is the question and context concatenated: [CLS] question [SEP] context [SEP]. Two linear layers over the token representations produce a start logit and an end logit for every position. The predicted answer is the span maximising start_logit[i] + end_logit[j], subject to i≤ji \le j, a maximum answer length, and both positions lying inside the context rather than the question.

Python
from transformers import AutoModelForQuestionAnsweringimport torchtok = AutoTokenizer.from_pretrained("distilbert-base-cased-distilled-squad")model = AutoModelForQuestionAnswering.from_pretrained(    "distilbert-base-cased-distilled-squad")question = "How many parameters does BERT-base have?"context = ("BERT-base has 12 transformer layers, a hidden size of 768, "           "and 110 million parameters. BERT-large has 340 million.")enc = tok(question, context, return_tensors="pt",          return_offsets_mapping=True, truncation="only_second", max_length=384)offsets = enc.pop("offset_mapping")[0]with torch.no_grad():    out = model(**enc)start = out.start_logits[0]end = out.end_logits[0]ctx = [k for k, s in enumerate(enc.sequence_ids(0)) if s == 1]   # context tokens only# Score every valid (i, j) pair inside the context, then take the best.best, span = -1e9, (0, 0)for i in ctx:    for j in ctx:        if i <= j < i + 30:                         # cap answer length            score = start[i].item() + end[j].item()            if score > best:                best, span = score, (i, j)char_start = offsets[span[0]][0].item()char_end = offsets[span[1]][1].item()print(context[char_start:char_end])       # "110 million"

Two details make the difference between working and not. truncation="only_second" truncates the context and never the question — truncate the question and the model is answering something you did not ask. And offset_mapping records the character range each token came from, which is the only reliable way to recover the exact original text: detokenising by joining pieces loses the original spacing and punctuation.

For contexts longer than the model's limit, use a sliding window with return_overflowing_tokens=True and stride=128, then take the highest-scoring span across all windows. The stride ensures an answer straddling a window boundary appears intact in at least one window.

The optimisation settings that decide whether it works

Learning rate and batch size

The opening failure was entirely a learning-rate problem. The mechanism is precise: your classification head is randomly initialised, so its early gradients are large and meaningless, and they flow back into an encoder whose weights encode syntax and semantics learned from 3.3 billion words. At 1e-3 those updates overwrite the pretrained weights within a few hundred steps. The model then converges to predicting the class prior, which for a balanced binary task is exactly 50%.

Training examplesBatch sizeLearning rateEpochsWarmup
Under 1,0008–162e-55–1010%
1,000–10,00016–322e-5 to 3e-53–510%
10,000–100,000323e-536%
Over 100,00032–643e-5 to 5e-52–36%

The usable range for fine-tuning a pretrained encoder is roughly 1e-5 to 5e-5. Above 1e-4 you are destroying pretrained weights; below 5e-6 you will need far more epochs than the schedule allows. Note that small datasets get more epochs and a smaller learning rate — more passes, each gentler.

Gradient accumulation

Transformer activations dominate memory, and batch size is usually limited by the GPU rather than by what you want. Gradient accumulation simulates a larger batch by summing gradients over several forward/backward passes before stepping.

Python
args = TrainingArguments(    per_device_train_batch_size=8,     # what actually fits in memory    gradient_accumulation_steps=4,     # effective batch = 8 x 4 = 32    gradient_checkpointing=True,       # trade ~30% speed for ~60% less memory    fp16=True,    learning_rate=3e-5,)

The effective batch size is per_device_batch × accumulation_steps × number_of_devices. Two things to keep straight: the learning rate should be chosen for the effective batch, not the physical one; and a "step" in the scheduler means an optimiser step, so accumulation steps of 4 makes your total step count four times smaller — set warmup as a ratio rather than an absolute number, or you will accidentally warm up for the entire run.

Warmup and decay

Linear warmup followed by linear decay is the standard schedule, and both halves earn their place. Warmup exists because AdamW's second-moment estimates are based on almost no data in the first steps, so its updates are poorly scaled exactly when the encoder is most vulnerable. Decay exists because large steps late in training bounce around the minimum instead of settling into it.

Work an example. With 3,000 training examples, batch size 16 and 4 epochs: ⌈3000/16⌉=188\lceil 3000/16 \rceil = 188 steps per epoch, so 188×4=752188 \times 4 = 752 total steps. A 10% warmup ratio gives 75 steps ramping the learning rate from 0 to 2e-5, then 677 steps decaying it back to 0.

Python
from transformers import get_linear_schedule_with_warmuptotal_steps = len(train_dataloader) * num_epochs // accumulation_stepsscheduler = get_linear_schedule_with_warmup(    optimiser,    num_warmup_steps=int(0.1 * total_steps),    num_training_steps=total_steps,)

Skipping warmup on a pretrained encoder is the fastest way to lose several points for no reason. The first hundred steps are where the damage happens, and warmup is what stops it.

Failure modes you should be able to recognise on sight

SymptomCauseFix
Loss flat at ln⁡(num classes)\ln(\text{num classes}), accuracy at chanceLearning rate far too high; encoder destroyed earlyDrop to 2e-5, add 10% warmup
NER F1 around 40% with no error messageWord labels not aligned to subword tokensUse word_ids() and -100 for non-first pieces
Token accuracy 85%, entity F1 near 0Model predicts O everywhere; wrong metricEvaluate entity-level F1, not token accuracy
Different answers on repeated identical inputsmodel.eval() not called; dropout still activeCall model.eval() and wrap in torch.no_grad()
Training fine, inference nonsenseTokenizer mismatched to checkpointLoad tokenizer and model from the same source
Validation peaks at epoch 2 then falls steadilyOverfitting — normal for small datasetsload_best_model_at_end=True, fewer epochs
Out of memory at batch size 32Activation memory scales with batch × sequence lengthAccumulate gradients, shorten sequences, enable checkpointing

What this means when you build something

The practical shape of a text fine-tuning project is narrower than it looks, because most of the decisions have known good answers and only a few genuinely depend on your problem.

  1. Start with DistilBERT, not BERT-large. It trains in roughly half the time, so you get through five experiments while a large model finishes one. Move up only after your pipeline is proven correct — and measure whether the extra points justify the inference cost.
  2. Set the learning rate to 2e-5 before you do anything else. This is not a hyperparameter to explore first; it is a default to depart from cautiously. Explore batch size and epochs instead.
  3. Measure your token length distribution and truncate accordingly. Attention cost is quadratic, so halving the sequence length quarters the attention compute. Most datasets do not need 512.
  4. For token-level tasks, verify alignment by printing it. Print ten examples as (token, label) pairs and read them. Misalignment produces no error and costs 40 F1 points, and it is invisible in every metric until you look at the tokens themselves.
  5. Choose the checkpoint's casing to match the task. Cased for NER and anything where capitalisation carries meaning; uncased is fine for sentiment and topic classification, where it slightly reduces vocabulary sparsity.
  6. Prefer a domain-matched checkpoint over a bigger general one. A base-sized model pretrained on biomedical text or on legal documents routinely beats a general large model on those domains, because the vocabulary and the writing conventions are already familiar. Look for one before assuming you need more parameters.

The team from the opening eventually shipped DistilBERT at 91.8% — half a point below their fine-tuned BERT-base, at 1.6× the throughput and 60% of the memory. The fifty-fold learning-rate error cost them a week; understanding why it happened means it costs them nothing ever again.