Skip to main content

Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism

Megatron-LM splits each transformer layer across GPUs with two all-reduces forward and two backward, and trains an 8.3B-parameter GPT-2 on 512 V100s.

TL;DR

  • By 2019 the largest language models no longer fit on one GPU together with their gradients and optimizer state. Data parallelism cannot help with that, because it keeps a full copy of the model on every GPU.
  • Megatron-LM splits the matrix multiplies inside each transformer layer across GPUs. In both the attention block and the MLP, the first multiply is split by columns and the second by rows, so the GPUs work independently in between and a layer needs only two all-reduces in the forward pass and two in the backward pass.
  • The method is a few communication calls added to an ordinary PyTorch transformer: two operators, f and g, that are the identity in one direction and an all-reduce in the other. It needs no compiler and no new framework.
  • With 8-way model parallelism and 64-way data parallelism, an 8.3-billion-parameter GPT-2 trains on 512 V100 GPUs and sustains 15.1 PetaFLOPs, which the paper reports as 76% scaling efficiency against a single-GPU baseline. Its results were state of the art at the time on WikiText103 (10.8 perplexity), LAMBADA (66.5% accuracy) and, with a 3.9-billion-parameter BERT, RACE (90.9%).

The model has to fit somewhere

Data parallelism gives every GPU a complete copy of the model and a different part of the batch. It raises throughput, but the whole model still has to fit on each GPU, along with its gradients and the extra state Adam keeps for every parameter. In the paper, the 1.2-billion-parameter configuration is the one that fits on a single 32 GB V100, and all of its models already train with mixed precision and with activation checkpointing after every transformer layer.

Two ways to split the model itself already existed. Pipeline parallelism, as in GPipe, places groups of layers on different devices and passes activations from one to the next; the paper notes that it needs extra scheduling logic and loses efficiency to pipeline bubbles. Mesh-TensorFlow splits individual tensor operations across devices, but through its own language and compiler. Megatron-LM takes the second idea and applies it by hand to one architecture, with a few targeted changes to an existing PyTorch transformer.

The paper calls the result intra-layer model parallelism. It is now usually called tensor parallelism, and it is one of the three strategies compared in data vs tensor vs pipeline parallelism.

One layer on two GPUs

The method fits in one picture of a transformer layer on two GPUs. The weight matrices inside the attention block and the MLP are split, so each GPU holds a different part of them. Layer normalization, dropout and the residual add are duplicated, so each GPU runs them itself. Two operators mark the only points where the GPUs exchange data: g after each block and f before it.

Press Play to follow the forward pass down the layer. Then switch to the backward pass to see the all-reduces move from the output of each block to its input.

The next three sections take the picture apart: the MLP, the attention block, and the embeddings with the loss.

Splitting the MLP: columns, then rows

The MLP block of a transformer layer is two matrix multiplies (GEMMs) with a GeLU between them:

Y = \text{GeLU}(XA), \qquad Z = \text{Dropout}(YB)

There are two ways to put A on two GPUs. The first splits it by rows, A = \begin{bmatrix} A1 \ A2 \end{bmatrix}, with the input split to match, X = [X1, X2]. Each GPU then computes one term of X1 A1 + X2 A2. GeLU is not linear, so

\text{GeLU}(X1 A1 + X2 A2) โ‰  \text{GeLU}(X1 A1) + \text{GeLU}(X2 A2)

and the GPUs have to add their halves before the activation. That is a synchronization point in the middle of the block.

The second way splits A by columns, A = [A1, A2]. Every GPU sees the whole input and computes a different set of hidden units:

[Y1, Y2] = [\text{GeLU}(XA1), \text{GeLU}(XA2)]

GeLU acts on each hidden unit separately, so nothing has to be exchanged. Megatron-LM uses this split for the first GEMM and then splits the second by rows, B = \begin{bmatrix} B1 \ B2 \end{bmatrix}, so each GPU multiplies the hidden units it already holds by its own rows of B. Each GPU ends up with a partial sum of the output, and a single all-reduce adds the partial sums before dropout.

