An annotated companion · AI Primer

Megatron-LM, annotated

About this page. This is a companion, not a copy. It follows the paper (version 4, March 2020) section by section, quotes at most a sentence or two per section (clearly marked and attributed), and explains everything in its own words. The paper is distributed under arXiv's standard licence, so its figures are redrawn from scratch and its tables are not reproduced; a few of their numbers are restated with attribution. Equations are reproduced with every symbol decoded; formulas that write out, in symbols, something the paper says in words are labelled as this page's own. Numbers made up for teaching are labelled illustrative. Read the original alongside: every section links to it.

How to read this page

  • Any dotted word explains itself on hover, focus or tap, and so does every symbol in every equation.
  • The heart of the paper is §3. Its redrawn Figure 3 is live, and the split playground next to it lets you cut a weight matrix by rows or by columns and see which one survives the GELU.

Every idea climbs the same ladder: everyday picture, tiny example, diagram, the math, why it matters today. This paper is the founding paper of tensor parallelism; the pretraining lesson builds its column and row splits in NumPy and checks them against the unsplit answer, and the ZeRO companion covers the other way to spread a model over GPUs, sharding its training state. Running example: two GPUs and matrices small enough to multiply by hand.

Abstract · original

“Our approach does not require a new compiler or library changes, is orthogonal and complimentary to pipeline model parallelism, and can be fully implemented with the insertion of a few communication operations in native PyTorch.”Shoeybi et al. (2019), Abstract

Everyday picture

A model too big for one GPU has to be cut up. One way is to give each GPU a few whole layers, like an assembly line. Megatron-LM cuts the other way: every layer is sliced lengthwise and each GPU in a server does its slice of every layer, like several cooks each chopping part of the same pile of onions, meeting briefly after each dish.

What the paper claims

  • A simple intra-layer model parallelism for transformers: split the matrix multiplies so each layer needs only two all-reduces in the forward pass and two in the backward.
  • Scaling: an 8.3-billion-parameter GPT-2 on 512 GPUs at 15.1 petaflops, 76% of perfect scaling from a single GPU that itself runs at 39 teraflops (30% of its peak).
  • Results: the 8.3B GPT-2 set records on WikiText103 (perplexity 10.8 against 15.8) and LAMBADA (66.5% against 63.2%); a 3.9B BERT set a record on RACE (90.9% against 89.4%).
  • A finding about architecture: BERT-style models stop improving with size unless the layer normalization and the residual connection are rearranged.

Why it matters today

The column-then-row split of this paper became the standard way to cut a transformer layer across GPUs, and it is still kept within one server, where the links are fastest. The ZeRO paper, which followed weeks later, built on it.

1 Introduction · original

Everyday picture

Bigger language models were getting better, but the memory of one GPU was not growing with them, and optimizers like Adam add extra memory for every weight. Earlier ways to split a model, GPipe and Mesh-TensorFlow, needed the model rewritten or a new compiler. The authors want something an ordinary PyTorch user can add in an afternoon.

Tiny example: what 76% means

The baseline is a 1.2-billion-parameter model on one 32 GB V100, sustaining 39 teraflops. Perfect scaling to 512 GPUs would be 512 × 39 = 19,968 teraflops, about 20 petaflops. The paper measures 15.1 petaflops: 15.1 / 19.97 = 75.6%, the “76% scaling efficiency”.

