BERT, explained · Part 5 of 6 · Covers §5, A.4, C.1, C.2

Ablations: What Really Matters

Section 5 and Appendix C of the BERT paper, line by line: which parts of BERT actually cause its gains. Next sentence prediction, bidirectionality, model size, frozen features versus fine-tuning, training length and the 80/10/10 masking recipe, every table row explained, with real parameter counts, a real perplexity and a real re-run of the feature-based NER experiment.

BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. Jacob Devlin, Ming-Wei Chang, Kenton Lee, Kristina Toutanova. NAACL 2019, 2018. arXiv:1810.04805

Part 4 showed that BERT wins almost everywhere. But a win does not tell you why. BERT changed many things at once compared with OpenAI GPT: two new pre-training tasks, bidirectional attention, more data, bigger batches, tuned learning rates. Which of these actually mattered?

This part reads the paper's answer: Section 5 (the ablation studies), plus Appendix A.4 and Appendix C, which hold the rest of the evidence. Every table row gets explained, and I check what can be checked in code.

What an ablation study is

The key rule of a good ablation is in the next screenshot: change one thing at a time.

Effect of the pre-training tasks

One detail in the screenshot deserves attention: for LTR & No NSP, "the left-only constraint was also applied at fine-tuning, because removing it introduced a pre-train/fine-tune mismatch that degraded downstream performance". In other words, you cannot take a model trained to look only left and suddenly let it look right during fine-tuning: it never learned what to do with that information, and it got worse, not better.

Each row of Table 5 changes one thing from the row beforeBERT-baseMLM + NSPdropNSPNo NSPMLM onlyMLM toLTRLTR & No NSPleft-to-right onlyadd aBiLSTM+ BiLSTMLTR, BiLSTM on topSame pre-training data, same fine-tuning recipe, same hyperparameters: only the named change differs.
The four models of Table 5. Each row changes exactly one thing from the row before, so each difference in scores can be blamed on that one change.

What differs between the four models is easiest to see in their attention masks and pre-training heads:

BERT-basemasked LM + NSP[CLS]mydogis[SEP][CLS][CLS] attends to [CLS][CLS] attends to my[CLS] attends to dog[CLS] attends to is[CLS] attends to [SEP]mymy attends to [CLS]my attends to mymy attends to dogmy attends to ismy attends to [SEP]dogdog attends to [CLS]dog attends to mydog attends to dogdog attends to isdog attends to [SEP]isis attends to [CLS]is attends to myis attends to dogis attends to isis attends to [SEP][SEP][SEP] attends to [CLS][SEP] attends to my[SEP] attends to dog[SEP] attends to is[SEP] attends to [SEP]who may look at whompre-training headsMLMNSP on CNo NSPmasked LM only[CLS]mydogis[SEP][CLS][CLS] attends to [CLS][CLS] attends to my[CLS] attends to dog[CLS] attends to is[CLS] attends to [SEP]mymy attends to [CLS]my attends to mymy attends to dogmy attends to ismy attends to [SEP]dogdog attends to [CLS]dog attends to mydog attends to dogdog attends to isdog attends to [SEP]isis attends to [CLS]is attends to myis attends to dogis attends to isis attends to [SEP][SEP][SEP] attends to [CLS][SEP] attends to my[SEP] attends to dog[SEP] attends to is[SEP] attends to [SEP]who may look at whompre-training headsMLMNSP removedLTR & No NSPleft-to-right LM, like GPT[CLS]mydogis[SEP][CLS][CLS] attends to [CLS]mymy attends to [CLS]my attends to mydogdog attends to [CLS]dog attends to mydog attends to dogisis attends to [CLS]is attends to myis attends to dogis attends to is[SEP][SEP] attends to [CLS][SEP] attends to my[SEP] attends to dog[SEP] attends to is[SEP] attends to [SEP]who may look at whompre-training headsnext wordleft-only mask keptduring fine-tuning too+ BiLSTMthe LTR model + a BiLSTM[CLS]mydogis[SEP][CLS][CLS] attends to [CLS]mymy attends to [CLS]my attends to mydogdog attends to [CLS]dog attends to mydog attends to dogisis attends to [CLS]is attends to myis attends to dogis attends to is[SEP][SEP] attends to [CLS][SEP] attends to my[SEP] attends to dog[SEP] attends to is[SEP] attends to [SEP]who may look at whompre-training headsnext wordBiLSTMnew, random startleft-only mask keptduring fine-tuning too
The four models of Table 5, side by side. BERT-base: no mask, MLM and NSP heads. No NSP: no mask, MLM head only. LTR & No NSP: a left-to-right mask (the upper triangle blocked, like GPT), next-word head, and the mask is kept during fine-tuning too. + BiLSTM: the LTR model with a new, randomly initialised BiLSTM added for fine-tuning.

Table 5, row by row

The five columns are tasks from Part 4:

ColumnTaskScore
MNLI-mdoes sentence B follow from A, contradict it, or neither? ("m" = matched: test sentences from the same genres as training)accuracy
QNLIdoes this sentence contain the answer to this question?accuracy
MRPCdo these two news sentences mean the same thing?accuracy
SST-2is this movie-review sentence positive or negative?accuracy
SQuADfind the answer span in a paragraph (v1.1)F1

Numbers in a table are easier to compare as differences. My script bert_part5.py subtracts each row from the row before (the Table 5 values are typed in from the paper):

plain text
Table 5 (Dev set): change in points
                                                MNLI-m    QNLI    MRPC   SST-2   SQuAD
