ZeRO: Memory Optimizations Toward Training Trillion Parameter Models
Data parallelism keeps a full copy of the training state on every GPU. ZeRO gives each GPU one slice and fetches the rest when a step needs it.
Rajbhandari, Rasley, Ruwase and He (Microsoft, 2019) partition the optimizer states, gradients and weights of data-parallel training across GPUs. Partitioning the optimizer states and gradients cuts per-GPU memory up to eightfold for the same communication as ordinary data parallelism; partitioning the weights too makes memory fall in proportion to the GPU count, for 1.5 times the communication.
Explaining the paperZeRO: Memory Optimizations Toward Training Trillion Parameter ModelsGPT-2's 1.5 billion weights take 3 GB in 16-bit floats, and the model still does not train on a 32 GB GPU in PyTorch or TensorFlow. The paper opens by asking where the rest of the memory goes.
A GPU has a fixed amount of memory, 32 GB on the NVIDIA V100s this paper used, and a training step runs only if everything it keeps alive fits. The weights are a small part of that. Training also keeps a gradient for every weight and, with the Adam optimizer, two running statistics per weight plus a full-precision copy of the weights; those last three are kept in 32-bit floats. The paper calls these per-parameter tensors the model states. For GPT-2's 1.5B parameters they add up to 24 GB, eight times the 3 GB of weights, before any activations are stored.
The usual way to train on many GPUs, data parallelism, does not help with this. Each GPU holds a complete replica of the model states and processes a different slice of the batch, so 64 GPUs hold 64 identical copies of the 24 GB. Adding GPUs adds throughput; the largest model you can train stays the one that fits on a single card.
ZeRO, short for Zero Redundancy Optimizer and released as part of Microsoft's DeepSpeed library, removes those copies. With data-parallel GPUs, each GPU stores of the model states and receives the rest from the GPUs that own it at the moment a step needs it. Every GPU still runs the full model on its own micro-batch, so the training math is that of data parallelism. A trillion-parameter model has about 16 TB of model states; split across 1024 GPUs, that is 15.6 GB each.
The page follows the paper's argument in order: count the bytes a training step keeps per parameter; see that data parallelism stores all of them on every GPU; partition them in three stages, optimizer states first; show that the first two stages move exactly as many bytes over the network as data parallelism already does and the third moves 1.5 times as many; and then handle what is left (activations, temporary buffers, fragmentation), which the paper calls ZeRO-R.
Where the memory goes in mixed-precision training
The paper splits training memory into two groups. Model states have one entry per parameter: the parameters themselves, their gradients, and the optimizer's per-parameter statistics. Residual states are everything else: activations saved in the forward pass for use in the backward pass, temporary buffers, and memory that is free but unusable because it is fragmented. For large models the model states dominate, so ZeRO starts there.
Large models in 2019 were trained in mixed precision.1 The forward and backward passes run in 16-bit floating point (fp16), which halves the memory of weights and activations and runs on the V100's tensor cores. The optimizer update cannot run in fp16: the update to a weight is often smaller than the gap between neighboring fp16 numbers near that weight, so adding it rounds back to the old value and training stalls. Mixed-precision training therefore keeps a second, fp32 master copy of every weight. The optimizer updates the master copy, and the fp16 weights used in the next forward pass are rounded from it.
Adam keeps two statistics for every weight: the momentum, a decaying average of past gradients, and the second moment, a decaying average of past squared gradients, which scales each weight's step by the size of its recent gradients. The paper calls the second one "variance". Both are kept in fp32. For a model with parameters the bytes add up as:
The paper names the optimizer-state term and calls the memory multiplier; for mixed-precision Adam, is 12. The fp32 master copy is one of the three blocks inside , next to momentum and variance, so the total is 16 bytes per parameter and not 20. Adam itself keeps two states; the third 4-byte block comes from mixed precision. The gradients here are fp16, 2 bytes each.
For GPT-2 at 1.5B parameters, (1) gives = 24 GB. The fp16 weights are 3 GB of that; the fp16 gradients another 3 GB; the fp32 optimizer states 18 GB. On a 32 GB card that leaves 8 GB, and the paper's own residual estimates for this model overrun it: about 8 GB of activations at batch 32 and sequence length 1024 even with activation checkpointing, and 6 GB for an fp32 buffer that fuses all gradients for one collective operation, which brings the total to 38 GB.
depends on the optimizer. SGD with momentum keeps the master copy and one velocity term, so is 8 and the total is 12 bytes per parameter. Plain SGD keeps only the master copy: is 4, 8 bytes in total. The paper states only Adam's value; the other two follow from its definition of . Switch optimizers in Figure 1 and move the model size to see where each one stops fitting on one 32 GB GPU: 2B parameters for Adam, about 2.7B for SGD with momentum, 4B for plain SGD, counting model states alone.
The 12 bytes of optimizer state are three quarters of the total, and they are used once per step, during the update. The forward and backward passes touch only the fp16 weights and gradients. ZeRO partitions the optimizer states first because they are the largest block and the one a GPU needs least often.
Data parallelism stores every byte on every GPU
In data parallelism each of the GPUs holds all bytes of model states. Each runs a full forward and backward pass on its own micro-batch, the GPUs average their gradients over the network, and every GPU applies the same optimizer step to its own copy, so the copies stay identical. The computation splits cleanly and the GPUs talk once per step, which is why data parallelism scales well. Its memory does not split at all. The paper's running example is a 7.5B-parameter model: 120 GB of model states on every GPU, whether there are 4 GPUs or 64.
In Figure 2 the five colors are the five blocks of (1). Under DP every GPU carries all five, and the redundancy counter reads . Switch to ZeRO and the same 120 GB is cut into disjoint slices, one per GPU, so the cluster stores each byte once and each GPU holds GB.
The other established way to fit a large model is model parallelism, as in Megatron-LM.2 It splits each layer's weight matrices across GPUs, so each GPU computes part of every layer and stores part of every weight. The GPUs then exchange partial results inside every layer: two all-reduces per transformer layer in the forward pass and two in the backward pass. That works inside one DGX-2 server, where GPUs talk over NVSwitch at about 300 GB/s per link, and degrades across servers on InfiniBand EDR at about 12.5 GB/s per link. The paper measured a 40B model run with Megatron-LM across two DGX-2 nodes at about 5 TFlops per V100, under 5% of the hardware's peak.
ZeRO keeps the computation of data parallelism. Every GPU runs the entire forward and backward pass on its own micro-batch; no matrix multiply is split. Only the stored states are partitioned, and a missing piece is gathered just before it is used and released right after. ZeRO also composes with model parallelism: run Megatron-style splitting inside a server and ZeRO across servers, and the per-GPU model-state memory can fall by up to , where is the model-parallel degree.
Two simpler ways to cut per-GPU memory exist, and the paper sets both aside on efficiency. Moving the model states to CPU memory sends them over PCIe every step; earlier systems that did this spent up to 50% of training time on GPU-to-CPU-to-GPU transfers. Raising the model-parallel degree until the model fits pushes Megatron's per-layer all-reduces across nodes, where the 40B measurement above ran at under 5% of peak.
How ZeRO partitions the model states
ZeRO-DP, the data-parallel half of the paper, removes the redundancy in three cumulative stages, in order of the bytes each one reclaims: optimizer states (12 bytes per parameter), then gradients (2), then fp16 parameters (2). In each stage the parameters are divided into equal shards and GPU owns shard .
Stage one, : optimizer states. GPU keeps the fp32 master weights, momentum and variance for its shard only. Every GPU still computes and stores a full fp16 gradient, but GPU needs only the cross-GPU sums for shard , and runs Adam on that shard alone. An all-gather at the end of the step then gives every GPU the full set of updated fp16 weights for the next forward pass. Per-GPU memory becomes:
For the 7.5B model on 64 GPUs: the fp16 weights and gradients stay whole, = 30 GB, and the optimizer states shrink to = 1.4 GB, for 31.4 GB per GPU instead of 120.
Stage two, : also the gradients. GPU updates only shard , so it needs only the gradients of shard , summed over all GPUs. As each layer's gradients are produced during the backward pass, they are summed onto the GPU that owns them, and the other GPUs free their copies. The paper groups gradients into buckets per shard and reduces a whole bucket at once, which amounts to a reduce-scatter (defined in the next section). Per-GPU memory:
Only the fp16 weights stay whole: 15 GB for the 7.5B model, plus = 1.6 GB, for 16.6 GB per GPU.
Stage three, : also the parameters. GPU stores only shard of the fp16 weights too. When the forward or backward pass reaches a layer, the GPUs gather that layer's full weights from their owners, compute, and discard the parts they do not own. Nothing remains replicated:
For the 7.5B model on 64 GPUs that is 120 / 64 = 1.875 GB, which the paper prints as 1.88. Equations (2) and (3) have a term that does not shrink with ; (4) has none, so doubling the GPU count halves per-GPU model-state memory at every scale. Because (4) has no floor, the paper can place a trillion parameters on 1024 GPUs, as the section on model size works out.
In Figure 3 each block of (1) keeps a fixed slot and fills only as much of it as one GPU stores. Step through the four stages at = 64 and the total reads 120, 31.4, 16.6 and 1.88 GB, the paper's Table 1. Then drag to 1024: only gets to 30.1 GB, because the 4 bytes of fp16 weights and gradients stay on every GPU.
The paper advertises the first two stages as 4× and 8× reductions. Those are the limits as grows without bound, when only the resident or is left. At 64 GPUs the reductions are = 3.82× and = 7.2×. Stage three has no resident term, so its 64× at 64 GPUs is exact.
Why stages one and two add no communication
Partitioning means GPUs must send each other the pieces they do not own. To judge the cost, the paper first counts what data parallelism already sends each step, then counts the same quantity for each stage. It counts elements per GPU and treats the transfers as bandwidth-bound, which holds for gradients this large.
Data parallelism averages gradients with an all-reduce: every GPU starts with its own gradient and ends with the sum over all GPUs (dividing by to get the average is a local step). The bandwidth-optimal way to do this, used by NVIDIA's NCCL library and by Horovod,3 runs in two phases over a ring of GPUs, with each gradient cut into chunks:
- Reduce-scatter. In each of steps, every GPU sends one chunk to the next GPU in the ring, which adds it to its own copy of that chunk. After the last step, GPU holds the full sum of chunk and nothing else that is complete.
- All-gather. In another steps, every GPU passes a completed chunk to the next GPU, which copies it without arithmetic. At the end every GPU holds every summed chunk.
Each GPU sends one chunk of elements per step, so each phase moves elements per GPU. With 4 GPUs that is 0.75Ψ per phase and 1.5Ψ for the all-reduce. Press Play in Figure 4 and watch the counts in the cells: amber cells hold partial sums, labeled with how many GPUs have contributed, and Σ marks a complete sum.
The exact volume is , a little under ; the paper rounds it to the large- value:
The 2 counts the two phases. It is not the 2 bytes of an fp16 number; here is a count of elements. For the 7.5B model with fp16 gradients, (5) is 15 billion elements, 30 GB sent per GPU per step.
In stage two, GPU needs only the summed gradients of shard , and the reduce-scatter alone delivers exactly that: elements. GPU updates its fp32 master shard and rounds it to fp16. Every GPU still needs all the fp16 weights for the next forward pass, so an all-gather circulates the updated weight shards: another . The total is , the same as (5), so stage two adds no communication over data parallelism. Stage one can run the same schedule, because GPU also reads only the summed gradients of its own shard; it differs from stage two only in keeping the whole gradient buffer allocated. Toggle Figure 4 to ZeRO and the arrows are identical step for step; what changes is the payload of the second phase (updated weights in place of summed gradients) and what each GPU keeps (one gradient chunk in place of all of them).
# One training step under ZeRO stage 2 (P_os+g) on GPU i of Nd.
# GPU i keeps: all fp16 params (2*Psi bytes), and for shard i only:
# fp16 grads, fp32 master params, Adam m and v.
loss = model(local_batch) # full forward, same as plain DP
loss.backward() # full backward, same as plain DP
reduce_scatter(grads) # Psi elements moved; GPU i keeps the
# SUM of grad shard i, frees the rest
master[i] = adam(master[i], grad[i]) # update 1/Nd of the parameters
params[i] = fp16(master[i]) # refresh this GPU's fp16 shard
all_gather(params) # Psi elements moved; everyone gets
# the full updated fp16 params
# total per step: Psi + Psi = 2*Psi, the same as plain DP's all-reduceStage three: parameters gathered per layer
Once the fp16 weights are partitioned too, each GPU must receive the weights it does not own whenever it computes a layer, and every layer is computed twice per step: once in the forward pass and once in the backward pass. The paper spreads the weight all-gather over the forward pass, one layer at a time, discarding each layer's gathered weights after use; the backward pass repeats the gathers in reverse order. Summed over all layers each pass gathers the whole model once, elements, and the gradient reduce-scatter adds a third:
The all-gather after the update that stages one and two needed is gone, because no GPU needs a full copy of the weights between steps; the forward gathers replace it. The extra is the second gather in the backward pass. The paper describes each transfer as the owner broadcasting its shard; across all shards that adds up to an all-gather.
If every GPU gathers the full weights of every layer, the memory saving comes from timing: only one layer's full weights exist on a GPU at any moment, next to the GPU's own shard of everything. Press Play in Figure 5 and follow the meter as the sweep goes down eight layers and back up.
# Stage 3 (P_os+g+p): GPU i stores only shard i of each layer's params.
for layer in layers: # forward pass
w = all_gather(layer.shards) # full weights of THIS layer only
x = layer(x, w)
free(w) # keep shard i, drop the rest
for layer in reversed(layers): # backward pass
w = all_gather(layer.shards) # gathered a second time
g = layer.backward(w)
reduce_scatter(g) # GPU i keeps the summed grad shard i
free(w)
# params moved: Psi forward + Psi backward; grads: Psi. Total 3*Psi.Figure 6 puts the three schedules side by side against the data-parallel baseline. Communication stays at or whatever is, while the memory reduction in the right column grows with it.
None of the three stages changes what is computed. Each GPU sums the same gradients and applies the same Adam update to each parameter as in data parallelism; ZeRO changes which GPU stores a value and when it is sent. The paper states that its optimizations do not change the optimization method or affect convergence, and contrasts this with PipeDream, a pipeline scheme that trains on stale weight copies, and with memory-saving optimizers that keep coarser per-parameter statistics. It does not claim bit-identical results, since summing in a different order can change floating-point rounding.
How large a model fits
Dividing a GPU's memory by the per-GPU bytes per parameter gives the largest model whose model states fit. With 32 GB per GPU and = 64, standard data parallelism fits 32 / 16 = 2B parameters; fits 32 / 4.19 = 7.6B; fits 32 / 2.22 = 14.4B; and fits = 128B. Model parallelism of degree divides every per-GPU term by , so each ceiling multiplies by . These are the paper's Table 2 numbers. They ignore activations and buffers, which is why the paper calls them theoretical maxima.
Figure 7 draws the four ceilings on a log axis next to models from the paper. At the default, standard data parallelism stops near GPT-2's 1.5B. Set to 16, one DGX-2 server's worth of model-parallel GPUs, and keep at 64: that is 1024 GPUs, and stage three reaches 2T. Then return to 1 and drag ; only the stage-three bar keeps moving, because the other stages have a floor in bytes per parameter.
The paper checked the stage-one ceiling against real runs. With alone (a configuration the paper calls ZeRO-OS), the largest models that ran were 6.2B on 64 GPUs and 100B on 1024 GPUs with 16-way model parallelism, against theoretical maxima of 7.6B and 121.6B. Without ZeRO the measured maxima were 1.3B and 20B, against 2B and 32B, which the paper takes as evidence that its analysis gives realistic upper bounds.
The trillion in the title is this analysis carried to 1024 GPUs. A trillion parameters at 16 bytes each is 16 TB of model states; with stage three on 1024 GPUs that is 15.6 GB per GPU (the introduction rounds it to 16 GB), which fits on a 32 GB V100. The paper separates fitting the model from training it in reasonable time. A 1T model does about 3000 times the computation per sample of BERT-Large (1 trillion / 330 million), and BERT-Large trained in 67 minutes on a 1024-GPU DGX-2H cluster. At the same efficiency, sample count and sequence length, 3000 × 67 minutes is about 140 days; with the larger datasets and longer sequences a bigger model would use, the paper estimates over a year, and says reasonable training times need an exaflop-scale system.
ZeRO-R: activations, buffers and fragmentation
With the model states partitioned, the residual states become the next limit. The paper's second set of optimizations, ZeRO-R, handles the three kinds in turn.
Activations. The paper estimates a GPT-2-like transformer's stored activations as about values, for hidden size , batch size , sequence length and layers. For the 1.5B GPT-2 configuration in the paper (48 layers, hidden size 1600) at batch 32 and sequence length 1024:
which is the paper's "about 60 GB". The standard remedy is activation checkpointing (Chen et al., 2016), which is prior work and not part of ZeRO: keep only some activations, the checkpoints, and recompute the rest during the backward pass. It cuts activation memory to roughly the square root of the total for about 33% extra computation, taking this model to about 8 GB. For a 100B model at batch 32 the paper still counts around 60 GB even with checkpointing.
Model parallelism adds a redundancy of its own. In Megatron-style splitting every GPU that shares a layer needs that layer's full input, so each of them stores the full checkpoint. ZeRO-R's partitioned activation checkpointing, , applies the stage-one idea to these copies: after a layer's forward pass its input checkpoint is split across the model-parallel GPUs, and an all-gather rebuilds the full copy just before the backward pass recomputes that layer. The paper's example is a 100B model at batch 32, sequence length 1024 and 16-way model parallelism, checkpointing one activation per transformer layer: about 33 GB of checkpoints per GPU, and about 2 GB with .4 A variant, , moves the partitioned checkpoints to CPU memory, which takes their GPU footprint to nearly zero.
The cost of is one extra all-gather per transformer block, of size sequence length × hidden size. Megatron-LM with checkpointing already moves 12 × sequence length × hidden size per block (two all-reduces in the forward pass, two in the recompute, two in the backward pass, each costing twice its message size), so the addition is 1/12 of that, about 8%, which the paper states as under 10%. Because frees activation memory, it allows a batch up to times larger, and data-parallel communication per sample falls in proportion.
Offloading to CPU crosses PCIe, which is far slower than the GPU's own memory, and it doubles the CPU data movement relative to . The paper's argument for doing it anyway is arithmetic intensity, the ratio of computation per iteration to activation-checkpoint bytes per iteration. A transformer layer's computation grows with the square of the hidden size and its checkpoint only linearly, so the ratio grows linearly with hidden size and is at least 10,000 for GPT-2-sized models and larger. With that much computation per byte, the transfers can overlap with computation instead of stalling it; the bytes still move. The paper turns on only when it helps, for models so large that the batch would otherwise be tiny.
Temporary buffers. Libraries such as NVIDIA Apex and Megatron fuse all gradients into one flat buffer before an all-reduce or a gradient-norm computation, because one large collective runs at much higher bandwidth than many small ones. That buffer grows with the model: in fp32, 6 GB for a 1.5B model and 12 GB for a 3B one. ZeRO-R's constant-size buffers () cap the fused buffer at a fixed size that is still large enough to run efficiently.
Fragmentation. Checkpointing interleaves short-lived tensors (activations that will be recomputed) with long-lived ones (the checkpoints), and the backward pass interleaves long-lived parameter gradients with short-lived activation gradients. The free memory ends up in pieces too small for the next large request, and the paper saw out-of-memory failures with more than 30% of memory still free. Memory defragmentation () pre-allocates contiguous regions for checkpoints and parameter gradients and copies each one there as it is produced.
What ZeRO-100B measured
The paper implemented a subset of ZeRO, called ZeRO-100B: stage two of ZeRO-DP () plus all of ZeRO-R, in PyTorch, wrapping any torch.nn.Module the way ordinary data parallelism does. The paper stopped there because it targeted models around 100B parameters, an order of magnitude past the largest published (T5's 11B), and still trainable in reasonable time on about a thousand V100s. Stage three and the trillion-parameter numbers are analysis, not measurement. All experiments ran on 400 V100 GPUs in 25 DGX-2 nodes with 800 Gbps between nodes, on GPT-2-like transformers of varying depth and width.
Model size. Combined with Megatron model parallelism inside each node, ZeRO-100B trained models up to 170B parameters. Megatron alone ran efficiently up to about 16 to 20B on one DGX-2 and fell off sharply once its model parallelism crossed nodes; the paper calls the gap over 8×. Its ablation, all at 16-way model parallelism, shows where the size comes from: 40B with , 60B after adding (16× less activation memory), 140B with (half the model-state memory of ), and 150B with CPU offload of the checkpoints.
Speed. On 100B models ZeRO-100B sustained over 38 TFlops per GPU, over 15 PetaFlops across the 400 GPUs, and averaged over 30% of hardware peak for models from 8B to 100B. That is up to 10× the Megatron baseline at the same model size. The paper notes the baseline used power-of-two GPU counts (256 for the 170B model) and compares per-GPU throughput, a setup it says gives the baseline an advantage, since fewer GPUs means better communication throughput for the baseline.
Super-linear scaling. For a 60B model, going from 64 to 400 GPUs more than doubled total throughput for each doubling of GPUs. Under , more GPUs means a smaller model-state share per GPU, the freed memory holds a larger per-GPU batch, and larger batches do more computation per byte communicated, so each GPU runs faster.
Without model parallelism. On 128 GPUs with data parallelism alone, ZeRO-100B trained models up to 13B parameters, more than T5's 11B, at over 40 TFlops per GPU. PyTorch's DistributedDataParallel ran out of memory above 1.4B and stayed under 20 TFlops per GPU. Users did not change their models. The paper also points out that without model parallelism's per-layer traffic, these runs do not need fast intra-node links such as NVLink or NVSwitch.
Turing-NLG. Microsoft trained Turing-NLG, a 17B-parameter language model, end to end with ZeRO-100B at 41.4 TFlops per GPU.5 It reached a WebText-103 perplexity of 10.21, a new state of the art at the time. Perplexity is the exponential of the average per-token loss, so lower means the model assigns higher probability to held-out text.
DeepSpeed added stage three after the paper, as its introduction planned, and later systems adopted the same partitioning. PyTorch's Fully Sharded Data Parallel cites ZeRO as its motivation; it shards each parameter across data-parallel GPUs, gathers a unit's full parameters before computing it and frees them after, and its authors describe a revised design for PyTorch rather than a port.6
Questions you might still have
Is ZeRO a form of model parallelism?
No. Model parallelism splits each layer’s matrix multiplies across GPUs, so every GPU computes part of every layer. Under ZeRO every GPU still runs the whole forward and backward pass on its own micro-batch, as in data parallelism; only the stored optimizer states, gradients and (in stage three) parameters are split, and pieces are fetched when a step needs them. The two compose: with data-parallel degree Nd and model-parallel degree Nm, per-GPU model-state memory can fall by up to Nd times Nm.
Is PyTorch FSDP the same thing as ZeRO stage three?
Close, but its authors do not call it the same system. The FSDP paper (Zhao et al., 2023) says it was motivated by DeepSpeed’s Zero Redundancy Optimizer and inspired by it, with a revised design built around PyTorch. Like stage three, FSDP keeps a shard of each parameter on each GPU, gathers a unit’s full parameters before computing it, and discards them afterwards, and every GPU still runs the full model on its own data.
If every GPU gathers the full parameters to compute, how is stage three different from replication?
The parameters are gathered one layer at a time and freed before the next layer is gathered, so a GPU holds its own 1/Nd shard of every layer plus the full weights of one layer. The figure on stage three shows the live-parameter meter staying at one layer of eight through the whole forward and backward pass.
Does ZeRO change the result of training?
It computes the same summed gradients and applies the same optimizer update as data parallelism; only where the bytes live and when they move changes. The paper states that its optimizations do not change the optimization method or affect convergence, which separates it from pipeline schemes that keep stale weight copies and from memory-saving optimizers that store coarser statistics. It does not claim bit-identical weights; a different reduction order can change floating-point rounding.
The paper says 4x and 8x, but 120 / 31.4 is 3.8. Which is right?
Both. The 4x and 8x are limits as Nd grows without bound, when only the resident 4 and 2 bytes per parameter remain. At Nd = 64 the reductions are 120 / 31.4 = 3.82x and 120 / 16.6 = 7.2x. Stage three has no resident term, so its 64x at Nd = 64 is exact.
Why did the paper evaluate only stage two?
The authors aimed the implementation at roughly 100B-parameter models, which current hardware could train in reasonable time, and stage two plus ZeRO-R reaches that. Stage three and the trillion-parameter numbers in the paper are memory and communication analysis; the introduction says DeepSpeed would add stage three later to support one trillion parameters.
Footnotes & further reading
- Micikevicius et al., Mixed Precision Training (ICLR 2018), which introduces the fp32 master copy and loss scaling. The two optimizer states are from Kingma & Ba, Adam: A Method for Stochastic Optimization (2015), explained here. ↩
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019), explained here. ↩
- The two-phase ring all-reduce was popularized for deep learning by Sergeev & Del Balso, Horovod (2018), and is implemented in NVIDIA's NCCL. ↩
- We could not reproduce the 33 GB from the paper's 100B configuration (125 layers, hidden size 8192, Table 4) as bytes: one fp16 checkpoint per layer at batch 32 and sequence length 1024 is values, about 67 GB at 2 bytes each, which is close to Section 3.2's "around 60 GB" for the same model. The 33 matches the count of values. The 16-fold reduction from is the same either way. Activation checkpointing itself is Chen, Xu, Zhang, Guestrin, Training Deep Nets with Sublinear Memory Cost (2016), explained here; the follow-on work on transformer activation memory is Reducing Activation Recomputation. ↩
- Turing-NLG: Microsoft Research blog (2020). ↩
- Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel (2023). The implementation of ZeRO ships in DeepSpeed. ↩
How could this explainer be improved? Found an error, or something unclear? I read every message.