PidokuInfra

Training vs Inference on a GPU

Advanced Intermediate 40 min Difficulty 3/5 Topic 04 of 04

Prerequisites 01, 03, IV.01

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 io

How it really works#

TrainingInference
PassesForward + backward + update (~3x the arithmetic)Forward only
MemoryWeights + gradients + optimizer state + all activations: ~4–8x the weightsWeights + KV cache
BatchLarge and fixed; you choose itWhatever traffic arrives
RegimeCompute-bound (high intensity)Often memory-bound (low intensity at small batch)
PrecisionBF16/FP16 mixed, FP32 accumulationsFP16 down to 4-bit
Cares aboutSamples per second, time to finishTime to first token, tokens per second, tail latency
DurationDays to months, then doneContinuous for the product’s life
Multi-GPUEssential; heavy GPU-to-GPU traffic every stepOften one GPU per replica; split only when the model does not fit
FailureCheckpoint and restartA 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#

Go
// 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#

  1. Run footprint.go. How many 80 GB GPUs does each case need?
  2. Halve the training batch. Which term shrinks? Which does not?
  3. 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#

  1. Why does training need several times more memory than inference for the same model?
  2. Which specs matter most for training? For inference?
  3. Why is inference often the larger lifetime cost?

↑↓ navigate↵ openesc close