cost of removing NSP                              -0.5    -3.5    -0.2    -0.1    -0.6
cost of left-to-right instead of MLM              -1.8    -0.6    -9.0    -0.5   -10.1
effect of adding a BiLSTM to the LTR model        +0.0    -0.2    -1.8    -0.5    +7.1
707580859095BERT-base, MNLI-m: 84.484.4No NSP, MNLI-m: 83.983.9LTR & No NSP, MNLI-m: 82.182.1+ BiLSTM, MNLI-m: 82.182.1MNLI-mBERT-base, QNLI: 88.488.4No NSP, QNLI: 84.984.9LTR & No NSP, QNLI: 84.384.3+ BiLSTM, QNLI: 84.184.1QNLIBERT-base, MRPC: 86.786.7No NSP, MRPC: 86.586.5LTR & No NSP, MRPC: 77.577.5+ BiLSTM, MRPC: 75.775.7MRPCBERT-base, SST-2: 92.792.7No NSP, SST-2: 92.692.6LTR & No NSP, SST-2: 92.192.1+ BiLSTM, SST-2: 91.691.6SST-2BERT-base, SQuAD F1: 88.588.5No NSP, SQuAD F1: 87.987.9LTR & No NSP, SQuAD F1: 77.877.8+ BiLSTM, SQuAD F1: 84.984.9SQuAD F1BERT-baseNo NSPLTR & No NSP+ BiLSTM
Table 5 as a chart (Dev set; the axis starts at 70 so the gaps are visible). Removing NSP mostly hurts QNLI. Switching to left-to-right crashes MRPC and SQuAD. The BiLSTM rescues part of SQuAD but nothing else.

Row 1 to row 2: removing NSP.

Looking at the computed differences, the NSP claim is weaker than the word "significantly" suggests. QNLI really drops (3.5 points), but MNLI drops 0.5 and SQuAD 0.6, and MRPC and SST-2 barely move (0.2 and 0.1). The paper does not report how much these Dev scores vary between runs, so we cannot tell how much of a 0.5-point gap is noise. Keep this in mind: in Part 6 we will see that later work (RoBERTa) questioned whether NSP is needed at all.

Row 2 to row 3: one direction instead of two. This is the big one. With everything else equal, going from the masked LM to a left-to-right LM costs 9.0 points on MRPC and 10.1 points of F1 on SQuAD, and loses something on every task. This is the cleanest evidence in the paper that the bidirectional design, not just more data or a bigger batch, explains much of BERT's gain.

Why does SQuAD suffer most? SQuAD needs a start and an end position for the answer (Part 4). Whether a word is the end of an answer depends heavily on the words after it. In an LTR model, the vector for each token was built without ever seeing those words. The BiLSTM adds right-side context, which is why SQuAD jumps by 7.1 points (77.8 to 84.9). But it is one small layer trained only on the task data, while BERT mixes both sides in all 12 layers during pre-training. The result is still 3.6 points below BERT-base (84.9 versus 88.5).

Question: who painted it? Passage: the mona lisa was painted by leonardo da vinciLeft-to-right model[CLS]whopaintedit?[SEP]themonalisawaspaintedbyleonardodavinci[SEP]hidden from "leonardo"Bidirectional model (BERT)[CLS]whopaintedit?[SEP]themonalisawaspaintedbyleonardodavinci[SEP]Teal arrows: right-side words. Only BERT can use them.Without "da vinci", the vector of "leonardo" cannot tell whether a longer name follows, or where the answer should end.The highlighted box is the token whose vector must say "the answer starts here". Green: the true answer span.
What the vector of the answer's first word can see. Question "who painted it?", passage "the mona lisa was painted by leonardo da vinci". In a left-to-right model, "leonardo" sees only what came before it; "da vinci" is hidden, so its vector cannot say where the answer ends. In BERT it sees both sides, in every layer.

On the GLUE tasks the BiLSTM does not help at all (MNLI unchanged, the others down by 0.2 to 1.8 points). A plausible reason, not tested in the paper: a randomly initialized layer has to be learned from small task datasets, and MRPC has only a few thousand training pairs (3,600, according to Section 5.2).

Why not just do what ELMo does?

Each argument in plain words:

  • (a) Cost. Two full models to pre-train and to run, instead of one.
  • (b) Question answering. In BERT's input the question comes first and the paragraph second (Part 4). A right-to-left model reading the paragraph has not reached the question yet (the question is to its left), so its vectors for the paragraph words know nothing about what is being asked.
  • (c) Depth. In ELMo the two directions meet only at the very end. In BERT, every layer mixes both sides, so later layers can build on information that already combines left and right. A small grammar slip in the paper can confuse readers here: in "this it is strictly less powerful ... since it can use both left and right context at every layer", the second "it" means the deep bidirectional model, not the concatenation.
Rows: passage tokens. Columns: question tokens. The question comes first in the input.left-to-rightkeywhopaintedit?paintedpainted attends to whopainted attends to paintedpainted attends to itpainted attends to ?byby attends to whoby attends to paintedby attends to itby attends to ?leonardoleonardo attends to wholeonardo attends to paintedleonardo attends to itleonardo attends to ?question is on the left: visibleright-to-leftkeywhopaintedit?paintedbyleonardoquestion is on the left: never visibleBERTkeywhopaintedit?paintedpainted attends to whopainted attends to paintedpainted attends to itpainted attends to ?byby attends to whoby attends to paintedby attends to itby attends to ?leonardoleonardo attends to wholeonardo attends to paintedleonardo attends to itleonardo attends to ?every token sees every token
Argument (b), as attention tables. Rows are passage tokens, columns question tokens. A left-to-right model lets passage tokens see the question (it comes first). A right-to-left model blocks all of it: the passage never sees the question. BERT allows everything.

Argument (a) is easy to measure. I built an LTR stack and an RTL stack of BERT-base size and timed them against one bidirectional model (bert_part5_math.py):

plain text
parameters: one bidirectional model 109,482,240; LTR + RTL = 218,964,480
vector per token: bidirectional (512, 768) ; LTR (512, 768) + RTL (512, 768) -> concat (512, 1536)
time for a batch of 8 x 512 tokens on mps (median of 5): one model 133 ms, LTR + RTL 261 ms -> 1.96x
in a sequence of 6 tokens, which positions can token 3 use (1-based)?
  bidirectional, any layer: [1, 2, 3, 4, 5, 6]
  LTR stack, any layer:     [1, 2, 3]    RTL stack, any layer: [3, 4, 5, 6]
  concatenation: both lists, but only side by side at the very top; no layer ever mixes them

