This post is for readers who care about LLM system design and want an overall picture. Starting from the LLM itself, it follows one causal chain through the main designs of training and inference systems, with the focus on the motivation of each design. Each topic covers, in order, what problem arises, how it is solved, and what is improved and what is sacrificed. Only the principles and numbers needed to understand the designs are kept; implementation details are left out. All numbers come from public hardware specifications and papers, with the sources listed at the end.
LLM
- Scope
The training and inference systems of LLMs (large language models), and where their main designs come from.
To understand an LLM, first look at its basic architecture
- This post is about the training and inference systems of LLMs.
- The model architecture determines the system's costs: analyzing the costs requires knowing the architecture first.
- Almost all current LLMs use the same architecture: the Transformer.
- So the analysis starts with the architecture of the Transformer and the data on the GPU during training.
Basic Architecture: Transformer
- Architecture
Each layer has two main modules: attention (lets each token gather information from the tokens before it; split internally into several heads computed independently) and FFN (a two-layer fully connected network, computed for each token separately). $$L$$ layers are stacked, with modules such as tokenization (splitting text into tokens, which are words or parts of words) and embedding added before and after. Most parameters are in the matrices of these $$L$$ layers.
- Data
During training, the data on the GPU is of four kinds: parameters (the matrices above, with count $$\Phi$$), gradients and optimizer states (one gradient per parameter, plus two statistics per parameter for Adam). These three together are the training state, and the model alone sets their size. Activations (intermediate results of each module's forward pass, needed in the backward pass) vary in size with the batch (the number of sequences processed at once) and the sequence length. The amount of input data is small, so it can be counted with the activations.
This data has only three states, being stored, computed or moved, and each state has its own cost
- During training, the data on the GPU is of four kinds: parameters, gradients, optimizer states and activations.
- Any piece of data, at any moment, is in one of three states: being stored, being computed on, or being moved.
- The three states use three different hardware resources: storing uses memory capacity, computing uses compute, and moving uses bandwidth. Each resource has a limit, usually called the memory wall, the compute wall and the bandwidth wall.
- The units of the three walls are Bytes, FLOP/s and Bytes/s.
Three costs
memory wall · compute wall · bandwidth wallBytesH100: 80 GBFLOP/sH100 BF16: 989 TFLOP/sBytes/sH100: HBM 3.35 TB/s, NVLink 450 GB/s, cross-machine network 50 GB/sThe three costs are of different kinds. Memory cost is a capacity: past the memory capacity, the program cannot run. Compute cost and data-movement cost are times, equal to operations ÷ achieved compute and bytes ÷ achieved bandwidth; when compute units sit idle or bandwidth is not fully used, the time grows. HBM is the GPU memory; SRAM is small, fast on-chip storage. H100 figures are for the SXM version: compute is the dense value without sparsity; NVLink and the cross-machine network are counted in one direction, and the cross-machine network assumes one 400 Gb/s NIC per GPU.
From here on, what each technique improves and sacrifices is described in terms of these three costs.
Two questions, in order: ① Is the technique waste removal, that is, does it only remove parts that served no purpose? ② If not, which cost does it improve and which does it sacrifice?
waste removal Removes parts that served no purpose, such as duplicated data or idle time of compute units: improves one cost and sacrifices no other; small implementation overheads are not counted.
trade-off Improves one cost and sacrifices another. Usually the improved cost is the bottleneck in the current setting, and the sacrificed one has room to spare in that setting.
Beyond the three costs: each kernel launch (a kernel is a function run on the GPU) or communication call also has a fixed cost on the order of microseconds, independent of data size; a few techniques sacrifice numerical accuracy, load balance, latency or simplicity of implementation. Cards mark these in gray.
Colors mark the cost that changes: memory cost compute cost data-movement cost beyond the three costs
With the three costs, trade-offs can be analyzed, starting inside one GPU
$$T$$: training time; $$D$$: number of training tokens; $$N$$: number of GPUs; $$R$$: achieved compute per GPU (FLOP/s), counting only the model's own operations and not recomputation; $$\eta$$: parallel efficiency. $$M$$: GPU memory used per GPU (byte); $$16\Phi$$ is the training state (parameters, gradients and Adam's two statistics, $$4\Phi$$ bytes each when all are in FP32) and $$A$$ the activations.
- Training takes about $$6\Phi D$$ FLOP in total: for each parameter and each token, about 2 FLOP in the forward pass and about 4 FLOP in the backward pass.
- The model and data set the numerator, so training time can be cut only by increasing the three factors in the denominator: the achieved compute of one GPU $$R$$ (mixed precision, kernel optimization), the number of GPUs $$N$$ (DDP and the parallel methods after it), and the parallel efficiency $$\eta$$ (waiting for communication and idle GPUs keep it below 1). The constraint is that each GPU's memory use $$M$$ does not exceed the memory capacity.
- First, two defaults on one GPU are changed, and both changes are trade-offs: giving up a cost with room to spare for the one that is the bottleneck, to raise $$R$$ or lower $$M$$.
- The first default is numerical precision: by default each number is stored in FP32, taking 4 bytes, and matrix multiplies do not go through Tensor Cores (when TF32 is not enabled). On H100, the FP32 units reach 67 TFLOP/s and the BF16 Tensor Cores 989 TFLOP/s.
- The second is how activations are kept: by default all are saved. A 1.3B model (the GPT-3 XL architecture) processing 32 sequences of 2048 tokens per step has about 110 GB of activations in BF16 without the attention score matrices (about 500 GB with them), while the training state is only 21 GB.
- The techniques that change these two defaults are mixed precision and activation checkpointing.
Trade-offs within one GPU
mixed precision · activation checkpointing- Problem
With FP32 throughout, the matrix-multiply peak is only about 1/15 of the BF16 Tensor Core peak, and each activation value takes 4 bytes; with all activations saved, they may exceed the memory capacity.
- Solution
- Mixed precisiontrade-off
Matrix-multiply inputs and activations use BF16 (2 bytes per number), with products accumulated in FP32; reductions such as softmax and normalization, and the parameter update, use FP32. The update stays in FP32 because an update smaller than about 1/256 of the parameter is rounded away in BF16. The training state is BF16 parameters and gradients at $$2\Phi$$ each, plus FP32 parameters and Adam's two statistics at $$12\Phi$$, for $$16\Phi$$ bytes in total.
Improvescompute: about 15 times the matrix-multiply peakImprovesmemory: activations halvedSacrificesnumerical accuracyTraining state still $$16\Phi$$: FP32 parameters keptActivation checkpointingtrade-offSaves one layer's activations every $$\sqrt{L}$$ layers or so; when the backward pass needs the others, it reruns the forward pass from the nearest saved point.
Improvesmemory: activations from $$L$$ layers down to about $$2\sqrt{L}$$ layersSacrificescompute: one extra forward pass, $$R$$ drops to about 3/4
Mixed precision raises peak compute, but achieved compute is far below the peak, and the kernel implementation sets the gap
$$F$$: peak compute; $$u$$: utilization. Mixed precision raises $$F$$; the implementation of the kernels sets $$u$$.
- The peak of BF16 Tensor Cores is about 15 times that of the FP32 units, but this is only an upper bound; achieved compute depends on each kernel.
- The gap has two sources. The first is kernels limited by memory bandwidth: RMSNorm does about 1 FLOP per byte read or written, while H100 needs about 295 FLOP per byte (989 TFLOP/s ÷ 3.35 TB/s) to use its full compute; memory bandwidth sets its time, independent of peak compute.
- These kernels (norm, softmax, activation functions, element-wise operations) do few operations but take much of the time: when training BERT-large with PyTorch on V100, they account for 0.2% of the operations and 39% of the time.
- The second is matrix multiplies: the whole matrix does not fit on chip, and compute units wait for data to be read from GPU memory into the chip.
- How to tell: a kernel that does $$W$$ FLOP and reads and writes $$Q$$ bytes of GPU memory takes at least the larger of $$W/F$$ and $$Q/B$$ ($$B$$ is the memory bandwidth; the derivation is in the roofline post). With arithmetic intensity $$W/Q$$ below the ridge point $$F/B$$, it is limited by memory bandwidth; above it, by compute.
- The two kinds are optimized separately: for kernels limited by memory bandwidth, cut the bytes read and written to GPU memory (kernel fusion → online softmax → FlashAttention); for compute-limited matrix multiplies, cut the time compute units wait for data (tiling, pipelining).
Kernel efficiency
roofline · kernel fusion · FlashAttention · tiling and pipelining- Problem
For some kernels, memory bandwidth sets the time, independent of peak compute; matrix multiplies also wait for data to be read from GPU memory into the chip.
- Solution
- Kernel fusionwaste removal
Merges several consecutive kernels limited by memory bandwidth into one kernel; intermediate results stay on chip and are not written back to GPU memory.
Improvesdata movement: bytes read and written to GPU memoryImprovesfixed cost of kernel launchesFlashAttentiontrade-offAttention computes a query, a key and a value vector for each token, stacks them into matrices $$Q$$, $$K$$ and $$V$$, then computes in turn $$S = QK^\top$$ (match scores between every pair of tokens), $$P = \text{softmax}(S)$$ and the output $$O = PV$$. $$S$$ and $$P$$ are both $$n \times n$$ matrices ($$n$$ is the sequence length), about 32 MB per head at $$n = 4096$$. Softmax needs the maximum and sum of a whole row; the standard implementation uses three kernels and writes $$S$$ and $$P$$ back to GPU memory.
FlashAttention uses online softmax to update each row's maximum and sum block by block, merges the three steps into one kernel and computes them in blocks on chip, so $$S$$ and $$P$$ are not written back to GPU memory; the backward pass recomputes $$S$$ and $$P$$ from $$Q$$, $$K$$ and these two per-row statistics. In the paper's example (GPT-2 medium, sequence length 1024, A100), the forward plus backward time of attention drops from 41.7 ms to 7.3 ms.
Improvesdata movement: GPU memory reads and writes 40.3 → 4.4 GBImprovesmemory: removes the $$n^2$$ term in activationsSacrificescompute: backward recomputes $$S$$ and $$P$$, 66.6 → 75.2 GFLOPTiling and pipeliningwaste removalIn an $$n \times n$$ matrix multiply, each number takes part in $$n$$ multiply-adds; if each number (BF16) is read or written only once, the arithmetic intensity is $$n/3$$ (about 1365 at $$n = 4096$$). But the whole matrix does not fit on chip. With tiling, data read into the chip is reused many times; the next tile is read while the current one is computed.
Improvesdata movement: bytes read repeatedly from GPU memoryImprovescompute: time compute units wait for data
With peak and utilization both raised, one GPU's compute still has a limit, so the only option is more GPUs
6 × 405B × 15.6T tokens ≈ 3.8×10²⁵ FLOP; one H100 at 989 TFLOP/s would take about 1200 years
- Both factors of $$R = F \times u$$ have been raised: mixed precision raises $$F$$ and kernel optimization raises $$u$$ (checkpointing goes the other way, trading about 1/3 extra compute for GPU memory).
- The model and data set the numerator $$6\Phi D$$. Llama-3 405B: $$\Phi = 4.05 \times 10^{11}$$, $$D = 1.56 \times 10^{13}$$ tokens, $$6\Phi D \approx 3.8 \times 10^{25}$$ FLOP.
- One H100 running continuously at its 989 TFLOP/s peak would need about 1200 years.
- The numerator is fixed, $$R$$ cannot exceed the peak and $$\eta$$ cannot exceed 1, so only the number of GPUs $$N$$ can keep growing.
DDP
Data Parallelism- Problem
One GPU's compute has a limit; frontier-scale training would take over a thousand years on one GPU.
- Solution
Each of $$N$$ GPUs holds a full copy of the model and processes different data. Each step averages the gradients with all-reduce (each GPU contributes one piece of data, and at the end every GPU has their sum), and all copies apply the same update. Communication can overlap with the backward computation; when communication takes less time than computation, $$\eta$$ is close to 1.
DDP does not reduce GPU memory, and each GPU still has to hold the full training state. What if it does not fit?
7B × 16 bytes = 112 GB, more than the 80 GB of H100
- DDP assumes that each GPU can hold the full training state of $$16\Phi$$ bytes.
- $$16\Phi$$ grows linearly with the parameter count: for a 7B model it is 112 GB, more than the 80 GB of H100, and DDP cannot run.
- Yet the $$N$$ copies of $$16\Phi$$ on the $$N$$ GPUs are identical.
- So each GPU needs to store only $$1/N$$ and fetch the rest from other GPUs when needed, with no information lost.
- ZeRO defines the order of sharding and the communication that fetching needs.
ZeRO / FSDP
Sharding the training state- Problem
$$16\Phi$$ exceeds one GPU's memory, while the $$N$$ GPUs store $$N$$ identical copies.
- Solution
Shards in stages, from the least to the most frequently used data.
ZeRO-1waste removalShards the optimizer states ($$12\Phi$$). All-reduce can be split into two steps: a reduce-scatter (each GPU gets $$1/N$$ of the gradient sum) and an all-gather (each GPU sends its $$1/N$$ to all GPUs). ZeRO-1 has each GPU update only its own $$1/N$$ of the parameters between the two steps and then all-gathers the updated parameters, so each GPU needs to store only $$1/N$$ of the optimizer states, and the communication volume is unchanged.
Improvesmemory: $$16\Phi \to 4\Phi + 12\Phi/N$$Data movement unchanged: $$2\Phi$$ communicated per stepZeRO-2waste removalAlso shards the gradients ($$2\Phi$$): after the reduce-scatter, each GPU needs to keep only its own $$1/N$$ of the gradient sum and can free the rest.
Improvesmemory: $$\to 2\Phi + 14\Phi/N$$Data movement unchangedZeRO-3trade-offAlso shards the parameters ($$2\Phi$$): the full parameters of each layer are fetched before it is computed, once in the forward pass and once in the backward pass. The FULL_SHARD mode of PyTorch FSDP corresponds to ZeRO-3.
Improvesmemory: $$\to 16\Phi/N$$Sacrificesdata movement: communication per step $$2\Phi \to 3\Phi$$
Memory is counted in bytes and communication in elements, following the convention of the ZeRO paper.
ZeRO splits storage, but each GPU still computes the full model, so with many GPUs, syncing parameters takes longer than computing
$$G$$: total tokens per step (global batch); $$B_\text{net}$$: cross-machine network bandwidth. $$\Phi$$ cancels out.
- With ZeRO-3, each GPU communicates $$3\Phi$$ elements per step, which does not shrink with $$N$$; its computation is $$6\Phi$$ times the number of tokens it gets. The two can overlap; when the ratio in the formula above is below 1, communication can overlap completely with computation.
- The total tokens per step $$G$$ has a limit: past the critical batch size, a larger batch saves fewer and fewer training steps.
- With $$G$$ fixed, each GPU gets $$G/N$$ tokens, and the ratio of communication time to computation time is proportional to $$N$$.
- Llama-3 405B uses 16384 H100s with $$G$$ = 16M tokens, only about 1000 tokens per GPU on average. If all 16384 GPUs used only ZeRO-3, at the measured rate of about 400 TFLOP/s per GPU, the ratio would be about 8.
- Also, each GPU must process at least one full sequence: with long sequences, the activations of one sequence can exceed one GPU's memory capacity.
- The first limit comes from each GPU computing the full model, and the second from each GPU processing full sequences. TP splits the matrix multiplies within a layer, so that several GPUs compute the same layer together; CP splits the sequence.
TP · SP · CP
Intra-layer parallelism- Problem
With many GPUs, syncing parameters takes longer than computing; with long sequences, the activations of one sequence exceed one GPU's memory capacity.
- Solution
- TPtrade-off
Tensor parallelism: splits each matrix multiply into $$t$$ parts across $$t$$ GPUs; the $$t$$ GPUs compute one copy of the model together, and $$N$$ in the formula above becomes the number of groups $$N/t$$.
Improvesmemory: training state and the activations inside attention and FFN drop to $$1/t$$Improvesdata movement: across machines, each GPU syncs only $$1/t$$ of the parametersSacrificesdata movement: within a machine, 2 all-reduces per layer in each pass; the forward pass waits on themSPwaste removalSequence parallelism, here meaning the Megatron-LM approach: TP does not split LayerNorm and Dropout; each GPU stores its own copy of their activations, about 3/4 of a layer's activations at $$t = 8$$ (not counting the attention score matrices). SP splits them along the sequence into $$t$$ parts and rewrites the original all-reduce as a reduce-scatter and an all-gather of the same total volume.
Improvesmemory: these activations drop to $$1/t$$Data movement unchangedCPtrade-offContext parallelism, the other class of sequence parallelism, for long sequences (Li et al. call it sequence parallelism; Ring Attention and DeepSpeed-Ulysses came later): the activations of the whole layer are split along the sequence into $$c$$ parts, and the K and V of other positions that attention needs are obtained through communication. Ring Attention passes them block by block around a ring, overlapped with computation; DeepSpeed-Ulysses switches to splitting by head with an all-to-all (each GPU sends different data to every other GPU); Llama-3 first all-gathers all K and V.
Improvesmemory: activations of the whole layer drop to $$1/c$$Sacrificesdata movement: K and V sent in every layer
TP's communication limits its own scale, so crossing machines needs a split with little communication
NVLink within a machine 450 GB/s, cross-machine network 50 GB/s (one direction)
- TP does 2 all-reduces per layer in each of the forward and backward passes; the forward pass must wait for each all-reduce to finish before it continues, so they are hard to overlap with computation.
- This communication has to run on NVLink within a machine; the cross-machine network has only about 1/9 of its bandwidth. So TP is limited to one machine; a machine has 8 GPUs, so $$t \le 8$$.
- With TP alone, 8 GPUs cannot hold a large model: the training state of 405B is about 6.5 TB, and 8 H100s have 640 GB in total. Combined with ZeRO-3, the number of groups is $$N/8$$, and at Llama-3's scale the ratio of communication time to computation time only drops from about 8 to about 1.
- A split across machines needs little communication: PP splits by layer and passes only the activations at the boundaries between stages, so the volume is far smaller than the parameter count.
PP
Pipeline Parallelism- Problem
TP is limited to one machine; TP alone cannot hold a large model, and combined with ZeRO-3, syncing parameters across machines still takes about as long as computing.
- Solution
Splits the $$L$$ layers into $$p$$ stages, which pass only activations between them. The batch is split into $$m$$ micro-batches fed in one after another, and the stages process different micro-batches at the same time; at the start and end, some GPUs wait, which is called the bubble, a fraction $$(p-1)/(m+p-1)$$ of the time. A larger $$m$$ shrinks the bubble, but with micro-batches that are too small, arithmetic intensity drops and fixed costs take a larger share.
None of the three kinds of parallelism is enough alone, so they are combined
- The limits of the three: TP is limited to one machine; PP has the bubble, and micro-batches cannot be too small; for DP (data parallelism, which covers both DDP and ZeRO), each doubling of the number of groups halves the tokens per group, while the global batch has a limit (the critical batch size).
- The three limits have different sources: TP lacks cross-machine bandwidth, PP lacks enough micro-batches, and DP lacks room for the global batch to keep growing.
- Because the sources differ, each kind of parallelism can be placed where its limit does not apply.
- This largely fixes the combination: TP within a machine, PP across machines, DP at the outermost level. This is 3D parallelism.
3D parallelism
TP × PP × DP- Problem
TP is limited to one machine, PP has the bubble, and DP is limited by the global batch; none is enough alone.
- Solution
$$N = t \times p \times d$$ ($$d$$ is the number of DP groups). One configuration of Llama-3 405B pretraining (sequence length 8K) uses $$8 \times 16 \times 128 = 16384$$ H100s, with FSDP for DP; the 128K-sequence stage changes to TP 8, CP 16, PP 16 and DP 8, which the report calls 4D parallelism. By the formula above, communication time ÷ computation time $$= R \cdot d/(B_\text{net} \cdot G)$$:
| Split | Groups $$d$$ | Communication time ÷ computation time |
|---|---|---|
| ZeRO-3 only | 16384 | about 8 |
| Add TP ($$t = 8$$) | 2048 | about 1 |
| Then add PP ($$p = 16$$) | 128 | about 0.06 |
Llama-3's FSDP does not free parameters after the forward pass and communicates $$2\Phi$$ elements per step; gradients are sent in FP32, for $$6\Phi$$ bytes in total, the same as the $$3\Phi \times 2$$ bytes used in the table.
As the parameter count keeps growing, can the operations per token stay the same?
- 3D parallelism solves the memory and training-time problems of large models.
- Scaling laws give a reason to keep adding parameters: with the same data, more parameters give a trained model with lower loss (the gap between predictions and correct answers).
- But in a dense model each token passes through all parameters: doubling the parameters doubles the operations per token, and doubles the cost of both training and inference.
- The goal is to decouple the parameter count from the operations per token.
- Most parameters are in the FFN (about 2/3 of each layer in the standard architecture), and each token passes through it separately: replacing the FFN with $$E$$ FFNs of the same structure and sending each token through only $$k$$ of them gives MoE.
- With many experts, one GPU cannot hold all of them; with ZeRO-3, each layer would fetch the parameters of all $$E$$ experts, while each token uses only $$k$$ of them. Keeping the parameters in place and sending tokens to the GPUs that hold their experts is EP.
MoE · EP
Mixture of Experts · Expert Parallelism- Problem
In a dense model, the operations per token are proportional to the parameter count; with MoE, one GPU cannot hold all the experts, and with ZeRO-3 each layer has to fetch the parameters of all experts.
- Solution
- MoEtrade-off
The FFN in each layer is replaced with $$E$$ FFNs of the same structure with separately trained parameters, called experts; a router picks $$k$$ of them for each token. The parameter count grows with $$E$$, and the operations per token grow only with $$k$$. DeepSeek-V3 has 671B parameters in total, and each token passes through 37B of them.
Improvescompute: $$\Phi$$ in $$6\Phi D$$ counts the parameters each token passes through, 671B → 37BSacrificesload balance: the router may concentrate tokens on a few experts, so extra balancing is neededEPtrade-offSpreads the experts across GPUs, so each GPU stores only some of them; each token is sent to the GPU holding the expert it picked and sent back after the computation. Compared with ZeRO-3, the parameters stay in place and tokens are sent instead. In DeepSeek-V3, each MoE layer has 256 routed experts (experts chosen by the router), spread over 64 GPUs in training, 4 per GPU; in decode, each GPU holds only 1.
Improvesmemory: each GPU stores only some of the expertsSacrificesdata movement: 2 all-to-alls per MoE layer in each of the forward and backward passes
The comparison here is with a dense model of the same parameter count, and the two use the same total GPU memory; compared with a dense model of the same operations per token, MoE stores about 634B more parameters, which EP spreads over more GPUs.
After deployment, the model generates tokens one by one for each request
Without storing intermediate results, the total operations to generate n tokens grow at least quadratically in n
- Once trained, the model is deployed and generates answers to requests; inference has only forward computation.
- Inference has two phases: prefill feeds the whole input into the model at once; decode generates new tokens one at a time, and computing the next token needs the previous one.
- To generate the $$n$$-th token, attention needs the key and value vectors (K, V) of every earlier position.
- Without storing them, each step recomputes the whole preceding context.
KV cache
Storing K and V of earlier tokens- Problem
Generating the $$n$$-th token recomputes the previous $$n - 1$$ tokens.
- Solution
Stores the K and V of each layer; each step computes only the new token, feeding in 1 token per request per step.
After the KV cache, each decode step computes only 1 token but still reads all the weights once, and the bottleneck moves from compute to memory bandwidth
Llama 2 7B: 13.5 GB of weights read per step; at batch 1, under 1% of compute is used
$$Q_\text{w}$$: bytes of the weights; $$b$$: number of requests computed together (batch); $$n$$: context length; $$q_\text{KV}$$: KV cache bytes per token.
- In training, many tokens pass through the model together; the matrix multiplies are large and limited by compute.
- Each decode step reads all the BF16 weights once and does 2 FLOP for every 2 bytes read; with $$b$$ requests computed together, the arithmetic intensity of the weight part is about $$b$$ FLOP/byte; when $$b$$ is far below the ridge point of 295, decode is limited by memory bandwidth.
- At batch 1, under 1% of compute is used (details in section 6 of the roofline post). With a different bottleneck, training-side optimizations cannot be reused directly.
- Following the formula above, the optimizations fall into two groups: increase $$b$$ so that each read of the weights generates more tokens; make each step shorter and the steps fewer.
Larger batches
Group 1- Problem
Each decode step reads all the weights once, only to generate 1 token for each of $$b$$ requests.
- Limit
$$b$$ is limited by memory capacity: $$Q_\text{w} + b \cdot n \cdot q_\text{KV} \le$$ memory capacity. Llama 2 7B, a 4096-token context, an 80 GB H100: weights 13.5 GB, KV cache 2.1 GB per request, $$b$$ at most 30. Each step then reads about 78 GB, 64 GB of it KV cache; each request reads its own KV cache, so a larger $$b$$ does not amortize these reads.
GQA (several query heads share one set of K and V) and MLA (K and V compressed into short vectors) both reduce $$q_\text{KV}$$.
- Solution
Of the four techniques, PagedAttention raises the $$b$$ that continuous batching can reach; chunked prefill and prefix caching are each independent.
Continuous batchingwaste removalRequests vary in length, and starting and ending a whole batch together leaves empty slots. Instead, at every step, finished requests leave and new requests join.
Improvesdata movement: more useful tokens per weight readImprovesqueueing time of new requestsPagedAttentionwaste removalThe limit of $$b$$ also depends on how well the KV cache memory is used. When memory is reserved for the maximum length, about 20% holds useful data (only about 38% even when the output length is known in advance). Instead, memory is allocated in fixed-size blocks, and the blocks need not be contiguous. Addressing by block makes the attention kernel about 20% to 26% slower, but $$b$$ can grow and end-to-end throughput rises 2 to 4 times, so this is counted as an implementation overhead.
Improvesmemory: about 96% useful dataChunked prefilltrade-offWhen a new request's prefill is inserted, requests in decode wait for it to finish. Chunked prefill splits the prefill into chunks and computes one chunk together with decode each step, using the compute that decode leaves spare.
Improvesdecode stallsSacrificeslatency of a new request's first tokenSacrificesdata movement: each chunk rereads the KV cache of the chunks before itPrefix cachingwaste removalMany requests share the same prefix (for example, a system prompt). Prefix caching keeps the KV cache of computed prefixes in GPU memory for reuse; the cache uses only spare GPU memory, and when the batch needs to grow, the least recently used parts are evicted first.
Improvescompute: repeated prefillImprovesmemory: one copy of each shared prefix
Independent of larger batches, Group 2 makes each step shorter and the steps fewer
Generation time of one request ≈ steps × bytes read per step ÷ achieved memory bandwidth
- Larger batches let each read of the weights serve more requests; Group 2 shortens the generation time of a single request.
- Each of the three factors above has one technique: Flash-Decoding raises the achieved bandwidth, quantization reduces the bytes read per step, and speculative decoding reduces the steps of the original model. The three can be used together.
Shorter and fewer steps
Group 2- Problem
Larger batches raise throughput but do not shorten the generation time of a single request: that time is set by the number of steps, the bytes read per step and the achieved memory bandwidth.
- Solution
- Flash-Decodingwaste removal
H100 has 132 SMs (groups of compute units that run independently); the more SMs read GPU memory at once, the closer the achieved bandwidth gets to the peak. FlashAttention divides work among SMs by batch, head and query block; in decode the query has only 1 position, and with a small batch there are fewer pieces of work than SMs. Flash-Decoding also splits the KV cache into chunks along its length, spreads them over more SMs, and at the end merges the chunks' results with a small kernel. The Flash-Decoding blog post reports that on A100 with long contexts, attention is up to about 50 times faster than with FlashAttention.
Improvesdata movement: idle SMs also read the KV cache, raising achieved bandwidthQuantizationtrade-offQuantizes weights from BF16 to a lower precision such as INT4; when only weights are quantized, they are read in low precision and converted back to BF16 before the computation.
Improvesdata movement: weight bytes read per step about 1/4Sacrificescompute: converting back to BF16, using spare computeSacrificesnumerical accuracy: GPTQ and similar methods use calibration data to reduce the loss of accuracySpeculative decodingtrade-offGeneration yields 1 token per step, while verifying several given tokens takes only one forward pass. A small model first generates $$k$$ candidates, the original model verifies $$k+1$$ positions in one forward pass, and a rejection-sampling rule decides how many are accepted; the output distribution is the same as generating one by one. The gain depends on the fraction of candidates accepted; with a large batch, compute is no longer spare and the gain shrinks.
Improvesdata movement: the original model reads its weights less often (the small model also reads its weights, but they are far smaller)Sacrificescompute: verifying $$k+1$$ positions
Prefill and decode have different bottlenecks and interfere with each other on the same GPUs
- Prefill processes the whole input at once and is limited by compute; decode is limited by memory bandwidth.
- Chunked prefill computes the two together on the same GPUs and reduces decode stalls; but each decode step still waits for the prefill chunk in the same batch to finish, and the two must use the same parallelism and batch.
- As with 3D parallelism, the two are placed separately according to the source of each limit.
P/D disaggregation
prefill and decode deployed separately- Problem
Prefill is limited by compute and decode by memory bandwidth; sharing the same GPUs, they interfere with each other.
- Solution
Prefill and decode run on two sets of GPUs; after prefill, the KV cache is sent to the decode set, once per request. DistServe reports that, with 90% of requests meeting latency targets, the requests served per second per GPU are up to 7.4 times those of the compared systems, or latency targets 12.6 times tighter can be met.
Finally, RL in post-training puts the training and inference systems into one loop
- At this point, training and inference each have their own system, each targeting its own bottleneck.
- One RL step: the model generates answers (rollout), the answers are scored, and the scored samples update the parameters.
- Generation is an inference workload, mostly decode, limited by memory bandwidth; the update is a training workload, limited by compute. One loop contains both workloads.
- The training engine lacks inference optimizations such as KV cache management and continuous batching, so generating with it is slow.
- The two systems share one set of parameters that is updated every step, so they must sync every step.
RL post-training
Training engine and inference engine- Problem
The training engine lacks inference optimizations such as KV cache management and continuous batching, so generation is slow.
- Solution
Rollout uses an inference engine (vLLM, SGLang) and the update a training engine (FSDP, Megatron-LM); each step, the new parameters are synced to the inference engine and resharded to its partitioning. When run synchronously, $$T_\text{RL} = T_\text{rollout} + T_\text{update} + T_\text{sync}$$.
Two deployments: colocated, the two engines take turns on the same GPUs; disaggregated, the two engines are on different GPUs. When disaggregated, the next batch of rollouts can also start early with the previous version of the parameters; the parameters that generate the data are then one or more steps behind (off-policy), and the training algorithm must tolerate this difference.
Summary
| Technique | Cause | Improves | Sacrifices | Type |
|---|---|---|---|---|
| Single-GPU training | ||||
| Mixed precision | FP32 matrix multiplies skip Tensor Cores; activations take 4 bytes per number | ComputeMemory: activations | Numerical accuracy | trade-off |
| Checkpointing | Activations exceed memory capacity | Memory: activations | Compute: one extra forward pass | trade-off |
| Kernel fusion | Intermediate results read from and written to GPU memory repeatedly | Data movementKernel launches | — | waste removal |
| FlashAttention | n × n S and P written back to GPU memory | Data movementMemory | Compute: backward recomputation | trade-off |
| Tiling and pipelining | Matrix multiplies wait for data | Data movementCompute | — | waste removal |
| Multi-GPU distributed training | ||||
| DDP | One GPU's compute has a limit | Compute: time about 1/N | Data movement: gradient sync | trade-off |
| ZeRO-1/2 | N GPUs store N identical copies of the training state | Memory | — | waste removal |
| ZeRO-3 | Parameters still not sharded | Memory | Data movement: 2Φ → 3Φ | trade-off |
| TP | With many GPUs, syncing parameters takes longer than computing | MemoryData movement: cross-machine | Data movement: within a machine | trade-off |
| SP (Megatron-LM) | Activations TP does not split are stored on every GPU | Memory | — | waste removal |
| CP | Long sequences: activations of one sequence do not fit | Memory | Data movement: K and V | trade-off |
| PP | TP is limited to one machine | MemoryData movement: cross-machine | Compute: bubble | trade-off |
| 3D parallelism | Each of the three kinds of parallelism has a limit | Data movement: cross-machine | — | combination |
| MoE | Operations per token proportional to parameter count | Compute | Load balance | trade-off |
| EP | One GPU cannot hold all the experts | Memory | Data movement: all-to-all | trade-off |
| Inference | ||||
| KV cache | Each step recomputes the whole preceding context | Compute | Memory | trade-off |
| Continuous batching | Empty slots in the batch | Data movementQueueing time | — | waste removal |
| PagedAttention | GPU memory reserved for the maximum length | Memory | — | waste removal |
| Chunked prefill | New requests' prefill stalls decode | Decode stalls | First-token latencyData movement: rereading the KV cache | trade-off |
| Prefix caching | Repeated prefill of the same prefix | ComputeMemory | — | waste removal |
| Flash-Decoding | With a small batch, few SMs read the KV cache | Data movement | — | waste removal |
| Quantization | Each step reads all the weights once | Data movement | Compute: converting back to BF16Numerical accuracy | trade-off |
| Speculative decoding | Each step yields only 1 token | Data movement: fewer steps | Compute: verification | trade-off |
| P/D disaggregation | Prefill and decode share GPUs and interfere | Latency | Data movement: transferring the KV cache | trade-off |
| RL | ||||
| Training engine + inference engine | The training engine generates slowly | ComputeData movement | Data movement: syncing parametersSimplicity | trade-off |
Sources
- H100 compute, memory capacity and bandwidth, NVLink: NVIDIA, H100 Tensor Core GPU specifications (SXM)
- H100 SM count (132): NVIDIA, NVIDIA Hopper Architecture In-Depth, 2022
- Mixed precision: Micikevicius et al., Mixed Precision Training, ICLR 2018
- BF16 training: Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training, 2019
- activation checkpointing: Chen et al., Training Deep Nets with Sublinear Memory Cost, 2016
- Activation estimate, Megatron-LM's SP: Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models, MLSys 2023
- GPT-3 XL architecture: Brown et al., Language Models are Few-Shot Learners, NeurIPS 2020
- Shares of operations and time by kernel type (BERT-large, V100): Ivanov et al., Data Movement Is All You Need: A Case Study on Optimizing Transformers, MLSys 2021
- online softmax: Milakov and Gimelshein, Online Normalizer Calculation for Softmax, 2018
- FlashAttention: Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS 2022
- Llama-3 405B parameters, token count, global batch, GPU count, network and parallel configuration: Llama Team, The Llama 3 Herd of Models, 2024
- ZeRO: Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, SC 2020
- TP: Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism, 2019
- CP (Li et al. call it sequence parallelism): Li et al., Sequence Parallelism: Long Sequence Training from System Perspective, 2021
- Ring Attention: Liu et al., Ring Attention with Blockwise Transformers for Near-Infinite Context, 2023
- DeepSpeed-Ulysses: Jacobs et al., DeepSpeed Ulysses: System Optimizations for Enabling Training of Extreme Long Sequence Transformer Models, 2023
- Combining TP, PP and DP, and communication analysis: Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM, SC 2021
- PP and the bubble: Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism, NeurIPS 2019
- critical batch size: McCandlish et al., An Empirical Model of Large-Batch Training, 2018
- scaling law: Kaplan et al., Scaling Laws for Neural Language Models, 2020
- MoE: Shazeer et al., Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer, ICLR 2017
- EP: Lepikhin et al., GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding, ICLR 2021
- DeepSeek-V3 parameter count and EP configuration: DeepSeek-AI, DeepSeek-V3 Technical Report, 2024
- Llama 2 7B layer count and width (32 layers, 4096): Touvron et al., LLaMA: Open and Efficient Foundation Language Models, 2023
- Llama 2 7B does not use GQA: Touvron et al., Llama 2: Open Foundation and Fine-Tuned Chat Models, 2023
- continuous batching: Yu et al., Orca: A Distributed Serving System for Transformer-Based Generative Models, OSDI 2022
- PagedAttention: Kwon et al., Efficient Memory Management for Large Language Model Serving with PagedAttention, SOSP 2023
- GQA: Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints, EMNLP 2023
- MLA: DeepSeek-AI, DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model, 2024
- chunked prefill: Agrawal et al., Taming Throughput-Latency Tradeoff in LLM Inference with Sarathi-Serve, OSDI 2024
- prefix caching: Zheng et al., SGLang: Efficient Execution of Structured Language Model Programs, NeurIPS 2024
- Flash-Decoding: Dao et al., Flash-Decoding for long-context inference (PyTorch blog), 2023
- Quantization: Frantar et al., GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers, ICLR 2023
- speculative decoding: Leviathan et al., Fast Inference from Transformers via Speculative Decoding, ICML 2023
- P/D disaggregation: Zhong et al., DistServe: Disaggregating Prefill and Decoding for Goodput-optimized Large Language Model Serving, OSDI 2024
- The two engines of RL and their deployment: Sheng et al., HybridFlow: A Flexible and Efficient RLHF Framework, EuroSys 2025
- Asynchronous RL with rollouts starting early: Noukhovitch et al., Asynchronous RLHF: Faster and More Efficient Off-Policy RL for Language Models, ICLR 2025