In words: “scaling efficiency is the throughput you got on n GPUs, divided by n times what one GPU achieved alone.” (This page's formula for what the paper describes in words.)

With the numbers: 15.1 × 1015 / (512 × 39 × 1012) = 0.756. The baseline itself is 30% of the GPU's peak, so the whole run delivers about 0.3 × 0.756 ≈ 23% of the cluster's peak arithmetic.

In Python:

F_1, n, F_n = 39e12, 512, 15.1e15
# perfect scaling, in petaflops
n * F_1 / 1e15  # → 19.968
# E = F_n / (n F_1)
E = F_n / (n * F_1)
round(E, 3)  # → 0.756
# share of the cluster's peak, given the baseline's 30%
round(0.3 * E, 2)  # → 0.23

Why it matters

A scaling number is only as good as its baseline. Choosing a baseline that already runs at 30% of peak, the authors call it a strong one, is what makes 76% meaningful: it is easy to scale well from a slow single-GPU program.

2 Background and challenges · original

Everyday picture

§2.1 recalls where pretraining came from: first reusing learned word vectors, then reusing whole pretrained networks and fine-tuning them end to end. Each step moved more of the model, and more of the compute, into pretraining.

2.2 Transformer language models and multi-head attention · original

Everyday picture

The paper works with both families of transformer: GPT-2, a decoder that writes left to right, and BERT, an encoder that reads in both directions. Each layer is the same two steps, self-attention (tokens exchange information) then a two-layer feed-forward network (each token alone), each wrapped in a residual connection. Both use GELU, and both normalize the input to each step rather than its output as the original transformer did.

Tiny example

GELU (the paper writes GeLU) is a smooth ReLU: x times the chance that a standard bell-curve variable falls below x. GELU(1) = 1 × 0.841 = 0.841; GELU(−1) = −1 × 0.159 = −0.159; GELU(2) = 1.954. Big positive inputs pass almost unchanged, big negative ones go to almost zero.

In Python:

import math
def gelu(x):
    # x times Φ(x), the bell curve's cumulative probability
    return x * 0.5 * (1 + math.erf(x / math.sqrt(2)))
round(gelu(1), 3), round(gelu(-1), 3), round(gelu(2), 3)  # → (0.841, -0.159, 1.954)

Why it matters

That GELU is not linear, GELU(a + b) ≠ GELU(a) + GELU(b), is the fact §3's whole design turns on. The transformer lesson builds the layer (TransformerBlock, with gelu).

2.3 Data and model parallelism in deep learning · original

Everyday picture

Data parallelism gives every GPU the whole model and a slice of the batch; growing the batch with the number of GPUs is weak scaling. It has a hard limit: the model must fit on one GPU. Model parallelism removes the limit, in one of two ways: pipeline parallelism (whole layers per GPU, fed by micro-batches, with idle bubbles) or distributed tensor computation (each operation split across GPUs), which is this paper's choice.

Tiny example

A 1.2-billion-parameter model trained with Adam in mixed precision needs about 16 bytes per parameter, 19.2 GB, before activations: it fits on a 32 GB V100. The 8.3B model needs 133 GB: no amount of data parallelism alone will fit it.

In Python:

# 16 bytes per parameter for Adam in mixed precision (the ZeRO paper's count), in GB
16 * 1.2e9 / 1e9  # → 19.2
round(16 * 8.3e9 / 1e9)  # → 133

Why it matters

The two kinds of model parallelism are complementary, as the paper stresses: pipeline across servers, tensor inside a server. The pretraining lesson's section 3 builds all three and shows where each belongs.

3 Model parallel transformers · original

Everyday picture

A transformer layer is two blocks, attention and an MLP, and each block is two matrix multiplies (the paper calls them GEMMs) with something in between. The trick is to cut the first multiply one way and the second the other way, so the pieces fit together with no talking in the middle.

3.1 The MLP block · original

Tiny example: two ways to cut

One token, X = [1, 2], and a weight matrix A = [[1, −1], [−1, 1]]. Unsplit: XA = [1 − 2, −1 + 2] = [−1, 1], and GELU gives [−0.159, 0.841].

In words: “multiply the tokens by the first weight matrix, then apply GELU to every number.” (The paper's equation 1.)

Cut by rows. GPU 1 takes A's first row and X's first entry; GPU 2 the second row and second entry.

In words: “cut by rows, each GPU gets a partial sum of every output; the partial sums must be added before GELU, because GELU of a sum is not the sum of GELUs.” (Equation 2 and the sentence after it.)

With the numbers: GPU 1: 1 × [1, −1] = [1, −1]. GPU 2: 2 × [−1, 1] = [−2, 2]. Added first: [−1, 1], GELU → [−0.159, 0.841], correct. GELU first, then added: [0.841 − 0.046, −0.159 + 1.954] = [0.796, 1.796], wrong. So the GPUs would have to talk before the GELU.

Cut by columns. GPU 1 takes A's first column, GPU 2 the second; both need all of X.

In words: “cut by columns, each GPU computes some of the output numbers completely, so it can apply GELU to them on its own.” (Equation 3.)

With the numbers: GPU 1: X · [1, −1]ᵀ = −1, GELU → −0.159. GPU 2: X · [−1, 1]ᵀ = 1, GELU → 0.841. Side by side: [−0.159, 0.841], exactly the unsplit answer, with no communication.

In Python:

import math
def gelu(x):
    return x * 0.5 * (1 + math.erf(x / math.sqrt(2)))
X = [1, 2]
A = [[1, -1], [-1, 1]]
# unsplit: Y = GELU(XA)
XA = [sum(X[i] * A[i][j] for i in range(2)) for j in range(2)]
XA  # → [-1, 1]
[round(gelu(v), 3) for v in XA]  # → [-0.159, 0.841]
# by rows: X_1 A_1 and X_2 A_2 are partial sums of every output
p1 = [X[0] * a for a in A[0]]
p2 = [X[1] * a for a in A[1]]
p1, p2  # → ([1, -1], [-2, 2])
# GELU before adding: wrong
[round(gelu(a) + gelu(b), 3) for a, b in zip(p1, p2)]  # → [0.796, 1.796]
# by columns: each GPU owns whole outputs, GELU runs locally
Y1 = gelu(X[0] * A[0][0] + X[1] * A[1][0])
Y2 = gelu(X[0] * A[0][1] + X[1] * A[1][1])
round(Y1, 3), round(Y2, 3)  # → (-0.159, 0.841)
Cut the first matrix A:

Reading it: the playground runs the same example on two simulated GPUs. Each card shows what one GPU holds and computes; the last line compares the result with the unsplit Y = GELU(XA) = [−0.159, 0.841]. Cut by rows and each GPU's product is a partial sum of both outputs, so applying GELU locally gives the wrong answer: the GPUs would have to add their partial sums first, a synchronization point. Cut by columns and each GPU finishes one whole output, so GELU runs locally and the answer is exact. That is why Megatron-LM cuts the first matrix by columns.

The second matrix: cut by rows, add once

After GELU, GPU 1 holds Y1 and GPU 2 holds Y2: the output is already split by columns. That is exactly the input a row-split of the second matrix B wants. Each GPU multiplies its piece by its rows of B, and one all-reduce adds the two partial results.

In words: “each GPU multiplies its own half of the hidden units by the matching rows of B; adding the two results (one all-reduce) gives the block's output.” (This page's formula for what the paper's Figure 3a draws.)

With the numbers: with B = [[2, 0], [1, 1]]: GPU 1 computes −0.159 × [2, 0] = [−0.317, 0]; GPU 2 computes 0.841 × [1, 1] = [0.841, 0.841]; the all-reduce adds them to [0.524, 0.841], the same as the unsplit YB.

In Python:

import math
def gelu(x):
    return x * 0.5 * (1 + math.erf(x / math.sqrt(2)))
Y1, Y2 = gelu(-1), gelu(1)
B = [[2, 0], [1, 1]]
# each GPU: its hidden unit times its row of B
part1 = [Y1 * b for b in B[0]]
part2 = [Y2 * b for b in B[1]]
[round(v, 3) for v in part1], [round(v, 3) for v in part2]  # → ([-0.317, -0.0], [0.841, 0.841])
# the all-reduce: add the partial results
[round(a + b, 3) for a, b in zip(part1, part2)]  # → [0.524, 0.841]
# unsplit Y B, for comparison
[round(Y1 * B[0][j] + Y2 * B[1][j], 3) for j in range(2)]  # → [0.524, 0.841]
GPU 1 GPU 2 X f X·A₁ X·A₂ GELU GELU Y₁·B₁ Y₂·B₂ g Drop-out, Z

Hover or tap a block. Start with X on the left.

Figure 3a of the paper, redrawn: the MLP block with 2-way model parallelism. Based on Shoeybi et al. (2019), Figure 3a.

Reading it: data flows left to right. The same X enters both GPUs (the dashed boxes) through f. Each GPU multiplies by its columns of A, applies GELU to the hidden units it owns, and multiplies by its rows of B, all without looking at the other GPU. The two partial outputs meet at g, the block's one all-reduce, before dropout. Read f and g as a pair: in the forward pass f does nothing and g adds; in the backward pass g does nothing and f adds the two GPUs' gradients with respect to X. One all-reduce each way per block.

f and g, in code

The paper's Code 1 defines f as a PyTorch autograd function whose forward returns x unchanged and whose backward all-reduces the gradient; g is the mirror image. They are conjugates: each is the other with forward and backward swapped.

OperatorForward passBackward passWhere
fidentity: pass X to every GPUall-reduce: add each GPU's gradient for Xentry of a parallel block
gall-reduce: add the partial outputsidentity: pass the gradient to every GPUexit of a parallel block

Why it matters today

Column-parallel then row-parallel is the tensor-parallel MLP. The pretraining lesson builds column_parallel_matmul, row_parallel_matmul and tensor_parallel_mlp, and checks that each matches the unsplit answer.

3.2 The attention block · original

Everyday picture

Multi-head attention is already split: each head is an independent small attention with its own query, key and value projections. Hand each GPU a group of whole heads and it can run them start to finish alone.

Tiny example

The 8.3B scaling model has 32 heads on 8 GPUs: 4 heads per GPU. The query, key and value projections are cut by columns so that GPU i's columns are exactly its heads' columns; softmax and the weighted sum stay local. The output projection after attention is cut by rows, and one all-reduce (g) adds the GPUs' partial results, exactly like the MLP.

In Python:

# heads per GPU, and each head's width, for the 8.3B scaling model
heads, gpus, hidden = 32, 8, 3072
heads // gpus, hidden // heads  # → (4, 96)
GPU 1: heads 1 to h/2 GPU 2: heads h/2 + 1 to h X f Q₁ K₁ V₁ Q₂ K₂ V₂ soft-max soft-max Y₁·B₁ Y₂·B₂ g Drop-out, Z

Hover or tap a block. The shape matches the MLP figure above on purpose.

Figure 3b of the paper, redrawn: the self-attention block with 2-way model parallelism. Based on Shoeybi et al. (2019), Figure 3b.

Reading it: the same skeleton as the MLP. X enters both GPUs through f. The column-split projections give each GPU the queries, keys and values of its own heads; each GPU runs softmax attention for those heads (with attention dropout) entirely locally; the row-split output projection turns each GPU's heads into a partial output; g adds them. The paper's word for what this does to each block is that it “fuses” two matrix multiplies and removes the synchronization point between them.

Why it matters today

Splitting by heads puts a ceiling on this kind of parallelism: you cannot use more GPUs than heads, and the heads must divide evenly. The paper's appendix also finds that more, narrower heads scale a little worse (Appendix D, below). The Attention Is All You Need companion explains the heads themselves.

3.3 Four all-reduces per layer · original

Everyday picture

Put the two blocks together and a whole transformer layer talks four times per training step: once after attention and once after the MLP on the way forward, and once at the start of each block on the way back.

Tiny example

Each all-reduce carries one activation tensor: batch × sequence × hidden numbers. For the 8.3B model with 8 sequences of 1,024 tokens at width 3,072, that is 25.2 million numbers, about 50 MB in 16-bit floats. Four per layer, 72 layers: 288 all-reduces per step, 14.5 GB of messages.

In words: “each layer all-reduces four activation-sized messages per step, two forward and two backward.” (This page's formula for the paper's Figure 4.)

With the numbers: 4 × 8 × 1,024 × 3,072 = 100.7 million numbers per layer; at 2 bytes, 201 MB; times 72 layers, 14.5 GB of messages per step.

In Python:

b, s, H, layers = 8, 1024, 3072, 72
message = b * s * H
round(message / 1e6, 1)  # → 25.2
# C_layer = 4 b s H
C_layer = 4 * message
round(C_layer / 1e6, 1)  # → 100.7
# bytes per step at 2 bytes per number, in GB
round(C_layer * 2 * layers / 1e9, 1)  # → 14.5

Why it matters today

Note what the traffic scales with: the activations, not the weights. Tensor parallelism's cost grows with batch and sequence length, and it happens inside every layer on the critical path, which is why it is kept on the fastest links; the hardware lesson puts numbers on those links.

3.4 The embedding and the loss · original

Everyday picture

The vocabulary table is big (GPT-2's has 50,257 rows) and it is used twice: to look tokens up at the input and, with the same weights (weight tying), to score every token at the output. Megatron-LM splits it by vocabulary: each GPU keeps a slice of the tokens.

Tiny example

At the output, each GPU can score its own slice of the vocabulary. The obvious next step, gathering all the scores (logits) onto every GPU for the cross-entropy loss, would move b × s × v numbers. Instead each GPU computes its slice's part of the loss, and only b × s numbers are exchanged. With b = 8, s = 1,024 and the padded vocabulary of 51,200, that is 419 million numbers against 8,192: 51,200 times less.

In words: “sending losses instead of logits divides the traffic by the size of the vocabulary.” (This page's formula for the paper's §3 argument.)

With the numbers: 8 × 1,024 × 51,200 = 419,430,400 logits against 8 × 1,024 = 8,192 losses; the ratio is v = 51,200.

In Python:

b, s, v = 8, 1024, 51200
V_logits = b * s * v
V_loss = b * s
V_logits, V_loss  # → (419430400, 8192)
V_logits // V_loss  # → 51200

Where does 51,200 come from? For fast matrix multiplies, each GPU's slice of the vocabulary should be a multiple of 128; with up to 8 GPUs, the vocabulary is padded up to a multiple of 128 × 8 = 1,024 (§5.1).

In words: “round the vocabulary up to the next multiple of 128 times the number of GPUs.”

With the numbers: 50,257 / 1,024 = 49.08, rounded up to 50; 50 × 1,024 = 51,200, so 943 padding rows that no real token uses.

In Python:

import math
v, n = 50257, 8
v_pad = math.ceil(v / (128 * n)) * 128 * n
v_pad, v_pad - v  # → (51200, 943)

Why it matters today

Vocabularies have since grown several times larger, so the same trick, never gathering all the logits in one place, matters more now; padding the vocabulary to a friendly multiple is routine.

3.5 Duplicate the cheap parts · original

Everyday picture

Some steps are too cheap to be worth splitting: layer norm, dropout and the residual addition. Rather than have one GPU compute them and send the result, every GPU simply computes them itself on its own copy. Arithmetic is cheap; waiting is not.

Tiny example

A layer norm has 2 × H parameters: 6,144 for H = 3,072, against the roughly 113 million weights of the layer's matrix multiplies (12H², from §5.1 below). Keeping 8 copies of the norm costs almost nothing, and because every GPU holds every value it needs, no updated parameters ever have to be sent: each GPU optimizes its own set.

In Python:

H = 3072
# a layer norm's gain and bias
2 * H  # → 6144
# the layer's matrix-multiply weights, about 12 H², in millions
round(12 * H ** 2 / 1e6)  # → 113

Why it matters

“Recompute or duplicate rather than communicate” recurs throughout large-scale training: activation checkpointing and the FlashAttention backward pass make the same bet.

4 Setup · original

Everyday picture

Before the scaling experiments, the paper fixes the ingredients: which text, cleaned how, and the training recipe for two model families, GPT-2 (left to right) and BERT (masked, both directions).

4.1 Training dataset · original

Everyday picture

Four large sources (Wikipedia, CC-Stories, RealNews and OpenWebText) are pooled; BERT also gets BooksCorpus, which GPT-2 skips because it overlaps with the LAMBADA test. Wikipedia articles that appear in the WikiText103 test set are removed, so the test stays unseen.

Tiny example

Two cleaning steps: drop every document under 128 tokens, and remove near-duplicates with locality-sensitive hashing at a Jaccard similarity above 0.7. Two documents whose combined set of distinct word-shingles (short runs of consecutive words) has 10 members, 7 of them in both, have Jaccard similarity 7/10 = 0.7: not above the threshold, so both stay. With 8 in both, one of them goes. What remains: 174 GB of text.

In Python:

# Jaccard = shared shingles / all distinct shingles
7 / 10, 8 / 10  # → (0.7, 0.8)
# removed only when above 0.7
[j > 0.7 for j in (7 / 10, 8 / 10)]  # → [False, True]

Why it matters today

The same recipe, MinHash and LSH at a Jaccard threshold around 0.7, is built step by step in the pretraining lesson, whose near_duplicate_pairs uses 0.7 as its default.

4.2 Training optimization and hyperparameters · original

Everyday picture

The recipe is standard for its time: mixed precision with dynamic loss scaling, Adam with weight decay 0.01, gradient clipping at a global norm of 1.0, dropout 0.1, and activation checkpointing after every layer. GPT-2 trains on 1,024-token sequences, batch 512, for 300,000 steps, with a peak learning rate of 1.5 × 10−4 after 3,000 steps of warmup, then cosine decay down to 10−5. BERT uses batch 1,024 and a learning rate of 10−4, warmed up over 10,000 steps and decayed linearly over 2 million.

Tiny example: shrinking the last layer of each block

One detail is new. Weights start random with spread 0.02 (the paper writes N(0, 0.02); this page reads 0.02 as the standard deviation), but the weights just before each residual addition are scaled down by 1/√(2N), N being the number of layers. Each layer adds two outputs (attention's and the MLP's) into the residual stream, so 2N contributions pile up. Random contributions add their variances: 144 of them with variance v sum to variance 144v. Scale each by 1/√144 = 1/12 and the sum's variance is back to v.

In words: “the layers that write into the residual stream start with their spread divided by the square root of how many such writes there are.”

With the numbers: the 8.3B model has N = 72 layers, so 2N = 144 residual writes and σout = 0.02 / 12 = 0.00167.

In Python:

import math
sigma, N = 0.02, 72
# 2N writes into the residual stream
2 * N  # → 144
sigma_out = sigma / math.sqrt(2 * N)
round(sigma_out, 5)  # → 0.00167
# summed variance of 144 writes, before and after scaling
round(144 * sigma ** 2, 4), round(144 * sigma_out ** 2, 4)  # → (0.0576, 0.0004)

Why it matters today

Scaling residual-branch outputs by depth is a common way to keep very deep stacks stable at the start of training. The deep networks lesson explains why the spread of the signal matters (init_std), and the optimizers lesson builds the schedule (warmup_cosine) and the clipping (clip_by_global_norm).

5 Experiments · original

Everyday picture

Up to 32 DGX-2H servers, 512 V100 GPUs with 32 GB each. Inside a server GPUs talk over NVSwitch at 300 GB/s; between servers, 8 InfiniBand adapters per server give 100 GB/s. The split in speed decides the split in work.

Tiny example

Each server holds 16 GPUs sharing 100 GB/s to the outside: 6.25 GB/s each, against 300 GB/s inside, 48× slower. (The ZeRO companion quotes the same network per link: 12.5 GB/s × 8 adapters = 100 GB/s.) Illustratively, the 14.5 GB of tensor-parallel messages per step from §3.3 would take about 0.05 s at 300 GB/s but 2.3 s at 6.25 GB/s, ignoring overlap and the ring's factor of about 2.

In Python:

# per-GPU share of the links out of a 16-GPU server, GB/s
100 / 16  # → 6.25
300 / 6.25  # → 48.0
# 14.5 GB of messages per step, in seconds (illustrative)
round(14.5 / 300, 2), round(14.5 / 6.25, 1)  # → (0.05, 2.3)

Why it matters

That factor is why the paper keeps model parallelism at 8 GPUs, inside one server, and uses data parallelism across servers.

5.1 Scaling analysis · original

Everyday picture

The test is weak scaling of the model: double the GPUs and double the model, about a billion parameters per GPU, and see whether each GPU stays as busy. Four GPT-2 configurations go from 1.2B on 1 GPU to 8.3B on 8, all with 96 numbers per attention head; then each runs again with 64-way data parallelism, up to 512 GPUs.

Tiny example: checking the parameter counts

A transformer layer holds about 12 × H² weights (4H² in attention, 8H² in the MLP), and the embedding adds v × H. For the largest configuration, 72 layers of H = 3,072 with the padded vocabulary: 12 × 72 × 3,0722 + 51,200 × 3,072 = 8.15 + 0.16 = 8.31 billion, the paper's 8.3B. (The rule of thumb is this page's check; the paper gives only the totals.)

In words: “twelve H-squared weights per layer, times the layers, plus one H-wide row per vocabulary entry.”

With the numbers: all four rows of the paper's Table 1: (40 layers, H = 1,536) → 1.21B; (54, 1,920) → 2.49B; (64, 2,304) → 4.19B; (72, 3,072) → 8.31B, against the paper's 1.2, 2.5, 4.2 and 8.3.

In Python:

v = 51200
# (layers, hidden size) from Table 1
configs = [(40, 1536), (54, 1920), (64, 2304), (72, 3072)]
# P ≈ 12 L H² + v H, in billions
[round((12 * L * H ** 2 + v * H) / 1e9, 2) for L, H in configs]  # → [1.21, 2.49, 4.19, 8.31]
# every configuration keeps 96 numbers per head
[H // heads for H, heads in [(1536, 16), (1920, 20), (2304, 24), (3072, 32)]]  # → [96, 96, 96, 96]

The results

Weak scaling efficiency against the 1-GPU, 1.2B baseline, restated from Shoeybi et al. (2019), §1 and §5.1.1
SetupEfficiency
8.3B, 8-way model parallel (8 GPUs, batch 8)77%
8.3B, 8-way model × 64-way data parallel (512 GPUs, batch 512)74% (§5.1.1); 76% (Abstract, §1)

The paper gives two figures for the same 512-GPU run: 76% in the abstract and introduction, 74% in §5.1.1. The arithmetic from its own throughput numbers (§1 above) gives 75.6%. Either way, data parallelism on top costs only a few points, because gradients are all-reduced once per step while the model-parallel traffic happens inside every layer.

Why it matters today

Weak scaling of the model, rather than the batch, was the right test for the question people cared about: can a bigger model be trained at nearly the same speed per GPU? The scaling-laws companion explains why people wanted the bigger model in the first place.

5.2 Language modeling results using GPT-2 · original

Everyday picture

Three GPT-2 models, 355M, 2.5B and 8.3B, are trained identically for 300,000 steps. Bigger models learn faster and end lower: the 8.3B reaches a validation perplexity of 9.27.

Tiny example: how long is training?

One pass over the data (an epoch) is 68,507 steps, so 300,000 steps is 4.4 epochs. Each step is 512 sequences of 1,024 tokens, 524,288 tokens, so training sees about 157 billion tokens. For the 8.3B model on 512 GPUs, at 2.10 days per epoch, that is about 9.2 days.

In Python:

steps, per_epoch = 300_000, 68_507
round(steps / per_epoch, 2)  # → 4.38
# tokens per step, and in total (billions)
512 * 1024  # → 524288
round(steps * 512 * 1024 / 1e9)  # → 157
# days for the 8.3B model at 2.10 days per epoch
round(2.10 * steps / per_epoch, 1)  # → 9.2

Reading it: two scores for each model with no fine-tuning (zero-shot), restated from the paper's Table 3. The top group is WikiText103 perplexity, where shorter is better; the bottom group is LAMBADA accuracy, where longer is better. The striped bars are the previous best results. Both scores improve steadily with size, and the 8.3B model passes the previous best on both: perplexity 10.81 against 15.79, accuracy 66.51% against 63.24%. The 2.5B model already beats the old perplexity record.

Checking the test was unseen

The authors count how many 8-word sequences of each test set also appear in the training data: at most 10.8% for WikiText103 (whose own training set already overlaps its test set by 9.09%) and 1.4% for LAMBADA. The paper judges these consistent with earlier work. The same paragraph notes that Turing-NLG, a 17-billion-parameter model, was later trained with Megatron; the ZeRO companion tells that story.

Why it matters today

“Bigger is better, and larger models converge faster” is an early instance of what the scaling laws made quantitative the next year; the GPT-3 companion takes the same curve to 175 billion.

5.3 Bi-directional transformer results using BERT · original

“We further investigated this behaviour and empirically demonstrated that rearranging the order of the layer normalization and the residual connections as shown in Figure 7 is critical to enable the scaling of the BERT-style models beyond BERT-Large.”Shoeybi et al. (2019), §5.3

Everyday picture

Earlier work (ALBERT) had found that BERT got worse beyond 336M parameters. Megatron-LM finds the culprit is plumbing: where the residual connection branches off. If the skip path carries the normalized signal, every layer re-normalizes the running total; if it carries the raw signal, the total flows untouched from bottom to top.

(a) original BERT (b) rearranged input LayerNorm Attention + LayerNorm MLP + output input LayerNorm Attention + LayerNorm MLP + output

Hover or tap a part. Compare where the two thick skip arrows start in (a) and in (b).

Figure 7 (left) of the paper, redrawn: the two layer arrangements. Based on Shoeybi et al. (2019), Figure 7.

Reading it: both columns have the same boxes in the same order, bottom to top: layer norm, attention, add, layer norm, MLP, add. The only difference is where each thick skip arrow leaves the main line. In (a) it leaves after the layer norm, so what is carried forward and added back is the normalized signal: each block re-normalizes the running total. In (b) it leaves before the layer norm, so the raw input travels up untouched and each block only adds a correction to it. The paper's training curves (Figure 7, right) show (a) training well at 336M parameters but becoming unstable at 752M, while (b) is stable and reaches a lower loss.

Tiny example: the scores

With arrangement (b), bigger is better again. On a 3% held-out set, masked-language-model perplexity falls from 1.58 (336M) to 1.30 (1.3B) to 1.16 (3.9B). On RACE, reading comprehension from exams, the 3.9B model scores 89.5% alone and 90.9% as a 5-model ensemble, against the previous best ensemble's 89.4%.

In Python:

# held-out perplexity for 336M, 1.3B, 3.9B
ppl = [1.58, 1.30, 1.16]
# each step's relative improvement
[round(1 - b / a, 3) for a, b in zip(ppl, ppl[1:])]  # → [0.177, 0.108]
# RACE ensemble: points above the previous best
round(90.9 - 89.4, 1)  # → 1.5

Why it matters today

Arrangement (b) is pre-norm, the layout almost every large transformer now uses. The transformer lesson compares pre-norm and post-norm, and the layer norm companion explains the normalization itself.

6 Conclusion and future work · original

Everyday picture

The authors list what comes next: better optimizer memory efficiency, and, for models over 16 billion parameters, more memory than one 16-GPU DGX-2H server offers, so a mix of intra-layer and inter-layer (pipeline) parallelism across servers. Distilling the large models into small ones is on the list too.

Tiny example

16 billion parameters at 16 bytes each is 256 GB of training state, half of the 512 GB that 16 GPUs of 32 GB hold, before any activations. Past that, one server is not enough and some parallelism has to cross the slow links.

In Python:

# training state for 16B parameters, and the server's memory, in GB
16 * 16, 16 * 32  # → (256, 512)

Why it matters

Both directions happened quickly: ZeRO attacked optimizer memory (see the ZeRO companion), and later Megatron-LM work composed tensor, pipeline and data parallelism across thousands of GPUs.

Appendices B, D and E · original

B.1 Hybrid model and data parallelism

Everyday picture: 512 GPUs arranged as 64 teams of 8. Each team holds one complete copy of the model, cut 8 ways; a GPU's opposite numbers in the other 63 teams hold the same slice, and they average their gradients together.

9

the chosen GPUits model-parallel group (one copy of the model)its data-parallel group (same slice, other copies)

Reading it: each small block of 8 cells is one model-parallel group, 64 in all, numbered left to right and top to bottom; the paper's Figure 8 draws the same grouping. Drag the slider to pick a GPU. Its whole block lights purple: the 8 GPUs that together hold one copy of the model and all-reduce inside every layer. The orange cells, one in every block at the same position, are its data-parallel group: the 64 GPUs that hold the same slice of the weights and all-reduce their gradients once per step. For GPU 9 that is GPUs 1, 9, 17, and so on up to 505, as in the paper.

In Python:

# GPUs numbered 1 to 512, model-parallel groups of 8
mp = 8
g = 9
# which model-parallel group, and which position in it
group = (g - 1) // mp + 1
position = (g - 1) % mp + 1
group, position  # → (2, 1)
# its data-parallel group: the same position in every model-parallel group
dp = [position + mp * k for k in range(512 // mp)]
dp[:3], dp[-1], len(dp)  # → ([1, 9, 17], 505, 64)

B.2 Random numbers for dropout

Everyday picture: dropout outside the parallel regions (before the residual additions) acts on tensors every GPU holds a full copy of, so all 8 GPUs must drop the same numbers, or their copies drift apart. Dropout inside a parallel region acts on different slices, so each GPU must drop different numbers, or the pattern repeats across slices. The fix: one generator seeded identically on every GPU for the first kind, and a second generator seeded differently per GPU for the second.

Tiny example (illustrative): with dropout 0.5 on a 4-number residual on 2 GPUs, both must zero, say, positions 2 and 3; if GPU 1 zeroed 2 and 3 and GPU 2 zeroed 1 and 4, the two copies of the residual would disagree from then on.

D Further scaling analysis

Attention heads: at 8.3B with 8-way model parallelism, 16 heads of width 192 scale at 82%, 24 heads of 128 at 80%, and 32 heads of 96 at 77% (the paper's Table 7): more heads mean smaller matrix multiplies and bigger softmaxes.

Strong scaling: the 1.2B model at a fixed batch of 8, spread over more GPUs. Two GPUs are 1.64× faster; eight, 2.98×.

Hover the chart, or focus it and use the arrow keys, to read the speedup at each GPU count.

Reading it: the horizontal axis is the number of GPUs splitting one fixed 1.2B model, the vertical axis how many times faster a training step runs. The straight line is perfect scaling; the lower curve is the paper's measurement (its Table 8). The gap widens quickly: at 8 GPUs, 2.98× out of a possible 8, 37% efficiency. With the model and batch fixed, each GPU's share of the work shrinks while the all-reduces stay the same size, so talking takes over. That is why the paper's main results use model parallelism to fit bigger models, not to speed up small ones.

In words: “strong scaling efficiency is the speedup divided by the number of GPUs.” (This page's formula.)

With the numbers: 1.64 / 2 = 82%, 2.34 / 4 = 58.5%, 2.98 / 8 = 37.2%.

In Python:

# speedups for 2, 4 and 8 GPUs (Table 8)
S = {2: 1.64, 4: 2.34, 8: 2.98}
{n: round(S_n / n, 3) for n, S_n in S.items()}  # → {2: 0.82, 4: 0.585, 8: 0.372}

E Evaluating perplexity and LAMBADA

Everyday picture: perplexity is e raised to the average surprise per token. But “per token” depends on how you cut the text into tokens: GPT-2's subword tokenizer makes 270,329 tokens of the WikiText103 test set, while earlier models counted 245,566 word-level tokens. To compare fairly, the total surprise is divided by the original count, To.

In words: “add up the surprise, minus the log of the probability the model gave each actual next token, over all T subword tokens; divide by the number of word-level tokens To; and exponentiate.” (The paper's equation 4.)

With the numbers: illustratively, if the model's average surprise were 2.2 per subword token, the total over T = 270,329 tokens is 594,724; divided by To = 245,566 that is 2.422 per word, and the perplexity is e2.422 = 11.27, rather than the e2.2 = 9.03 a per-subword average would give. Normalizing by To makes the score comparable with word-level models, and higher.

In Python:

import math
T, T_o = 270_329, 245_566
# illustrative: average surprise of 2.2 per subword token
total = 2.2 * T
round(total)  # → 594724
round(total / T_o, 3)  # → 2.422
# PPL = exp(total / T_o), against the per-subword exp(2.2)
round(math.exp(total / T_o), 2), round(math.exp(2.2), 2)  # → (11.27, 9.03)

Two more details. The model sees at most 1,024 tokens of context, so the test text is read in overlapping windows that advance 32 tokens at a time, scoring only each window's last 32 tokens: every scored token has at least 992 tokens of context, at 1/32 of the cost of a fresh window per token. And LAMBADA counts a word as right only if every subword token of it is predicted correctly.

In Python:

window, overlap, T = 1024, 32, 270_329
# context before the first scored token of a window
window - overlap  # → 992
# windows needed, against one window per token
T // overlap + 1  # → 8448

Why it matters

Perplexities computed with different tokenizers are not comparable until they are put on the same denominator; this appendix is a clear worked instance. The losses lesson builds cross-entropy and perplexity from scratch.

What happened next

DevelopmentWhat it does
Megatron-LM codeThe open-source training code the paper released
ZeRO (Rajbhandari et al., 2019)Shards the optimizer state, gradients and weights across data-parallel GPUs, and combines with Megatron-LM's model parallelism to train 100B-parameter models
Megatron-LM at scale (Narayanan et al., 2021)Composes tensor, pipeline and data parallelism to train a trillion-parameter model on 3,072 GPUs at 52% of peak
Reducing activation recomputation (Korthikanti et al., 2022)Counts a tensor-parallel transformer's activation memory and cuts what has to be stored or recomputed
On layer normalization in the transformer (Xiong et al., 2020)Explains why the pre-norm arrangement of Figure 7(b) trains more stably

To build the pieces yourself, the pretraining lesson splits a matrix multiply both ways and assembles the tensor-parallel MLP, then adds pipeline and data parallelism; the hardware lesson explains why tensor parallelism lives inside one machine.

Glossary

Every term with hover guidance on this page, in one place.