Twice the parameters and 1.96 times the time, as the paper says. And the last lines are argument (c) in miniature: in the glued model, no single layer ever combines token 1 with token 6 when building token 3.

ELMo-style: two one-way models, glued at the topBERT: one model, both sides in every layerthekidsmilesLTR stackthekidsmilesRTL stacksmiles = [LTR ; RTL] = 1,536 numbersthekidsmilessmiles = 768 numbersMeasured, BERT-base size, batch of 8 x 512 tokens on an Apple GPU (bert_part5_math.py):LTR + RTL261 ms, 219.0M parametersone bidirectional133 ms, 109.5M parameters
ELMo-style versus BERT, with the measured cost. Left: an LTR stack and an RTL stack, each one-directional, glued at the top into 1,536 numbers per token. Right: one stack where every layer mixes both sides, 768 numbers per token. Glued: 219.0M parameters and 261 ms per batch; bidirectional: 109.5M and 133 ms.

How ELMo actually combines its layers, in its own paper:

where hk,jh_{k,j} is the (forward and backward, concatenated) vector of token kk at layer jj, and the task model learns the L+1L + 1 weights sjs_j and the single number γ\gamma; the LSTMs themselves stay frozen.

The other differences between BERT and GPT

Table 5's LTR & No NSP row matters for one more reason, which the paper explains in Appendix A.4.

OpenAI GPTBERT
Pre-training dataBooksCorpus, 800M wordsBooksCorpus 800M + Wikipedia 2,500M words
[SEP], [CLS], segment A/Badded only at fine-tuninglearned during pre-training
Batch size32,000 words, 1M steps128,000 words, 1M steps
Fine-tuning learning ratealways 5e-5best of several, chosen per task on Dev
Appendix A.4: the four non-architecture differences between OpenAI GPT and BERTOpenAI GPTBERTpre-training textBooksCorpus, 800M wordsBooksCorpus + Wikipedia, 3,300M wordsbatch size32,000 words per step, 1M steps128,000 words per step, 1M steps[CLS], [SEP], A/BGPT: added only at fine-tuningBERT: learned during pre-trainingfine-tuning rateGPT: always 5e-5BERT: best of 5e-5, 3e-5, 2e-5 on DevThe "LTR & No NSP" row of Table 5 = GPT's objective + all four BERT advantages above.So its gap to BERT-base (MRPC 77.5 vs 86.7, SQuAD 77.8 vs 88.5) comes from the masked LM and NSP.
The four non-architecture differences between OpenAI GPT and BERT, as bars and boxes: training text (800M against 3,300M words), batch size (32,000 against 128,000 words per step), when [CLS], [SEP] and A/B are learned, and the fine-tuning learning rate. The "LTR & No NSP" row of Table 5 is GPT's objective with all four BERT advantages.

The logic: the LTR & No NSP model is GPT's recipe (left to right, no NSP) but trained with all four of BERT's advantages (BERT's data, BERT's input format, BERT's batch size, BERT's fine-tuning). So any gap between LTR & No NSP and BERT-base cannot come from those four things. It must come from the masked LM, which gives bidirectionality, and NSP. And that gap is large: up to 9.2 points on MRPC (77.5 versus 86.7) and 10.7 on SQuAD (77.8 versus 88.5).

Effect of model size

A small inconsistency: the text says "all four datasets", but Table 6 shows three tasks (MNLI-m, MRPC, SST-2). The paper does not say what the fourth one is.

You met L, H and A in Part 2. The new column is perplexity.

For a masked LM, the perplexity is computed over the positions the model has to predict:

ppl=exp⁡ ⁣(−1N∑i=1Nlog⁡p(xi∣context))\text{ppl} = \exp\!\left( -\frac{1}{N} \sum_{i=1}^{N} \log p(x_i \mid \text{context}) \right)

where:

  • NN is the number of predicted (masked) positions;
  • xix_i is the true token at the ii-th predicted position;
  • p(xi∣context)p(x_i \mid \text{context}) is the probability the model gave to that true token, from its softmax over the vocabulary;
  • log⁡\log is the natural logarithm, and the minus sign makes the sum positive: the inside of the bracket is the average cross-entropy loss (Part 3);
  • exp⁡\exp ("e to the power of") undoes the logarithm, turning an average loss into a number of "equally likely choices".

So BERT-base's 3.99 means: on its held-out training text, when it predicts a masked token, it is on average about as unsure as a choice between 4 words.

How many parameters is each row?

Where the per-layer term 12H2+13H12H^2 + 13H comes from, piece by piece (all for one layer with hidden size HH and feed-forward size 4H4H):

PartMatricesBiases and LayerNormsCount
Q, K, V, O projections4×H×H4 \times H \times H4H4H4H2+4H4H^2 + 4H
Feed-forwardH×4H+4H×HH \times 4H + 4H \times H4H+H4H + H8H2+5H8H^2 + 5H
Two LayerNorms2×2H2 \times 2H4H4H
One layer12H2+13H12H^2 + 13H

For H=768H = 768: 12×7682+13×768=7,077,888+9,984=7,087,87212 \times 768^2 + 13 \times 768 = 7{,}077{,}888 + 9{,}984 = 7{,}087{,}872 parameters per layer.

