PidokuInfra

Training vs Inference

Foundations Beginner 45 min Difficulty 1/5 Topic 02 of 11

Prerequisites 01-what-is-inference.md


1. What is it?#

Training and inference are the two phases of a model’s life.

  • Training — figuring out what the numbers in the model should be. You show the model examples, it makes predictions, you measure how wrong it was, and you nudge every parameter slightly in the direction that would have made it less wrong. Repeat a few trillion times.
  • Inference — using those numbers. Input in, output out. Nothing changes.

They share the mathematics of the forward pass and almost nothing else operationally.


2. Why does it exist?#

The distinction exists because the two phases have opposite engineering constraints, and conflating them leads directly to bad systems.

Training is a batch job: you own the input, you choose the batch size, nobody is waiting, and you optimize purely for throughput and cost-to-completion. If a training step takes 900 ms instead of 800 ms, you shrug.

Inference is a service: someone is waiting, you don’t control the input, the load varies by the hour, and a 100 ms regression may violate an SLO and lose customers.

Most tooling, most tutorials, and most ML engineers’ instincts come from the training world. Applying them unmodified to inference is the single most common source of bad LLM deployments.


3. Simple analogy#

Training is writing a textbook. Inference is answering a student’s question using it.

Writing the textbook takes a year, a team, and an enormous library. You can take breaks. Nobody is standing at your desk.

Answering a question takes seconds, happens constantly, and the student is standing right there. You need the book open on your desk (weights resident in memory), you need to find the relevant part fast, and you need to do it for thirty students at once without making any of them wait too long.

Another framing engineers find useful: training is compiling; inference is running the binary. You compile once, expensively, and then execute billions of times. Nobody would accept a program that recompiles itself on every request — yet a surprising number of naive inference deployments effectively do the equivalent (reloading models, re-tokenizing, rebuilding graphs).


4. Tiny example#

The same two-parameter model from file 01, now with training shown:

Go
package main

import "fmt"

func main() {
	// ---------- TRAINING ----------
	w, b := 0.0, 0.0 // start with nothing learned
	// inputs are in units of 1000 sqft so one learning rate suits both w and b
	data := [][2]float64{{1.0, 170000}, {1.5, 245000}, {2.0, 320000}}
	const lr = 0.05

	for epoch := 0; epoch < 20000; epoch++ {
		for _, d := range data {
			x, price := d[0], d[1]
			pred := w*x + b     // FORWARD  (this is inference!)
			err := pred - price // LOSS
			dw := 2 * err * x   // BACKWARD (gradient)
			db := 2 * err
			w -= lr * dw // UPDATE
			b -= lr * db
		}
	}
	fmt.Printf("w=%.1f per 1000 sqft, b=%.1f\n", w, b) // ≈ 150000, ≈ 20000

	// ---------- INFERENCE ----------
	infer := func(sqft float64) float64 {
		return w*(sqft/1000) + b // FORWARD only. No loss. No gradient. No update.
	}
	fmt.Printf("%.0f\n", infer(1500)) // ≈ 245000
}

Count the work per example:

PhaseStepsExtra memory needed
Inferenceforwardjust weights + one activation
Trainingforward + loss + backward + updateweights + gradients + optimizer state + all intermediate activations

That last item is why training memory is so much larger: the backward pass needs every intermediate value the forward pass produced. Inference can throw each intermediate away the instant the next layer consumes it. Inference is memory-cheap per token but must hold the weights; training is memory-expensive but amortizes over huge batches.


5. Technical explanation#

The full comparison#

DimensionTrainingInference
Directionforward + backwardforward only
Weightschange every stepfrozen
Memory for weights (7B, FP16)14 GB14 GB
Plus gradients+14 GB—
Plus optimizer state (Adam)+56 GB (FP32 m, v, master)—
Plus stored activationslarge, ∝ batch × seq × layerssmall, transient
Typical total (7B)~100+ GB~16-20 GB
Batch sizeyou choose (e.g. 4M tokens)dictated by arriving traffic
Sequence lengthfixed / bucketedarbitrary, adversarial
Optimizes forthroughput, cost to convergelatency and throughput and cost
Failure of one stepretry, no user impactuser-visible error
PrecisionBF16/FP8 compute, FP32 master weightsas low as quality permits (FP8/INT8/INT4)
Durationdays-weeks, then doneforever
Numerical determinismnot requiredoften desired, rarely guaranteed
Hardwarebiggest available, tightly coupledheterogeneous, cost-optimized

The memory arithmetic (important, do it yourself)#

For a model with P parameters trained in mixed precision with Adam:

weights (BF16)          2P bytes
gradients (BF16)        2P bytes
Adam m (FP32)           4P bytes
Adam v (FP32)           4P bytes
FP32 master weights     4P bytes
                       ---------
                       16P bytes   ≈ 16 bytes per parameter

For 7B: 112 GB before a single activation. That’s why training a 7B model needs multiple 80 GB GPUs while inference fits comfortably on one.

For inference:

weights (FP16)           2P bytes
KV cache                 varies (Section V — this becomes the dominant term)
activations              small, transient
                        ---------
                        ~2P + KV

For 7B: 14 GB + KV cache. On an 80 GB GPU you have ~64 GB left for KV cache, which — as you will compute in Section V — determines your maximum concurrency. In inference, memory not spent on weights is capacity. That single reframing motivates quantization (Section VII), GQA/MLA (Section XIII), and PagedAttention (Section V).

