The idea in one minute#
Training and inference run the same matrix multiplications on the same GPUs, but they stress the hardware in opposite ways. Training is a long, steady, compute-bound batch job that wants maximum arithmetic and does not care about any individual sample. Inference is millions of small, impatient, often memory-bound requests that care about latency, arrive unpredictably, and never end.
Knowing which one you are doing tells you which specs to pay for.
A picture#
flowchart TB
subgraph T["Training step"]
direction LR
F1["Forward pass<br/>keep every activation"] --> L["Loss"] --> BK["Backward pass<br/>gradients"] --> U["Update weights"]
end
subgraph I["Inference step"]
direction LR
F2["Forward pass only<br/>discard activations"] --> O["Output token"]
end
class F1,F2,BK compute
class L,U neutral
class O ioHow it really works#
| Training | Inference | |
|---|---|---|
| Passes | Forward + backward + update (~3x the arithmetic) | Forward only |
| Memory | Weights + gradients + optimizer state + all activations: ~4–8x the weights | Weights + KV cache |
| Batch | Large and fixed; you choose it | Whatever traffic arrives |
| Regime | Compute-bound (high intensity) | Often memory-bound (low intensity at small batch) |
| Precision | BF16/FP16 mixed, FP32 accumulations | FP16 down to 4-bit |
| Cares about | Samples per second, time to finish | Time to first token, tokens per second, tail latency |
| Duration | Days to months, then done | Continuous for the product’s life |
| Multi-GPU | Essential; heavy GPU-to-GPU traffic every step | Often one GPU per replica; split only when the model does not fit |
| Failure | Checkpoint and restart | A user sees an error |
What that means for hardware#
- Training buys FLOPs and interconnect. Large batches sit above the ridge point, so arithmetic is the limit; and gradients must be exchanged between GPUs every step, so NVLink and the cluster network matter enormously (module VI).
- Inference buys memory capacity and bandwidth. Capacity sets concurrency (V.03). Bandwidth sets per-user speed (II.03). Peak FLOPs matter once batching pushes you over the ridge.
- Performance per watt and per dollar matter more for inference, because it runs forever and its total cost over a model’s life usually exceeds the cost of training it.
This is why the same vendor sells both 1,000 W flagship parts and 70 W inference cards, and why accelerators aimed purely at inference exist at all (VII.02).
The activation memory difference, concretely#
To compute gradients, training must remember the output of every layer for every sample in the batch until the backward pass reaches it. Inference can throw each layer’s output away as soon as the next layer has consumed it. That one difference accounts for most of the memory gap, and it is why a GPU that cannot train a model can often serve it comfortably.
Code#
// footprint.go — the same model, two very different memory bills.
package main
import "fmt"
func main() {
const (
params = 8e9
layers = 32.0
hidden = 4096.0
seqLen = 2048.0
batch = 8.0
fp16 = 2.0
fp32 = 4.0
actsPerLy = 12.0 // tensors of size batch×seq×hidden kept per layer for backward (rough)
)
gb := func(b float64) float64 { return b / 1e9 }
weights := params * fp16
// Training (AdamW, mixed precision): FP32 master weights + gradients + two optimizer moments.
trainState := params * fp32 * 4
trainActs := layers * actsPerLy * batch * seqLen * hidden * fp16
fmt.Printf("TRAINING weights+grads+optimizer %5.0f GB activations %4.0f GB total %4.0f GB\n",
gb(trainState), gb(trainActs), gb(trainState+trainActs))
// Inference: weights + KV cache + one layer's activations at a time.
kv := 2 * layers * 8 * 128 * fp16 * seqLen * batch
inferActs := 4 * batch * seqLen * hidden * fp16
fmt.Printf("INFERENCE weights %5.0f GB KV + activations %4.1f GB total %4.0f GB\n",
gb(weights), gb(kv+inferActs), gb(weights+kv+inferActs))
}The numbers are rough by design. The ratio is the point.
Remember this#
- Training: compute-bound, memory-hungry, throughput-oriented, finite.
- Inference: frequently memory-bound, latency-sensitive, unpredictable, endless.
- Training hardware priorities: FLOPs and interconnect. Inference: memory capacity, bandwidth, efficiency.
- Training keeps all activations for the backward pass; inference discards them.
Try it#
- Run
footprint.go. How many 80 GB GPUs does each case need? - Halve the training batch. Which term shrinks? Which does not?
- A team has H100s bought for training and wants to serve a small model on them. What would you measure to decide whether that is wasteful?
Check yourself#
- Why does training need several times more memory than inference for the same model?
- Which specs matter most for training? For inference?
- Why is inference often the larger lifetime cost?