Where the parameters of each Table 6 model live (counted, term by term)embeddingsattentionfeed-forwardLayerNorm + poolerL=3, H=768, A=12embeddings: 23,837,184attention: 7,087,104feed-forward: 14,167,296LayerNorm + pooler: 599,80845.7ML=6, H=768, A=3embeddings: 23,837,184attention: 14,174,208feed-forward: 28,334,592LayerNorm + pooler: 609,02467.0ML=6, H=768, A=12embeddings: 23,837,184attention: 14,174,208feed-forward: 28,334,592LayerNorm + pooler: 609,02467.0ML=12, H=768, A=12embeddings: 23,837,184attention: 28,348,416feed-forward: 56,669,184LayerNorm + pooler: 627,456109.5ML=12, H=1024, A=16embeddings: 31,782,912attention: 50,380,800feed-forward: 100,724,736LayerNorm + pooler: 1,098,752184.0ML=24, H=1024, A=16embeddings: 31,782,912attention: 100,761,600feed-forward: 201,449,472LayerNorm + pooler: 1,147,904335.1MEach layer adds 12H² + 13H numbers: 7.09M at H = 768, 12.60M at H = 1024. The embedding tables do not grow with L.
Where the parameters of each Table 6 model live. Blue: embeddings (23.8M or 31.8M, fixed by H, not by L). Orange: attention. Green: feed-forward. The embedding tables do not grow with L, so in the 3-layer model they are half of everything.

The paper gives parameter counts only for BERT-base and BERT-large. I built every model of Table 6 with the Hugging Face BertConfig (30,522-token vocabulary, 512 positions, feed-forward size 4H, as in Part 2) on PyTorch's meta device, which creates the shapes without allocating any memory, and counted:

python
import torch
from transformers import BertConfig, BertModel

for L, H, A in [(3, 768, 12), (6, 768, 3), (6, 768, 12), (12, 768, 12), (12, 1024, 16), (24, 1024, 16)]:
    cfg = BertConfig(vocab_size=30522, hidden_size=H, num_hidden_layers=L, num_attention_heads=A,
                     intermediate_size=4 * H, max_position_embeddings=512, type_vocab_size=2)
    with torch.device("meta"):              # shapes only, no memory
        model = BertModel(cfg)
    print(L, H, A, sum(p.numel() for p in model.parameters()))
plain text
 #L    #H  #A        params       formula  embeddings   ppl  MNLI-m  MRPC  SST-2
  3   768  12    45,691,392    45,691,392  23,837,184  5.84    77.9  79.8   88.4
  6   768   3    66,955,008    66,955,008  23,837,184  5.24    80.6  82.2   90.7
  6   768  12    66,955,008    66,955,008  23,837,184  4.68    81.9  84.8   91.3
 12   768  12   109,482,240   109,482,240  23,837,184  3.99    84.4  86.7   92.9
 12  1024  16   183,987,200   183,987,200  31,782,912  3.54    85.7  86.9   93.3
 24  1024  16   335,141,888   335,141,888  31,782,912  3.23    86.6  87.8   93.7

The "formula" column is the count from the equation of Part 2 (embeddings, plus 12H2+13H12H^2 + 13H per layer, plus the H×HH \times H pooler), and it matches the real model exactly in every row. Three things stand out:

  • BERT-base is 109.5M and BERT-large is 335.1M parameters, counted this way. The paper rounds them to "110M" and "340M". (Part 2 discusses the 340M.)
  • The two 6-layer models have exactly the same number of parameters (66,955,008), yet 12 heads beat 3 heads on all three tasks (MNLI-m 81.9 versus 80.6). The number of heads only changes how the same H×HH \times H matrices are split up, not how many numbers they hold.
  • The embeddings are a big share of small models. In the 3-layer model, 23.8M of its 45.7M parameters (52%) are embeddings, the word-lookup table alone.
768084889296SST-2, L=3 H=768 A=12: 88.4SST-2, L=6 H=768 A=3: 90.7SST-2, L=6 H=768 A=12: 91.3SST-2, L=12 H=768 A=12: 92.9SST-2, L=12 H=1024 A=16: 93.3SST-2, L=24 H=1024 A=16: 93.7SST-2 93.7MRPC, L=3 H=768 A=12: 79.8MRPC, L=6 H=768 A=3: 82.2MRPC, L=6 H=768 A=12: 84.8MRPC, L=12 H=768 A=12: 86.7MRPC, L=12 H=1024 A=16: 86.9MRPC, L=24 H=1024 A=16: 87.8MRPC 87.8MNLI-m, L=3 H=768 A=12: 77.9MNLI-m, L=6 H=768 A=3: 80.6MNLI-m, L=6 H=768 A=12: 81.9MNLI-m, L=12 H=768 A=12: 84.4MNLI-m, L=12 H=1024 A=16: 85.7MNLI-m, L=24 H=1024 A=16: 86.6MNLI-m 86.6L=3, H=768A=1246M paramsLM ppl 5.84L=6, H=768A=367M paramsLM ppl 5.24L=6, H=768A=1267M paramsLM ppl 4.68L=12, H=768A=12109M paramsLM ppl 3.99L=12, H=1024A=16184M paramsLM ppl 3.54L=24, H=1024A=16335M paramsLM ppl 3.23Dev accuracy (Table 6, average of 5 fine-tuning runs)
Table 6 as a chart. Accuracy rises with every step up in size, for all three tasks, including MRPC with its few thousand training examples. Parameter counts are my measurements; everything else is from the paper.

Was this big for 2018?

The two comparison models, in their own papers:

Can we check these two numbers with our formula? Only roughly, because both models differ from BERT in their embeddings and outputs:

plain text
Vaswani et al. 2017, big encoder: L=6, H=1024: the layers alone, if shaped like BERT's (12H^2 + 13H each) = 75,577,344   (paper says 100M for the encoder)
Al-Rfou et al. 2018: L=64, H=512: the layers alone, if shaped like BERT's (12H^2 + 13H each) = 201,752,576   (paper says 235M)