What carries over and what doesn’t#

Carries over: the operator set (matmul, attention, normalization), the numerics knowledge, the model architecture.

Does not carry over: batch size intuitions, memory profiles, the idea that you can retry freely, the assumption of fixed shapes, the assumption of uniform request cost, DDP-style scaling intuitions, and — critically — the assumption that more GPUs always help.


6. Under the hood#

During training, the framework builds an autograd graph: every operation records what inputs it consumed so gradients can flow backward later. Every intermediate tensor is kept alive by that graph.

During inference you turn that off:

Python
with torch.inference_mode():      # or torch.no_grad()
    out = model(x)

inference_mode (PyTorch ≥1.9) is stronger than no_grad: it also skips version-counter bookkeeping on tensors. Forgetting it is a classic bug — memory grows request over request until OOM, because every activation of every request is still being retained for a backward pass that will never happen.

Also switched: layers that behave differently in the two phases.

Python
model.eval()   # Dropout becomes identity; BatchNorm uses running statistics

Forgetting model.eval() gives you nondeterministic, subtly wrong outputs. For transformers without BatchNorm or dropout at inference the effect is smaller, but it is still a hard requirement of correct deployment.


7. Performance implications#

Training is throughput-only, so every trick that raises throughput is unambiguously good. Bigger batch? Yes. Gradient accumulation? Yes. Slower per-step but better hardware use? Yes.

Inference has a latency constraint, so throughput tricks have a cost. Batch size 256 gives excellent tokens/sec per GPU and terrible TTFT if requests must wait to fill the batch. This tension is the reason continuous batching (Section V.09) was invented: it captures most of the throughput benefit without making requests wait for batch formation.

Training runs at high arithmetic intensity by construction (huge batches, huge sequence blocks). Inference decode runs at catastrophically low arithmetic intensity (one token per sequence per step). This is the single largest performance difference between the two phases, and file 08 makes it quantitative.


8. Production implications#

  • Do not size inference hardware from training hardware. A model trained on 1,024 H100s may serve happily on 2. Conversely, a model that trained fine may be un-servable at your latency target because of its attention design.
  • Inference-aware model design is a real lever. GQA instead of MHA cuts KV cache by 4-8x. A smaller vocabulary shrinks the output projection. A shorter maximum context reduces worst-case memory. These are training-time decisions with permanent inference consequences, and inference engineers should be in the room when they are made (Section XIV).
  • Two different teams, two different on-call rotations. Training failures are batch-job failures. Inference failures are outages. Different tooling, different urgency, different skills — which is exactly why this specialty exists.
  • Fine-tuning muddies the line. LoRA adapters let you serve many variants from one base model (Section VIII.09), which is an inference-architecture question, not a training one.

9. Common mistakes#

Deploying training code as inference code. Training code carries autograd, checkpointing hooks, distributed wrappers, and dataloader machinery. All of it is overhead or outright bugs in production. Export a clean inference path.

Forgetting eval() / inference_mode(). Causes memory leaks and wrong numbers. Check it first whenever inference memory grows over time.

Assuming training precision must equal inference precision. A model trained in BF16 usually serves fine in FP8 or INT8 with proper calibration. Insisting on training precision leaves an easy 2-4x on the table.

Benchmarking inference with training-shaped inputs. Fixed length 2048, batch 32, all sequences identical — that benchmark tells you nothing about production, where lengths vary by 100x and arrivals are Poisson-ish. Benchmark with realistic distributions (Section X.07).

“The model was 95% accurate in training so it’ll be 95% accurate in production.” Different question entirely, and quantization, batching-induced nondeterminism, and prompt distribution shift all matter. Always evaluate the deployed configuration.


10. Hands-on exercise#

A. Memory arithmetic. For a 13B parameter model, compute:

  1. Training memory with Adam mixed precision (use 16 bytes/param). Does it fit on one 80 GB GPU?
  2. Inference memory for weights at FP16, FP8, INT4.
  3. How much is left over on an 80 GB GPU for KV cache in each case?

B. Prove the autograd leak. In PyTorch, run the same small model 100 times in a loop — once with torch.inference_mode() and once without — printing torch.cuda.memory_allocated() (or psutil RSS on CPU) each iteration. Plot both. Explain the shape of each curve.

C. Measure the eval() difference. Take any model with dropout. Run the same input 10 times in train() mode and 10 times in eval() mode. Compare the outputs. Write down what you see and why.


11. Interview questions#

  1. Why does training need ~8x more memory per parameter than inference?
  2. What does torch.inference_mode() do that torch.no_grad() doesn’t?
  3. Your inference service’s memory grows steadily until OOM after ~2 hours. Give three hypotheses in order of likelihood.
  4. When would you deliberately serve at a different precision than you trained at? What could go wrong?
  5. Name a model architecture decision made during training that you cannot undo at inference time, and explain its cost.

12. Further reading#

  • [REFERENCE] PyTorch docs: torch.inference_mode, Module.eval
  • [ESTABLISHED] Rajbhandari et al., “ZeRO” (2020) — for where the 16 bytes/param comes from
  • [ESTABLISHED] Micikevicius et al., “Mixed Precision Training” (2018)
  • Next: 03 — Model anatomy

↑↓ navigate↵ openesc close