The backward pass mirrors this. The paper wraps the block in two operators that are conjugates of each other. g sits at the output: an all-reduce in the forward pass and the identity in the backward pass. f sits at the input: the identity in the forward pass, because every GPU already holds the same X, and an all-reduce in the backward pass, because each GPU computes only its own share of the gradient with respect to X. In PyTorch, f is a custom autograd function of a few lines (the paper's Code 1):

code
class f(torch.autograd.Function): def forward(ctx, x): return x def backward(ctx, gradient): all_reduce(gradient) return gradient

The figure runs one token through a toy block on two GPUs, with a hidden size of 2 and made-up weights. Every hidden unit is a vertical lane: its column of A above and its row of B below. Splitting A by columns and B by rows gives each GPU whole lanes, so nothing crosses between the GPUs until the end of the block. Step through both passes, then switch to the row split to see the extra synchronization it forces.

The distributed parallelism page draws the column and row splits tensor by tensor.

Attention: whole heads on each GPU

Multi-head attention is parallel by construction: each head has its own part of the query, key and value projections. Megatron-LM splits those three matrices by columns so that every multiply belonging to one head runs on one GPU, which divides the heads among the GPUs. Nothing is exchanged to complete the attention itself. The output projection that follows is split by rows, like the second MLP matrix, and ends in the same all-reduce.

The figure colors each head by the GPU that owns it, for the four configurations of the paper's scaling study.

Every pair of GEMMs is therefore fused into a column-split multiply followed by a row-split one, with no synchronization in between. A whole transformer layer costs two all-reduces in the forward pass and two in the backward pass (the paper's Figure 4). From the shapes involved, each one carries a block's output for the batch: batch size ร— sequence length ร— hidden size values.

Embeddings and the loss

The other large matrix is the embedding table: hidden size times vocabulary size, with 50,257 tokens in GPT-2's vocabulary. The output layer shares its weights with the input embedding, so the two have to be split the same way. Megatron-LM splits the table along the vocabulary dimension, E = [E1, E2]. On the input side each GPU holds only part of the table, so an all-reduce (g again) follows the lookup.

On the output side each GPU computes the logits for its own part of the vocabulary, [Y1, Y2] = [XE1, XE2]. Gathering them in one place for the loss would move batch size ร— sequence length ร— vocabulary size values. The paper instead fuses the parallel GEMM with the cross-entropy loss, which cuts what has to be communicated to batch size ร— sequence length values.

The paper gives that result without the derivation. The reason it can work is that cross-entropy needs only two numbers per position, the logit of the target token and the log of the sum of exponentials over the vocabulary, and a sum over the vocabulary is a sum of per-GPU sums.

The vocabulary is also padded. The logit GEMMs are efficient when the per-GPU vocabulary is a multiple of 128, so for up to 8-way model parallelism the paper pads 50,257 tokens to 51,200, the next multiple of 128 ร— 8 = 1,024.

How much the fused loss saves follows from the shapes. The figure draws what would cross between GPUs in one training step as areas to scale, for a batch of 8 sequences of 1,024 tokens per copy of the model and four all-reduces per layer. At that scale the fused loss is a dot about one pixel wide. The counts are this page's arithmetic, not numbers from the paper.

For the 8.3B model on 8 GPUs, a step makes 288 layer all-reduces of 25.2 million values each. Gathering the logits would move 419.4 million values, which is 17 times one layer all-reduce and the largest single transfer of the step. With the loss fused, 8,192 values cross instead.

What stays local

Much of the design is about not communicating. Instead of having one GPU compute layer normalization, dropout or the residual connection and broadcast the result, every GPU keeps its own copy of the layer normalization parameters and runs dropout and the residual add itself on the output of each model-parallel region. Each GPU also optimizes only the parameters it holds. Every value is either local to one GPU or duplicated on all of them, so no updated weights are ever exchanged inside a group.

Duplicated computation only works if it gives the same answer everywhere, which makes dropout the delicate part (Appendix B.2). Dropout on the residual path sits outside the model-parallel regions and has to produce the same mask on every GPU, so those random number generators are seeded identically at the start of training. Dropout inside a model-parallel region, in the attention block, has to differ between GPUs to be random across the whole operation, so each GPU keeps a second generator with its own seed.

Adding data parallelism

Model parallelism decides how one copy of the model is spread over GPUs. Data parallelism still decides how many copies train at once, and the paper uses both (Figure 8):

  • The GPUs that share one copy of the model form a model-parallel group, placed inside one server. They all-reduce activations among themselves.
  • The GPUs that hold the same slice of the parameters, one from each model-parallel group, form a data-parallel group. They all-reduce their weight gradients during back propagation, and the gradient all-reduces of the different data-parallel groups run in parallel.

The total is the product of the two. For the 8.3-billion-parameter model that is 8 GPUs per model-parallel group and 64-way data parallelism, 512 GPUs in all. All communication is PyTorch calling NCCL.

In the figure each column is one copy of the model, and each row is the set of GPUs that hold the same slice of the weights. Pick a GPU to see both of its groups.

Scaling results

The experiments run on up to 32 DGX-2H servers: 512 V100 GPUs with 32 GB each. GPUs inside a server are connected through NVSwitch at 300 GB/s, and servers are connected by InfiniBand at 100 GB/s (see multi-GPU communication for the difference). The baseline is the 1.2-billion-parameter model on one GPU. It sustains 39 TeraFLOPs over the whole training process, which the paper puts at 30% of the theoretical peak for one GPU in a DGX-2H.

For weak scaling the model grows with the number of GPUs, at roughly a billion parameters per GPU, while the hidden size per attention head stays at 96 (Table 1):

Hidden sizeHeadsLayersParametersModel-parallel GPUsWith 64-way data parallelism
153616401.2B164
192020542.5B2128
230424644.2B4256
307232728.3B8512

Runs with model parallelism alone use a batch of 8. Runs that add data parallelism use a global batch of 512, which is 64 copies of the model with a batch of 8 each. Measured against the single-GPU baseline (Figure 5), efficiency is 95%, 82% and 77% on 2, 4 and 8 GPUs with model parallelism alone, and 96%, 83%, 79% and 74% on 64, 128, 256 and 512 GPUs once data parallelism is added. The paper attributes the further drop to the gradient communication that data parallelism adds.

Two measurements in Appendix D show where the method stops paying:

  • Strong scaling (Table 8). Spreading the same 1.2B model over more GPUs at a fixed batch of 8 gives speedups of 1.64ร—, 2.34ร— and 2.98ร— on 2, 4 and 8 GPUs. The paper attributes the diminishing returns to per-GPU computation shrinking while memory bandwidth and communication overheads begin to dominate. The method is a way to fit a larger model, not a way to make a small one fast.
  • More heads scale worse (Table 7). For the 8.3B model on 8 GPUs, efficiency is 82% with 16 heads, 80% with 24 and 77% with 32. More heads mean smaller GEMMs inside the attention block and more elements in the attention softmax.

The headline number needs care. The abstract and introduction report 15.1 PetaFLOPs sustained on 512 GPUs and call it 76% scaling efficiency; 15.1 PetaFLOPs against 512 ร— 39 TeraFLOPs is 75.6%. Section 5.1 and Figure 5 give 74% for the 8.3B model on 512 GPUs. The paper does not reconcile the two figures.

Model results

The training corpus is 174 GB of deduplicated text drawn from Wikipedia, CC-Stories, RealNews and OpenWebText, with BooksCorpus added for the BERT models.

GPT-2 (Tables 2 and 3). Three left-to-right models train for 300k iterations at a batch of 512 on sequences of 1,024 tokens. Both evaluations are zero-shot.

ModelLayersHiddenHeadsGPUsDays per epochWikiText103 perplexityLAMBADA accuracy
355M24102416640.8619.3145.18%
2.5B541920201282.2712.7661.73%
8.3B723072245122.1010.8166.51%
Previous best15.7963.24%

The figure sets the three models against the previous best results. The 2.5B model already beats the previous best perplexity, and only the 8.3B model beats the previous best LAMBADA accuracy.

An epoch is 68,507 iterations. The 8.3B model reaches a validation perplexity of 9.27, and the larger models also converge faster per iteration (Figure 6). The previous bests are from Khandelwal et al. (2019) for WikiText103 and from the original GPT-2 for LAMBADA. This 8.3B model has 24 attention heads, not the 32 of the scaling study's 8.3B configuration.

The paper also checks for leakage by counting test-set 8-grams that occur in its training data: at most 10.8% for WikiText103, whose own training set already overlaps the test set by 9.09%, and at most 1.4% for LAMBADA.

BERT (Tables 4 and 5). Earlier work on ALBERT had found that BERT got worse when scaled past the 336M parameters of BERT-large. The paper traces the problem to where layer normalization sits. In the original BERT layout (Figure 7a) the skip connection branches off after the normalization, so the normalization sits on the main path between sublayers. The rearranged layout (Figure 7b) moves it inside each residual branch, ahead of the attention and the MLP, and leaves the skip path untouched; this is the arrangement now usually called pre-norm. A 752M model whose loss blew up partway through training with the original layout trained stably, and to a lower loss, with the rearranged one.

The figure redraws the two layouts of the paper's Figure 7.

With that change the paper trains three BERT-style models: 336M (24 layers, hidden size 1024), 1.3B (24 layers, hidden size 2048) and 3.9B (48 layers, hidden size 2560), on 128, 256 and 512 GPUs. Results improve with size on every task (Table 5; development-set medians over five seeds, except RACE, which is the test set):

ModelMNLI m/mmQQPSQuAD 1.1 F1/EMSQuAD 2.0 F1/EMRACE (m/h)
RoBERTa90.2 / 90.292.294.6 / 88.989.4 / 86.583.2 (86.5 / 81.8)
ALBERT90.892.294.8 / 89.390.2 / 87.486.5 (89.0 / 85.5)
XLNet90.8 / 90.892.395.1 / 89.790.6 / 87.985.4 (88.6 / 84.0)
Megatron-336M89.7 / 90.092.394.2 / 88.088.1 / 84.883.0 (86.9 / 81.5)
Megatron-1.3B90.9 / 91.092.694.9 / 89.190.2 / 87.187.3 (90.4 / 86.1)
Megatron-3.9B91.4 / 91.492.795.5 / 90.091.2 / 88.589.5 (91.8 / 88.6)

A 5-way ensemble of the 3.9B model reaches 90.9% on RACE (93.1 / 90.0), against 89.4% for the ALBERT ensemble. That is the 90.9% in the abstract; the single model scores 89.5%.

Critical analysis

Strengths:

  • A small change with a large effect. Two operators and a few all-reduce calls inside an existing PyTorch transformer, with the code released. The paper notes that Microsoft's 17-billion-parameter Turing-NLG was trained with it.
  • Communication is designed out rather than tuned. Fusing each pair of GEMMs, duplicating layer normalization, dropout and the residual add, and fusing the loss leave four all-reduces per layer and little else.
  • A strong baseline. Scaling is measured against a single GPU that already runs at 30% of its theoretical peak, not against a slow reference.
  • The models show why scale matters. Perplexity and accuracy improve with every increase in size, for GPT-2 and for BERT, and the layer normalization finding is a result in its own right.

Limitations:

  • It stays inside one server. A model-parallel group runs four all-reduces per layer, so it needs the bandwidth of NVSwitch. The paper says a model beyond 16 billion parameters would not fit in the 16 GPUs of a DGX-2H, and that such models would need inter-layer parallelism and model parallelism across nodes as well.
  • Efficiency falls with every doubling. It is 95%, 82% and 77% with model parallelism alone, and the strong-scaling speedup on 8 GPUs is 2.98ร—.
  • No measured comparison. GPipe and Mesh-TensorFlow are discussed but not benchmarked against.
  • Loose ends in the numbers. The abstract's 76% and Figure 5's 74% are not reconciled. The 3.9B BERT had run 1.5 million iterations and was still training when the paper was written, where the two smaller models ran 2 million. The 8.3B GPT-2 of the results is not the 8.3B configuration of the scaling study.
  • Not a like-for-like comparison with other BERT variants. Table 5 lists a trained-tokens ratio of 1 for the Megatron models against 2 to 3 for RoBERTa, ALBERT and XLNet, so the baselines saw two to three times as many tokens in pretraining. That gap runs in the baselines' favor, yet the comparison still mixes model size with training budget and recipe.

What came after

The future work the paper names, combining this method with inter-layer parallelism across nodes, is the subject of the 2021 follow-up by Narayanan et al. It composes tensor, pipeline and data parallelism and reports training iterations on a 1-trillion-parameter model at 502 PetaFLOPs on 3,072 GPUs.

The two all-reduces per layer did not go away either. They sit on the critical path of every token when a large model is served across several GPUs: NanoFlow overlaps them with computation, and Every ยตs Matters makes the all-reduce itself shorter.

  • Every ยตs Matters: shortens the small all-reduces that tensor parallelism puts after every attention and MLP block
  • NanoFlow: overlaps those collectives with computation when serving a tensor-parallel model
  • Collective Communication for 100k+ GPUs: the same tensor-parallel collectives at the scale of Llama 4, overlapped with the matrix multiply and moved off the GPU's threads
  • HiveD: why a shared cluster has to keep whole nodes available for jobs that, like a model-parallel group, need their GPUs in one server
  • BERT: the model whose layer normalization placement this paper changes to scale it to 3.9B parameters
  • Attention Is All You Need: the transformer layer whose attention heads and MLP are being split
  • Switch Transformers: a different route to very large models, through sparsely activated experts

If you found this paper review helpful, consider sharing it with others.

Mastodon