The layers alone give 75.6M and 201.8M. The rest of the published totals is embeddings and other parts (the big Transformer's shared vocabulary table of about 37,000 tokens × 1,024 adds roughly 38M). So the orders of magnitude agree; the exact numbers depend on details these papers do not spell out in the same way.

Note the careful word "hypothesize": the paper offers an explanation for why fine-tuning scales where feature-based approaches did not, but does not test it directly. The measured fact is Table 6 itself.

A real perplexity, for scale

First, the definition with a toy you can follow. For NN predicted words with probabilities p1,…,pNp_1, \dots, p_N given to the true words:

NLL=−1N∑n=1Nlog⁡pn,perplexity=eNLL\text{NLL} = -\frac{1}{N}\sum_{n=1}^{N} \log p_n, \qquad \text{perplexity} = e^{\text{NLL}}

where NLL is the mean negative log-likelihood in nats (natural-log units). Sanity check: a model that is uniform over 4 words gives every word p=1/4p = 1/4, so NLL =log⁡4=1.386= \log 4 = 1.386 and perplexity =4= 4: "as unsure as a 4-way guess".

A real sentence, three masked words, bert-base-uncased:

plain text
input: she opened the [MASK] with her [MASK] and walked into the [MASK] .
door     p = 0.9586   -log p = 0.0423   (top 3 guesses: door, gate, box)
key      p = 0.1324   -log p = 2.0222   (top 3 guesses: key, hand, keys)
kitchen  p = 0.1823   -log p = 1.7022   (top 3 guesses: room, kitchen, house)
mean NLL = (0.0423 + 2.0222 + 1.7022) / 3 = 1.2554 nats -> perplexity = exp(1.2554) = 3.51
GPT-2 mean NLL = 3.7125 -> perplexity 40.96  (same three words, left context only)
she opened the [MASK] with her [MASK] and walked into the [MASK] .1. probability of the true word2. surprise = minus log p3. average, then e to thatdoor0.95860.0423key0.13242.0222kitchen0.18231.7022mean = 1.2554 natsppl = e^1.2554 = 3.51about as unsure as a pickamong 3.5 equal wordsSame three words for GPT-2, which sees only the left side: perplexity 40.96. "key" alone gets p = 0.0006.
Perplexity in three steps on one sentence. 1: the probability of each true word. 2: the surprise, minus the log of each. 3: average the surprise and raise e to it: 3.51, "about as unsure as a pick among 3.5 equal words".

BERT's 3.51 here is close to Table 6's numbers (3.23 to 5.84), but it is one sentence. (GPT-2's 40.96 is not a fair comparison: it predicts each word from the left side only, a harder task.)

What does a masked-LM perplexity of about 4 look like on real text? I measured the released bert-base-uncased on the test split of WikiText-2, a public set of Wikipedia articles, using the paper's masking recipe (15% of tokens chosen; of those 80% [MASK], 10% random, 10% unchanged), sequences of 512 tokens, and a fixed random seed:

python
chosen = torch.rand(x.shape) < 0.15                    # pick 15% of positions
chosen[:, 0] = chosen[:, -1] = False                   # never [CLS] or [SEP]
r = torch.rand(x.shape)
inp = x.clone()
inp[chosen & (r < 0.8)] = tok.mask_token_id            # 80%: [MASK]
rnd = chosen & (r >= 0.8) & (r < 0.9)                  # 10%: a random token
inp[rnd] = torch.randint(len(tok), x.shape)[rnd]       # the other 10% stay unchanged
logp = torch.log_softmax(model(input_ids=inp).logits[chosen], -1)
nll += -logp.gather(1, x[chosen][:, None]).sum()       # cross-entropy on the chosen positions only
plain text
WikiText-2 test: 261,428 WordPiece tokens -> 512 sequences of 512 ([CLS] + 510 + [SEP])
predicted positions: 38,984 of 261,120 (14.93%): 31,282 [MASK], 3,947 random, 3,755 unchanged
mean cross-entropy = 1.9967 nats -> masked-LM perplexity = 7.36   (top-1 accuracy 63.6%)

So on this text the released model has a perplexity of 7.36, and its first guess is right 63.6% of the time. That is higher (worse) than the 3.99 in Table 6. I cannot reproduce the paper's number, and the comparison is not like for like:

  • The paper measured "held-out training data": text from the same BooksCorpus and Wikipedia mix BERT was trained on. WikiText-2 is a different sample, with its own formatting (I undid its @-@ style marks and dropped headings, but other differences remain).
  • Table 6's models are the ablation models, trained with the same procedure; the paper does not say the 3.99 was measured on the released checkpoint.
  • My random replacements draw from the whole vocabulary, including rarely used tokens.

What the measurement does show: the 15% / 80-10-10 recipe behaves as described (14.93% of positions chosen, split 80.2% / 10.1% / 9.6%), and a single-digit perplexity on unseen text is in the same range as Table 6's numbers.

Feature-based approach with BERT

Part 1 introduced two ways to reuse a pre-trained model: feature-based (freeze it, use its vectors) and fine-tuning (train everything). Every BERT result so far was fine-tuned. Section 5.3 asks whether frozen BERT features are good too.

Use case. A team wants to try twenty different classifiers on a million documents. Fine-tuning means twenty full BERT training runs. Feature-based means running BERT once over the million documents, saving the vectors, and training twenty small models on the saved vectors, each in minutes.

The tags come from the CoNLL-2003 data format, described in its own paper:

A real CoNLL-2003 dev sentence (#99): one tag per word, one vector per wordB-LOCLithuaniaLithuaniaO--B-PERDaniusDani##usI-PERGleveckasG##lev##eck##asO((O13rd13##rdO))top: BIO tags (B = beginning of a name, I = inside it, O = outside). middle: WordPiece pieces (bert-base-cased).bottom: the tagger reads only the vector of the first piece of each word (highlighted); the other pieces are ignored.
A real CoNLL-2003 dev sentence (#99), "Lithuania - Danius Gleveckas ( 13rd )": one tag per word (B-LOC, O, B-PER, I-PER, O, O, O). WordPiece splits "Danius" into Dani ##us and "Gleveckas" into G ##lev ##eck ##as; the tagger reads only the first piece of each word.

Why would a CRF help? A per-token classifier picks each tag on its own, and can produce impossible sequences. A toy with hand-made scores for three words (rows) and three tags (columns):

plain text
scores (rows = words ['met', 'Ada', 'Lovelace'], columns = ['O', 'B-PER', 'I-PER']): 3.0, 0.5, 0.2; 1.2, 1.0, 0.4; 0.3, 0.6, 2.5
per-token argmax: ['O', 'O', 'I-PER']  (sum 6.7; O followed by I-PER is invalid)
best valid sequence: ['O', 'B-PER', 'I-PER']  (sum 6.5)
Which tag may follow which (BIO rules)Hand-made tag scores for three wordsnext tag (column)OB-PERI-PEROO attends to OO attends to B-PERB-PERB-PER attends to OB-PER attends to B-PERB-PER attends to I-PERI-PERI-PER attends to OI-PER attends to B-PERI-PER attends to I-PERprevious tag (row)OB-PERI-PERmet3.00.50.2Ada1.21.00.4Lovelace0.30.62.5each word on its own: O O I-PER (sum 6.7): invalidbest valid sequence: O B-PER I-PER (sum 6.5)
Left: which tag may follow which (an I-PER cannot follow O). Right: picking the best tag per word gives O O I-PER (score 6.7), which breaks the rule; the best valid sequence is O B-PER I-PER (score 6.5). A CRF finds the best valid sequence; BERT's tagger in Section 5.3 does not use one.

The six feature choices in Table 7 read different hidden states:

which of the 13 hidden states each feature choice usesE123456789101112Embeddingsone layerpaper Dev F1 91.0Second-to-last hiddenone layerpaper Dev F1 95.6Last hiddenone layerpaper Dev F1 94.9Weighted sum last fourweighted sum of 4paper Dev F1 95.9Concat last four4 x 768 = 3,072 numberspaper Dev F1 96.1Weighted sum all 12 layersweighted sum of 12paper Dev F1 95.5E = embedding output (before any layer); 1 to 12 = outputs of the 12 layers
The six feature choices of Table 7. A weighted sum learns one weight per layer and adds the layers up; concatenation glues the last four 768-number vectors into one of 3,072 numbers.

Reading Table 7 carefully:

  • Fine-tuning: BERT-large 96.6 Dev / 92.8 Test, BERT-base 96.4 / 92.4. Note the word "competitively": on the Test set, CSE (93.1) is higher than BERT-large (92.8). NER is not one of BERT's eleven new records.
  • Embeddings alone: 91.0. The embedding output is the word-lookup vector plus position and segment, before any attention. It knows nothing about the sentence. Context is worth more than 5 points here.
  • Last hidden 94.9 is worse than second-to-last 95.6. One common explanation (not tested in the paper): the last layer is shaped most strongly by the pre-training tasks, predicting masked words, so the layer before it holds more general information.
  • Combining layers helps: weighted sum of the last four 95.9, concatenating them 96.1. Summing all 12 layers (95.5) is worse than using only the top four.
  • 96.1 versus 96.4: the best frozen choice is 0.3 points behind fine-tuned BERT-base on Dev.

The frozen-feature tagger of Section 5.3, shape by shape, on "Ada Lovelace met Charles Babbage in London .":

plain text
WordPiece tokens               [13]
all hidden states              [13, 13, 768]
first sub-token of each word   [8, 13, 768]
concat last four layers        [8, 3072]
BiLSTM output                  [8, 768]
tag scores                     [8, 9]
BiLSTM layer 1 (input 3072): 2 directions x 4 x (384*3072 + 384*384 + 2*384) = 10,622,976
BiLSTM layer 2 (input 768):  2 directions x 4 x (384*768 + 384*384 + 2*384)  = 3,545,088
classifier 9 x 768 + 9 = 6,921;  tagger total 14,174,985 trainable, BERT-base-cased frozen (108,310,272 parameters never change)

The LSTM count: each direction has 4 gates, each gate a matrix from the input (3,072 numbers) and one from its own previous state (384 numbers), plus two bias vectors in the PyTorch layout: 4×(384×3072+384×384+2×384)4 \times (384 \times 3072 + 384 \times 384 + 2 \times 384) per direction.

The frozen-feature tagger of Section 5.3, shape by shape (n = number of words)BERT hidden statesfrozenn x 13 x 768concat last fourno weightsn x 3,072BiLSTM, 2 layers14.17M weightsn x 768linear W, 9 x 7686,921 weightsn x 9softmaxper word9 tag probabilities[h9 ; h10 ; h11 ; h12]384 forward + 384 backward9 rows, one per tagTrainable: 14,174,985 numbers (BiLSTM + classifier). BERT-base-cased: 108,310,272 numbers, never updated.
The Section 5.3 tagger, shape by shape. BERT's 13 hidden states (n × 13 × 768, frozen) → concatenate the last four (n × 3,072) → a 2-layer BiLSTM with 768 outputs (14.17M weights, trained) → a linear layer to 9 tags (6,921 weights) → softmax per word.
### A smaller re-run of the feature-based experiment

Can we see the same pattern ourselves? I ran a smaller version of the experiment in bert_part5_ner.py:

  • Model: bert-base-cased, frozen (no weight of BERT changes).
  • Data: CoNLL-2003 from the Hugging Face hub (eriktks/conll2003). Training: the first 5,000 of the 14,041 training sentences. Evaluation: the whole Dev set, 3,250 sentences (51,362 words). Each sentence is read on its own: this copy of the dataset has no document boundaries, so there is no "maximal document context".
  • Features: all 13 hidden states, first sub-token of every word, saved once.
  • Classifier: a 2-layer BiLSTM with 384 units per direction (768 in total, my reading of "768-dimensional"), then a linear layer over the 9 tags; Adam, learning rate 0.001, batch 32, 4 epochs. Each choice trained with 3 seeds, averaged (the paper used 5).
  • Score: entity-level F1 on Dev, computed by my own code (an entity counts as correct only if its type and its exact span match).

The heart of the feature extraction:

python
enc = tok(words, is_split_into_words=True, return_tensors="pt")
with torch.no_grad():                                                  # BERT is frozen
    hs = torch.stack(bert(**enc, output_hidden_states=True).hidden_states, 2)[0]   # (tokens, 13, 768)
wid = enc.word_ids(0)                                                  # which word each token belongs to
first = [t for t, w in enumerate(wid) if w is not None and (t == 0 or wid[t - 1] != w)]
features = hs[first]                                                   # (words, 13, 768): first sub-token only

concat_last_four = features[:, 9:13].flatten(1)                        # (words, 3072)
weighted_last_four = (torch.softmax(w, 0)[None, :, None] * features[:, 9:13]).sum(1)   # w: 4 learned weights

The results:

plain text
CoNLL-2003: 5000 training sentences used, 3250 dev sentences; labels: ['O', 'B-PER', 'I-PER', 'B-ORG', 'I-ORG', 'B-LOC', 'I-LOC', 'B-MISC', 'I-MISC']
features: 67634 train words, 51362 dev words, 13 layers x 768 numbers each, 34 s
Embeddings                   dev F1 = 82.66   (seeds: 82.73, 82.42, 82.82; 113 s)
Second-to-last hidden        dev F1 = 89.94   (seeds: 88.72, 90.25, 90.84; 127 s)
Last hidden                  dev F1 = 88.64   (seeds: 88.13, 89.03, 88.77; 161 s)
Weighted sum last four       dev F1 = 90.37   (seeds: 90.08, 90.00, 91.05; 311 s)
Concat last four             dev F1 = 91.17   (seeds: 90.94, 91.41, 91.16; 484 s)
Weighted sum all 12 layers   dev F1 = 90.49   (seeds: 91.28, 89.60, 90.59; 450 s)
paper, Table 7 (BiLSTM, document context, 5 runs)our smaller re-run (5,000 sentences, 3 runs)7580859095100EmbeddingsEmbeddings: 91.091.0Embeddings: 82.6682.66Second-to-last hiddenSecond-to-last hidden: 95.695.6Second-to-last hidden: 89.9489.94Last hiddenLast hidden: 94.994.9Last hidden: 88.6488.64Weighted sum last fourWeighted sum last four: 95.995.9Weighted sum last four: 90.3790.37Concat last fourConcat last four: 96.196.1Concat last four: 91.1791.17Weighted sum all 12 layersWeighted sum all 12 layers: 95.595.5Weighted sum all 12 layers: 90.4990.49entity-level F1 on the CoNLL-2003 Dev set (axis starts at 75)
Dev F1 for each feature choice: the paper's Table 7 and my smaller re-run. The absolute numbers are lower in my run (fewer training sentences, no document context, a smaller training budget), but the shape agrees.

Lined up against Table 7:

Feature choicePaper, Dev F1 (Table 7)Our re-run, Dev F1
Embeddings91.082.66
Second-to-last hidden95.689.94
Last hidden94.988.64
Weighted sum last four95.990.37
Concat last four96.191.17
Weighted sum all 12 layers95.590.49

What agrees with the paper:

  • The embeddings alone are clearly worst (82.66), about 6 to 8.5 points below every choice that uses BERT's layers. Context is what BERT adds.
  • Concatenating the last four layers is best (91.17), in both the paper and our run.
  • The second-to-last layer beats the last layer (89.94 versus 88.64), as in the paper (95.6 versus 94.9).
  • Combining the top four layers beats any single layer, whether summed or concatenated.

What does not agree: in our run the weighted sum of all 12 layers (90.49) edges out the weighted sum of the last four (90.37), the opposite of the paper's order. The gap is 0.12 points, while the three seeds of the 12-layer choice alone range from 89.60 to 91.28, so this run cannot separate those two.

Our absolute numbers are about 5 to 8.5 points lower than the paper's. That is expected, not a contradiction: we trained on 5,000 of the 14,041 sentences, for 4 epochs, with no hyperparameter search, and without the document context the paper adds to every sentence. We also did not fine-tune BERT, so this experiment says nothing about the 96.4 of fine-tuning; it only checks the shape of the feature-based rows.

How long to pre-train

The last two ablations are in Appendix C. The first asks whether BERT's long pre-training is necessary.

The total pre-training in Question 1 is the paper's own product: 128,000 words per batch × 1,000,000 steps = 128 billion words processed (with the paper's rounding of 256 × 512 = 131,072 to 128,000, as in Part 3).

How to read it:

  • Both curves rise, then flatten. More pre-training always helped, but with smaller gains later on.
  • The MLM curve is still climbing between 500k and 1M steps: that is the "almost 1.0%" of Question 1.
  • At the very first checkpoint the left-to-right model is ahead, which fits "converges slightly slower": it gets a training signal from every word, the MLM from only 15%. But the MLM curve passes it straight away and the gap then widens to about 2 points by the end (reading the figure, not a number printed in the paper).

The trade-off is worth it. The MLM sees fewer training signals per batch, but each one uses context from both sides.

How many predictions does each objective get from one batch?

plain text
one batch = 256 x 512 = 131,072 tokens; masked LM predicts 15% = 19,661; left-to-right LM predicts every next token, about 131,072
over 1,000,000 steps: 19,661,000,000 masked-LM predictions versus about 131,072,000,000 for LTR (6.67x more)
Appendix C.1: how many predictions one batch givesmasked LM15% of tokens: 19,661 predictions per batchleft-to-right LMevery token: about 131,072 per batchSame batch of 131,072 tokens. The masked LM gets 6.67x fewer training signals, yet overtakes LTR almost at once (Figure 5).
Training signal per batch. The masked LM predicts 15% of tokens (19,661 per batch); a left-to-right LM predicts almost every token (about 131,072). The masked LM gets 6.67 times fewer training signals, yet Figure 5 shows it overtakes the left-to-right model almost at once.

That is the "converges slightly slower" of Question 2, made concrete: per step the masked LM learns from 6.67 times fewer predictions. Figure 5 shows it is still the better choice.

Masking strategies

Part 3 explained the 80/10/10 masking recipe and why the paper uses it. Appendix C.2 tests whether it was a good choice.

Why would the mismatch be worse for frozen features? With fine-tuning, every weight of BERT is adjusted on the task's real sentences, so the model can unlearn any habit tied to [MASK]. With frozen features, BERT is used exactly as pre-trained. If its vectors for ordinary, unmasked words are not very informative (because it was never asked to predict an unmasked word), nothing can fix that afterwards.

what happens to a chosen tokenDev results (Table 8)MASKSAMERNDMNLINER fine-tuneNER feature-based80%10%10%BERT84.295.494.9100%84.394.994.080%20%84.195.294.680%20%84.495.294.720%80%83.794.894.6100%83.694.994.6highlighted: best and worst value in each column
Table 8 as a picture. Each bar shows what happens to the tokens chosen for prediction. In each results column, the best value is drawn in the accent colour and the worst in red.

Row by row:

  1. 80 / 10 / 10 (BERT's recipe): 84.2, 95.4, 94.9. The best NER scores in both columns.
  2. 100 / 0 / 0 (always [MASK]): MNLI 84.3 is fine, but feature-based NER drops to 94.0, the worst value in that column and 0.9 below BERT's recipe. This is the predicted mismatch: the model only ever had to predict at [MASK], and its frozen vectors for normal words are less useful.
  3. 80 / 0 / 20: 84.1, 95.2, 94.6. Random replacement without any "keep" cases.
  4. 80 / 20 / 0: 84.4, 95.2, 94.7. Interestingly the best MNLI score in the table, 0.2 above BERT's recipe.
  5. 0 / 20 / 80 (no [MASK] at all): 83.7, 94.8, 94.6.
  6. 0 / 0 / 100 (always random): 83.6 on MNLI, the worst. Here the model is trained to find and fix wrong words, and 15% of every training input is noise.

Putting numbers on "robust": across all six recipes, MNLI moves within 0.8 points (83.6 to 84.4) and fine-tuned NER within 0.6 (94.8 to 95.4). Feature-based NER moves within 0.9 (94.0 to 94.9), with always-[MASK] at the bottom. These are small differences, single numbers without error bars, so the honest summary is: the masking recipe matters little when you fine-tune, and a bit more when you freeze.

One more detail, which the paper does not explain: BERT's own row here scores 84.2 on MNLI, while the BERT-base row of Table 5 scores 84.4. They are probably different training runs, but the paper does not say.

Next, in Part 6: what happened after BERT. Later work revisited exactly these ablations (RoBERTa dropped NSP; ALBERT, DistilBERT and ELECTRA changed size and the masked LM itself), BERT's limits, how it is used today, and a summary of the whole paper.

Run it yourself

The scripts behind every measured number in this part:

bash
pip install torch transformers datasets
python bert_part5.py        # writes results/part5.json
python bert_part5_ner.py    # writes results/part5_ner.json
Terminal output of bert_part5.py: Table 5 differences, Table 6 parameter counts matching the formula, and the WikiText-2 masked-LM perplexity of 7.36
The real output of bert_part5.py.
Terminal output of bert_part5_ner.py: dataset sizes, feature extraction, and Dev F1 for six frozen-feature choices, each averaged over three seeds
The real output of bert_part5_ner.py.

References

The BERT paper

  1. J. Devlin, M.-W. Chang, K. Lee, K. Toutanova. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. NAACL 2019 (ACL Anthology).
  2. Google Research. BERT code and models, including extract_features.py and the cased models.

Papers the BERT paper cites in this part

  1. A. Radford, K. Narasimhan, T. Salimans, I. Sutskever. Improving Language Understanding by Generative Pre-Training (OpenAI GPT). OpenAI, 2018.
  2. M. E. Peters et al. Deep contextualized word representations (ELMo, "2018a"). NAACL 2018.
  3. M. E. Peters, M. Neumann, L. Zettlemoyer, W. Yih. Dissecting Contextual Word Embeddings: Architecture and Representation ("2018b", the mixed results on bi-LM size). EMNLP 2018.
  4. O. Melamud, J. Goldberger, I. Dagan. context2vec: Learning Generic Context Embedding with Bidirectional LSTM. CoNLL 2016.
  5. A. Vaswani et al. Attention Is All You Need. NeurIPS 2017.
  6. R. Al-Rfou, D. Choe, N. Constant, M. Guo, L. Jones. Character-Level Language Modeling with Deeper Self-Attention. AAAI 2019.
  7. E. F. Tjong Kim Sang, F. De Meulder. Introduction to the CoNLL-2003 Shared Task: Language-Independent Named Entity Recognition. CoNLL 2003.
  8. K. Clark, M.-T. Luong, C. D. Manning, Q. V. Le. Semi-Supervised Sequence Modeling with Cross-View Training (CVT, Table 7). EMNLP 2018.
  9. A. Akbik, D. Blythe, R. Vollgraf. Contextual String Embeddings for Sequence Labeling (CSE, Table 7). COLING 2018.

Other sources used in this part

  1. L. A. Ramshaw, M. P. Marcus. Text Chunking using Transformation-Based Learning (the IOB tagging scheme). Workshop on Very Large Corpora, 1995.
  2. S. Merity, C. Xiong, J. Bradbury, R. Socher. Pointer Sentinel Mixture Models (WikiText-2). ICLR 2017.
  3. Code for this part: bert_part5.py, bert_part5_ner.py, bert_part5_math.py.