Pre-training: Masked LM and Next Sentence Prediction
Section 3.1 and Appendices A.1 and A.2 of the BERT paper, line by line: why a normal language model cannot look both ways, the masked language model and its 80/10/10 rule, next sentence prediction, the training data and the full training recipe. Every claim run on the real model or measured in code.
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 2 built the model: a stack of Transformer layers that turns tokens into vectors. But a freshly built model knows nothing. Its millions of weights are random numbers.
This part is about how BERT learns language from plain text, before any real task. The paper uses two training games: fill in the blanks (the masked language model) and does this sentence come next? (next sentence prediction). We go through Section 3.1 paragraph by paragraph, then the two appendices that give the exact recipe, and we run every piece on the real model.
Here is the problem. Say we want every position to predict its own word from all the other words. In one layer, that is fine: we can simply forbid each position from looking at itself. But BERT has 12 layers, and information moves sideways at every layer.
How a word sees itself through two layers. Position 2 must predict "cat" and never looks at "cat" directly. But in layer 1, the neighbour "sat" reads "cat". In layer 2, position 2 reads "sat", and the answer comes back in.
So with two or more layers, every word can find its own answer by going through a neighbour. Training would then learn to copy, not to understand. That is what "trivially predict the target word in a multi-layered context" means.
I trained three tiny Transformers on the same real text (WikiText-2, a public set of Wikipedia articles) and measured how well each predicts words. Each has 128 numbers per token and 4 attention heads, and trains for 3,000 steps on the Apple GPU of a laptop.
A. Left-to-right, 2 layers. A normal language model: predict the next token from the tokens before it.
B. Both sides, 1 layer. Predict each token from every other token. Position t never sees token t.
C. Both sides, 2 layers. Exactly like B, with one more layer.
python
# simplified from bert_part3_seeitself.py# B and C: position t may look at every position except itselfnot_self = torch.arange(T)[None, :] != torch.arange(T)[:, None]# in layer 1, the query at position t is built from "position t" only, never from token tx = blocks[0](x, not_self, query_from=position_embedding, residual=False)for block in blocks[1:]: # model C has one more layer here x = block(x, not_self)
A: left-to-right, 2 layers training loss: step 1 9.30 -> step 3000 4.65 validation loss 5.05 (perplexity 156.1)B: both sides, 1 layer training loss: step 1 9.20 -> step 3000 3.95 validation loss 4.37 (perplexity 79.1)C: both sides, 2 layers training loss: step 1 9.17 -> step 3000 0.87 validation loss 0.90 (perplexity 2.5)
Training loss of the three tiny models. A and B learn slowly and honestly. C's loss drops far lower, and stays low on text it never trained on, because it reads its own answer back through a neighbour.
Three things to notice:
C is cheating. Its loss on unseen text is 0.90 (perplexity 2.5). No two-layer model with 128 numbers per token can predict real Wikipedia text that well. It is reading the answer.
B is honest, and it beats A (4.37 against 5.05 on unseen text). With one layer there is no way back to the answer, so B really predicts from both sides. Seeing both sides genuinely helps, which is the paper's point from Part 1.
The problem appears exactly at the second layer. That is the "multi-layered context" of the paper. BERT has 12 layers, so it needs another way.
The masked language model. The whole sentence goes in, with the chosen token hidden. Only the output vector of the hidden position (here T₄) goes through a softmax over the vocabulary and counts in the loss.
The model is trained with cross-entropy loss on the masked positions only:
LMLM=−∣M∣1i∈M∑logP(xi∣x~)
where:
x is the original sequence of tokens, and x~ (read "x tilde") is the same sequence after masking;
M is the set of positions chosen for prediction, and ∣M∣ is how many there are;
xi is the original token at position i, the answer;
P(xi∣x~) is the probability the model gives to that answer, from the softmax at position i, having seen only x~;
log is the natural logarithm. If the model is sure and right, P=1 and −logP=0. If it gives the answer a probability of 0.01, −logP=4.6.
How is P computed from the final vector Ti? The paper only says "an output softmax over the vocabulary, as in a standard LM". In the released model it is a small head: one more dense layer with GELU and LayerNorm, then a multiplication by the token embedding matrix (the same matrix used at the input, shared to save weights) plus a bias, then softmax.
Here is that head as equations, for one masked position i:
Ti is BERT's final 768-number vector at position i;
Wt (768×768) and bt are the head's own small "transform" layer;
E is the token embedding table from the input, V×768 with V=30,522, reused here; hiE⊤ is the dot product of hi with every token's input vector;
bo holds one extra learned number per vocabulary token;
z is the list of V scores (often called logits), and P(v∣x~) is the softmax probability of token v.
From Tᵢ to a probability for every token of the vocabulary. Tᵢ (768 numbers) goes through a dense layer, GELU and LayerNorm to give h. h is multiplied by the transposed token table Eᵀ (768 × 30,522) and a bias is added, giving 30,522 scores z. Softmax turns them into probabilities; the loss is minus the log of the right token's probability.
The whole computation on one real example (bert_part3_math.py). The paper's own Appendix A.1 sentence "the man went to the store to buy a gallon of milk", with "store" masked:
plain text
the masked-LM head of bert-base-uncased transform: dense (768, 768) + GELU + LayerNorm(768) decoder weight (30522, 768), bias (30522,) decoder weight is the input token-embedding matrix (same memory): Truesentence: [CLS] the man went to the [MASK] to buy a gallon of milk . [SEP] (masked position i = 6, original word "store")T_i: (768,), first 4 numbers: +0.388, +0.178, -0.210, +0.139scores z = h E^T + b: (30522,); largest score 12.814; smallest -12.503top 5 scores and their probabilities: store z = 12.814 exp(z - max) = 1.0000 p = 0.4743 refrigerator z = 10.991 exp(z - max) = 0.1615 p = 0.0766 fridge z = 10.915 exp(z - max) = 0.1496 p = 0.0710 kitchen z = 10.707 exp(z - max) = 0.1215 p = 0.0576 counter z = 10.278 exp(z - max) = 0.0792 p = 0.0375log of the softmax denominator, log sum_v exp(z_v) = 13.5602p("store") = exp(12.814 - 13.560) = 0.4743cross-entropy = -log p("store") = 0.7459 (library: 0.7459)for scale: a uniform guess over 30,522 tokens gives -log(1/30522) = 10.3262
Follow it by hand. The softmax needs ∑vezv over all 30,522 scores, and its log is 13.5602. So
Two useful reference points: a model that guessed uniformly would get −log(1/30,522)=10.33, and a perfect model 0. Notice that refrigerator, fridge and kitchen also score high: places where milk lives. The model knows the scene, not just the word.
Cross-entropy as a function of the probability the model gave the right word: loss = −log p. A uniform guess (p = 3.3 × 10⁻⁵) costs 10.33; the real examples in this part sit along the curve, from "##s" in "penguins" (p = 0.003, loss 5.76) to "of" (p = 0.9986, loss 0.0014).
The paper's contrast with denoising auto-encoders (Vincent et al., 2008) is about where the loss is computed:
Where the loss is. A denoising auto-encoder (Vincent et al., 2008) rebuilds the whole input, so every position has a loss. BERT's masked LM puts a loss only on the hidden position; the other positions produce no loss.
Use case. This is the same "fill in the blank" you saw in Part 1, where BERT put 0.901 on "bank". The same masked-LM head is still a quick way to probe what a model knows: give it "The capital of France is [MASK]." and read its guesses.
Appendix A.1 walks through the rule on one sentence, and explains the reasons.
The 80/10/10 rule. Each chosen token becomes [MASK] 80% of the time, a random word 10% of the time, and stays itself 10% of the time. The target is always the original word.
What share of all tokens ends up in each state? Combine the two random choices. A token is first chosen with probability 0.15, and a chosen token is then masked with probability 0.8, replaced with 0.1, or kept with 0.1:
and the other 1−0.15=0.85 are untouched and not predicted. So in a typical batch, 12% of tokens are [MASK], 1.5% are wrong words, 1.5% are unchanged but still predicted, and 85% are plain context. The model sees 85%+1.5%=86.5% of tokens exactly as they were written. For one 512-token sequence that is about 76.8 predictions: 61.4 masks, 7.7 random and 7.7 unchanged.
Every 200 tokens of training text, on average. 170 untouched (no loss), 24 [MASK], 3 random words and 3 unchanged-but-predicted. A loss on 12 + 1.5 + 1.5 = 15%; 85 + 1.5 = 86.5% look like the original text.The procedure on two real sentences, step by step. Choose 15% of positions (here "store", "milk" and "bag"). Of the chosen ones, one becomes [MASK], one becomes a random word ("piano"), and one stays the same. The model must predict all three originals; the other 17 tokens carry no loss.
In plain words, the three cases do three jobs:
[MASK] (80%) is the main exercise: fill in a blank from both sides.
A random word (10%) teaches the model that a visible word might be wrong, so it must check every word against its context.
Unchanged (10%) teaches the model that a visible word is usually right, so its vector for a normal word should stay close to that word. This is the case that keeps the model useful when there are no [MASK] tokens at all, as in fine-tuning.
Here is the rule exactly as Section 3.1 states it:
python
def mask_tokens(ids, rng, rate=0.15): """Choose 15% of the positions. Each chosen one: 80% [MASK], 10% random token, 10% unchanged.""" cand = [i for i, t in enumerate(ids) if t not in SPECIAL] # never [CLS], [SEP] or padding chosen = sorted(rng.sample(cand, max(1, round(len(cand) * rate)))) out, case = list(ids), {} for i in chosen: r = rng.random() if r < 0.8: out[i], case[i] = tok.mask_token_id, "mask" elif r < 0.9: out[i], case[i] = rng.randrange(VOCAB_SIZE), "random" # any id in the vocabulary else: case[i] = "same" # left as it is, but still predicted return out, chosen, case
I ran it over the whole test split of WikiText-103 (62 Wikipedia articles), cut into sequences of 128 WordPiece tokens, with a fixed random seed:
plain text
1. the masking procedure (paper version), seed 0 word-piece tokens (not counting [CLS]/[SEP]): 261,324 chosen for prediction: 39,406 = 15.08% of tokens of the chosen: [MASK] 31,548 (80.06%), random 3,933 (9.98%), unchanged 3,925 (9.96%) random replacements as a share of ALL tokens: 1.51% (paper: 10% of 15% = 1.5%)
The procedure, measured. Every share lands within 0.1 points of the paper's numbers (the dashed lines): 15% chosen; of those 80% [MASK], 10% random and 10% unchanged; 1.5% of all tokens replaced at random.
How does this compare with Google's released code (create_pretraining_data.py in the BERT repository)? I read it line by line. It does the same thing, with three small differences that the paper does not mention:
It picks an exact count per sequence, round(0.15×length), which for 128 tokens is 19. Re-running its logic on the same text chose exactly the same number of tokens as above, 39,406.
It has a cap, max_predictions_per_seq (20 in the README's example for length 128). The README says to set it to about length × 15%, so the cap and the rate agree.
It never masks [CLS] or [SEP], and the random replacement is drawn from the whole vocabulary file, special tokens included.
Now the three cases of Appendix A.1, on the real bert-base-uncased. One honest change: the paper writes the sentence without a full stop. Without it, the model spends its guess at position 4 on the missing full stop (real sentences end with punctuation), so I added one.
plain text
3. "my dog is hairy." with token 4 chosen (Appendix A.1). What the model predicts at position 4: input: [CLS] my dog is [MASK] . [SEP] P(hairy) = 1.7e-05 P(apple) = 1.5e-06 top 5: dead 0.131, here 0.096, fine 0.080, gone 0.069, missing 0.036 input: [CLS] my dog is apple . [SEP] P(hairy) = 2.3e-10 P(apple) = 0.9997 top 5: apple 1.000, apples 0.000, oak 0.000, orchard 0.000, orange 0.000 input: [CLS] my dog is hairy . [SEP] P(hairy) = 0.9999 P(apple) = 9.0e-10 top 5: hairy 1.000, furry 0.000, shaggy 0.000, ugly 0.000, covered 0.000
This short example teaches more than it seems:
[MASK]: "hairy" gets a probability of only 0.000017. That is correct behaviour: nothing in "my dog is ___." points to "hairy". The model offers sensible guesses ("dead", "here", "fine").
Random word: the model believes "apple" (0.9997). In a four-word sentence there is no context to overrule it, and in pre-training only 1.5% of all tokens were random replacements, so a visible word is almost always the real one.
Unchanged: "hairy" gets 0.9999. This is the "bias towards the actual observed word" that the appendix describes.
On a short sentence the model has little to go on. On real paragraphs it has a lot. I masked 400 of the 128-token sequences above with the same procedure and checked how often the real model's top guess is the original token:
plain text
4. top-1 accuracy of bert-base-uncased on the chosen positions of 400 sequences (7,600 predictions) all chosen positions: 61.1% [MASK] 58.3% (3,555 of 6,095) random token 47.7% (368 of 771) unchanged 97.7% (717 of 734)
Top-1 accuracy of the real model on 7,600 chosen positions, split by what the position showed.
With a whole paragraph of context, BERT finds the hidden word 58.3% of the time, out of 30,522 possible tokens. And it corrects almost half of the random replacements (47.7%), which is exactly the skill the 10% random case trains. One caution: WikiText is made of Wikipedia articles, and BERT was pre-trained on Wikipedia, so the model may have seen some of this text. These numbers show the mechanism, not a clean test score.
How one next-sentence example is made. Sentence A comes from a document. Half of the time B is the real next sentence (IsNext); half of the time it is a sentence from a random document (NotNext). The vector C of the [CLS] token makes the two-way decision.
Here is where NSP lives in the paper's Figure 1: the left-most output, C, feeds the box marked NSP.
The NSP loss is cross-entropy over two classes:
LNSP=−logP(y∣C),P(⋅∣C)=softmax(WNSPC)
where:
y is the true label, IsNext or NotNext;
C is the final vector of [CLS], with 768 numbers in BERT-base;
WNSP is a small learned matrix with 2 rows (one score per class) and 768 columns.
(In the released model, C first passes through one extra layer, called the pooler, a 768-by-768 dense layer with a tanh function, and the classifier has a bias. The idea is the same.)
The released bert-base-uncased still contains its trained NSP layer, so we can ask it. Here are the paper's two examples and four pairs of my own:
plain text
1. the pre-trained NSP head on single pairs: probability of IsNext 0.9999 [paper A.1, labelled IsNext] A: the man went to [MASK] store | B: he bought a gallon [MASK] milk 0.0011 [paper A.1, labelled NotNext] A: the man [MASK] to the store | B: penguin [MASK] are flight ##less birds 1.0000 [ours, true next] A: she opened the fridge. | B: there was nothing left but an old lemon. 3.5e-06 [ours, random] A: she opened the fridge. | B: the treaty was signed in 1648 by both parties. 1.0000 [ours, true next] A: the match was delayed by rain. | B: play finally started two hours late. 1.1e-05 [ours, random] A: the match was delayed by rain. | B: photosynthesis turns light into chemical energy.
All six are right, and very confident. The paper reports how good the final model is in a footnote:
The NSP head, worked by hand. In the released model, C first passes through the pooler (Part 2), then a two-way linear classifier:
c=tanh(CWp+bp),s=cW⊤+b,P(IsNext)=es0+es1es0
where Wp is 768×768, W is 2×768 (one row per class: row 0 IsNext, row 1 NotNext), s=(s0,s1) are the two scores, and tanh squeezes every number into the range -1 to 1. The loss is −log of the right class's probability. On the paper's two Appendix A.1 examples:
plain text
A: the man went to [MASK] store | B: he bought a gallon [MASK] milk | label IsNext scores W C + b = [5.111, -4.188] (library: [5.111, -4.188]) softmax: P(IsNext) = 0.999908, P(NotNext) = 0.000092; loss -log P(IsNext) = 0.000092A: the man [MASK] to the store | B: penguin [MASK] are flightless birds | label NotNext scores W C + b = [-2.080, 4.691] (library: [-2.080, 4.691]) softmax: P(IsNext) = 0.001145, P(NotNext) = 0.998855; loss -log P(NotNext) = 0.001146
For the first pair, the gap between the two scores is 5.111−(−4.188)=9.299, so
P(IsNext)=1+e−9.2991=1+0.0000921=0.999908
(A two-way softmax only depends on the difference of the two scores, which is why it can be written with a single exponential.)
Next sentence prediction reads one vector, C, and makes two scores. C goes through the pooler, then W (2 × 768) and a bias. Softmax over the two scores gives IsNext 0.99991 for the first A.1 example and NotNext 0.99885 for the second, exactly as the library computes them.
**Checking footnote 5.** I built 1,000 pairs from real Wikipedia text: for 500, B is the true next sentence; for 500, B is a random sentence from a different article. One sentence on each side, no masks.
plain text
2. NSP accuracy on 1000 real sentence pairs from 61 WikiText-103 test articles (one sentence each side, no masks) IsNext pairs: 486 of 500 right (97.2%) NotNext pairs: 478 of 500 right (95.6%) overall: 96.4%
96.4% is close to the paper's 97% to 98%, but the two numbers are not the same test. The paper does not say what data its figure comes from, and its "sentences" are long spans of text, while mine are single sentences cut by a simple splitter. Treat it as "the same ballpark", nothing more. Note also that a random sentence from another article is usually about a different topic, which makes NotNext easy to spot. Part 6 comes back to this.
Checking footnote 6. If C were a good summary of a sentence's meaning, similar sentences would get similar vectors. I compared the raw C vectors with cosine similarity:
plain text
3. cosine similarity of the raw final [CLS] vectors (no fine-tuning) 0.923 related A man is playing a guitar on stage. | A musician performs a song for the crowd. 0.909 related The stock market fell sharply today. | Share prices dropped a lot this afternoon. 0.766 unrelated A man is playing a guitar on stage. | The stock market fell sharply today. 0.807 unrelated Penguins cannot fly. | The invoice is due next Monday. 0.977 opposite I loved this movie. | I hated this movie.
Footnote 6, measured. Raw [CLS] vectors of any two sentences point in nearly the same direction. "I loved this movie" and "I hated this movie" get the highest similarity of all.
Every pair scores above 0.76, and the two sentences with opposite meanings score highest of all (0.977). Related pairs do score a bit higher than unrelated ones, but the gaps are small and unreliable. Footnote 6 is right: without fine-tuning, C is not a meaning vector. (Part 6 shows models fine-tuned specially to fix this.)
What is in the BooksCorpus? Its paper (Zhu et al., 2015) describes it:
And the Billion Word Benchmark, which the paper rules out, says this about itself:
Document-level text versus shuffled sentences. In a document-level corpus, the sentence after sentence 1 really is sentence 2, so IsNext pairs and long coherent spans exist. In a shuffled corpus, the "next" sentence comes from another document.
Why must the text be **document-level**? The Billion Word Benchmark is a big set of single sentences in **shuffled** order. In it, the sentence after "She opened the fridge." is some unrelated sentence from somewhere else. That makes next sentence prediction impossible to learn, and it stops the model from ever seeing long stretches of connected text. BERT needs sequences of up to 512 tokens that really belong together, so it needs whole documents.
Appendix A.2 gives the exact recipe. First, how one training example is built.
Google noticed this too. In May 2019 the BERT repository added Whole Word Masking models (its README: "New May 31st, 2019: Whole Word Masking Models"), which always mask every piece of a word together. The released masking code has a switch for it, and its comment notes that the training itself does not change: each piece is still predicted on its own.
Next, the training run itself. The sentence starts at the bottom of the left column:
batch: 256 sequences x 512 tokens = 131,072 tokens (the paper rounds this to 128,000)1,000,000 steps x 128,000 tokens = 128,000,000,000 tokensdivided by the 3.3 billion words of BooksCorpus + Wikipedia = 38.8 passes ("approximately 40 epochs")with the exact 131,072 tokens per batch: 39.7 passes
So "approximately 40 epochs" checks out. Two honest caveats about this calculation, both from the paper itself. It divides tokens by words, and WordPiece makes slightly more tokens than words. And it assumes every step uses 512 tokens, but (as we will see in a moment) 90% of the steps used sequences of 128 tokens, and the paper does not say what batch size was used then. So the true number of tokens seen is not stated anywhere in the paper; "about 40 epochs" is the paper's own rough figure.
where η (the Greek letter "eta") is the learning rate, ηmax=10−4 is the peak from the paper, and s is the step number.
The learning-rate schedule as the released code computes it: a straight climb from 0 over the first 10,000 steps, then a straight fall to 0 at step 1,000,000.
A small detail you only see in code: the decay is counted from step 0, so when warmup ends at step 10,000 the rate is already 10−4×0.99=9.9×10−5, not quite the peak.
The paper calls it "L2 weight decay". The released optimizer does not add a squared-weight penalty to the loss (a code comment explains that this interacts badly with Adam). Instead it subtracts a small fraction of each weight directly at every update, and it skips LayerNorm weights and biases. This later became known as decoupled weight decay, the "W" in the name of the AdamW optimizer.
The usual choice was ReLU, which keeps positive numbers and turns negative ones into 0. BERT uses GELU (Hendrycks and Gimpel, 2016), as OpenAI GPT did:
ReLU(x)=max(0,x),GELU(x)=x⋅Φ(x)
where:
x is the input number;
Φ(x) (the Greek letter "phi") is the probability that a random number from a standard bell curve (a normal distribution with mean 0 and spread 1) is smaller than x. It goes smoothly from 0 (for very negative x) to 1 (for very positive x).
So GELU keeps x in proportion to how large x is. Big positive inputs pass almost unchanged, big negative inputs become almost 0, and small inputs are scaled down smoothly.
Φ has no simple closed form, so the released code uses a fast approximation based on tanh:
where erf is the "error function" of statistics. Worked values:
plain text
x = -1.0: Phi(x) = 0.1587, GELU = -1.0 x 0.1587 = -0.1587; tanh formula -0.1588; ReLU +0.0x = -0.5: Phi(x) = 0.3085, GELU = -0.5 x 0.3085 = -0.1543; tanh formula -0.1543; ReLU +0.0x = +0.5: Phi(x) = 0.6915, GELU = +0.5 x 0.6915 = +0.3457; tanh formula +0.3457; ReLU +0.5x = +1.0: Phi(x) = 0.8413, GELU = +1.0 x 0.8413 = +0.8413; tanh formula +0.8412; ReLU +1.0
Read the x=+0.5 row: a standard bell-curve number is below 0.5 with probability 0.6915, so GELU lets 69% of the input through: 0.5×0.6915=0.346. ReLU would pass all of it.
ReLU and GELU. For large positive inputs they agree. GELU bends smoothly through zero and lets small negative inputs out as small negative numbers.
plain text
GELU: largest gap between the exact form x*Phi(x) and the tanh formula of the released code: 4.7e-04 x = -3.0: ReLU +0.000 GELU -0.0040 x = -1.0: ReLU +0.000 GELU -0.1587 x = -0.5: ReLU +0.000 GELU -0.1543 x = +0.0: ReLU +0.000 GELU +0.0000 x = +0.5: ReLU +0.500 GELU +0.3457 x = +1.0: ReLU +1.000 GELU +0.8413 x = +3.0: ReLU +3.000 GELU +2.9960
The released code computes Φ with a fast formula based on tanh. It differs from the exact GELU by at most 0.00047.
The training loss is the sum of the two tasks' average losses:
L=LMLM+LNSP
where LMLM is the mean cross-entropy over all masked positions in the batch and LNSP is the mean cross-entropy over all sentence pairs. (The paper says "likelihood"; training minimises the negative log-likelihood, which is the same cross-entropy.)
I checked this on a real batch: 8 sentence pairs from Wikipedia (half IsNext, half NotNext), masked with the procedure above, through BertForPreTraining:
plain text
5. pre-training loss on one real batch of 8 sentence pairs (58 masked positions) mean masked-LM loss 2.1302 mean next-sentence loss 0.0001 sum 2.1303 loss computed by BertForPreTraining: 2.1303
The two parts add up exactly to the loss the library computes. Notice how lopsided they are: next sentence prediction is nearly solved (0.0001), while filling in blanks is still hard (2.1302). After pre-training, almost all of the remaining learning signal comes from the masked LM.
The same check on the paper's own two Appendix A.1 examples, where every number can be followed:
plain text
pair 1, position 5, target "the": p = 0.8087, loss 0.2123pair 1, position 12, target "of": p = 0.9986, loss 0.0014pair 2, position 3, target "went": p = 0.0755, loss 2.5830pair 2, position 9, target "##s": p = 0.0032, loss 5.7557mean MLM loss over 4 masked positions: 2.1381NSP loss per pair: 0.000092, 0.001146; mean 0.000619total = 2.1381 + 0.000619 = 2.1387 (BertForPreTraining: 2.1387)
(The paper does not print the hidden words in A.1; the script restores the obvious ones. "##s" is the last piece of "penguins": its probability is low because, seeing "penguin [MASK]", the model mostly expects other continuations.)
The pre-training loss for the two A.1 examples as one batch. Four masked-LM losses, averaged to 2.1381; two NSP losses, averaged to 0.000619; their sum, 2.1387, is exactly what BertForPreTraining reports.
attention scores per sequence: 128^2 = 16,384, 512^2 = 262,144 -> 16x more for a 4x longer sequenceper token: 4x more attention work at length 512
The other parts of the model (the feed-forward layers) cost the same per token at any length, so attention is the part that blows up. Training mostly at length 128 saves a lot. But the position embeddings for positions 128 to 511 (Part 2) are only trained in the final 10% of steps, which is the reason for that last phase.
How big is the effect for BERT-base, exactly? Count multiply-adds per layer for a sequence of n tokens with H=768:
Q, K, V, O and the feed-forward12nH2+scores QK⊤ and mixing AV2n2H
where the first term grows like n (each token does the same matrix work) and the second like n2 (every pair of tokens).
plain text
n = 128: projections + feed-forward 905,969,664; attention scores and mixing 25,165,824 attention share of a layer's work: 2.7%; work per token 7,274,496n = 512: projections + feed-forward 3,623,878,656; attention scores and mixing 402,653,184 attention share of a layer's work: 10.0%; work per token 7,864,320512 vs 128: score matrix 16x, whole-sequence work 4.32x (4x the tokens), work per token 1.08x
Attention cost at the two lengths. One head's score table grows from 128 × 128 = 16,384 to 512 × 512 = 262,144 entries (16 times). But for BERT-base the n² part is only 2.7% of a layer's work at length 128 and 10.0% at 512, so the work per token grows by just 1.08 times.
So the paper's statement is true (the n2 part does grow 16-fold), but at n≤512 it is a modest share of the total. The bigger practical costs of long sequences on 2018 hardware were memory (every layer stores n×n attention weights per head for the backward pass) and the 4 times larger batch in tokens. The short-then-long schedule saves time either way:
The pre-training schedule of Appendix A.2. Steps 0 to 900,000 use sequences of 128 tokens; the last 100,000 steps use 512. Position embeddings 0 to 127 are trained the whole time; positions 128 to 511 only get a training signal in the last 10%.
Next, in Part 4: fine-tuning. How the same pre-trained BERT becomes a classifier, a question-answering system and a multiple-choice solver, the paper's results tables, and a real fine-tuning run.
Run it yourself
The scripts behind every number in this part are in code/papers/bert/. They need Python with torch, transformers and datasets, and download bert-base-uncased (about 440 MB) and the WikiText data the first time.
bash
pip install torch transformers datasetspython bert_part3_seeitself.py # the three tiny models (about 6 minutes per model on a laptop GPU)python bert_part3_mlm.py # masking statistics, "my dog is hairy", MLM accuracy, the loss on a batchpython bert_part3_nsp.py # the NSP head, 1,000 real pairs, raw [CLS] similaritiespython bert_part3_schedule.py # batch arithmetic, learning rate, GELU, attention cost
The real output of bert_part3_seeitself.py. The run times depend on what else the GPU is doing.The real output of bert_part3_mlm.py.The real output of bert_part3_nsp.py.The real output of bert_part3_schedule.py.