PidokuInfra

Training vs Inference: The Math

Foundations Intermediate 45 min Difficulty 2/5 Topic 10 of 12

Prerequisites 04, 09, I.02


1. What is it?#

A precise accounting of what training does that inference doesn’t, in FLOPs and bytes. File I.02 gave the operational differences; this gives the arithmetic.


2. Why does it exist?#

Because the “3x rule” (backward pass costs 2x forward) and the “16 bytes per parameter” rule are constantly used in capacity discussions, and you should be able to derive both rather than quote them.


3. Simple analogy#

Building a road versus driving on it. Driving (inference) traverses the route once. Building (training) requires surveying the route, laying it, and then walking back along it to check and adjust every segment — roughly three passes for every one the driver makes.


4. Tiny example#

For y = W x with loss L:

FORWARD:
  y = W x                          FLOPs: 2·M·K  (for W: M×K, x: K)

BACKWARD (two things needed):
  dL/dx = Wᵀ (dL/dy)               FLOPs: 2·M·K   ← to propagate to earlier layers
  dL/dW = (dL/dy) xᵀ               FLOPs: 2·M·K   ← to update this layer

  Total backward: 4·M·K = 2× forward

Hence:

Training step ≈ 3 × forward pass       (1 forward + 2 backward)
FLOPs per token, training ≈ 6P
FLOPs per token, inference ≈ 2P

That’s the derivation of the two most-quoted numbers in the field. 6P for training, 2P for inference.


5. Technical explanation#

Memory accounting for training#

Weights (BF16)              2P
Gradients (BF16)            2P
Optimizer m (FP32)          4P
Optimizer v (FP32)          4P
FP32 master weights         4P
                          ─────
                           16P    bytes

Plus activations stored for backward:
  ≈ L · B · S · d · (10 to 30) bytes, depending on what's recomputed

For a 7B model: 112 GB of state, plus activations. Activation checkpointing trades compute for memory: store only layer boundaries, recompute the rest during backward. Costs ~30% more FLOPs, saves most activation memory.

Memory accounting for inference#

Weights                     P · bytes_per_weight     (1-4 bytes)
KV cache                    2 · L · h_kv · d_head · bytes · S · B
Activations                 transient, small (decode) or large (long prefill)
Framework overhead          1-3 GB

The ratios#

                        Training        Inference (decode)
FLOPs per token         6P              2P
Memory per parameter    16 bytes        1-2 bytes
Batch size              huge (1M+ tok)  small (1-256 sequences)
Arithmetic intensity    very high       very low
Bottleneck              compute         memory bandwidth

Training is compute-bound; decode is memory-bound. That single inversion explains why training hardware choices (maximize FLOPs) and inference hardware choices (maximize bandwidth) diverge, and why an H200 (same FLOPs as H100, 43% more bandwidth) is a better inference chip and an equal training chip.

Why inference can use lower precision#

Training needs precision because gradients are small and accumulate over millions of steps; rounding errors compound. Inference does one forward pass — errors don’t compound across steps (within a single token), and the model’s output is then discretized by sampling anyway.

Training:  BF16 compute, FP32 master weights, FP32 optimizer states
Inference: FP8 or INT8 weights and activations often lose <1% quality

Caveat for autoregressive generation: errors do compound across tokens, since each token conditions the next. This is why aggressive quantization can degrade long generations more than short ones — a real effect that short benchmarks miss (Section VII.02).


6. Under the hood#

Why the backward pass is exactly 2x: each weight matrix participates in two matmuls during backward (gradient w.r.t. input, gradient w.r.t. weight) versus one during forward. The chain rule requires both: one to continue backward, one to update.

The activation memory is the other half of the story. During forward, PyTorch’s autograd retains every tensor an operation will need for its backward. For attention, that historically included the full (B,h,S,S) matrix — which is why FlashAttention’s memory savings mattered for training even more than inference.


7. Performance implications#

  • Training clusters are FLOP-optimized; inference fleets should be bandwidth-optimized.
  • Inference can run on much smaller hardware than training. Do not size from training.
  • Precision choices diverge. Serving at training precision leaves 2-4x on the table.
  • Fine-tuning sits in between. LoRA fine-tuning has training’s structure but inference-like memory (base weights frozen, only small adapters have gradients and optimizer state).

8. Production implications#

  • Multi-LoRA serving exploits this: one base model in memory, many small adapters swapped per request. Serving 50 fine-tuned variants costs barely more than serving one (Section VIII.09).
  • Do not let the training team size your inference cluster with 6P FLOPs and 16P bytes.
  • Evaluate quantized models on long generations, not just short benchmarks, because of error compounding.

9. Common mistakes#

Using 6P for inference FLOPs. It’s 2P.

Sizing inference memory at 16 bytes/param. It’s 1-2 plus KV.

Assuming inference precision must match training precision.

Evaluating quantization only on short-answer benchmarks. Misses compounding degradation.


10. Hands-on exercise#

A. Derive it. For a 2-layer MLP, write out forward and backward by hand and confirm backward is 2x the forward FLOPs.

B. Memory table. For 7B, 13B, 70B: compute training memory (16P + activations) and inference memory (2P + KV at 8k, batch 32). How many GPUs for each?

C. Compounding. Take a model, quantize it to INT4, and compare outputs against FP16 for generations of 10, 100, and 1000 tokens. Measure divergence (e.g. fraction of matching tokens, or perplexity on the generated text). Does degradation grow with length?


11. Interview questions#

  1. Derive why the backward pass costs 2x the forward pass.
  2. Where does “16 bytes per parameter” for training come from?
  3. Why is training compute-bound and decode memory-bound?
  4. Why can inference use lower precision than training?
  5. What is the caveat about error compounding in autoregressive generation?
  6. How does LoRA change the serving economics of many fine-tuned models?

12. Further reading#

  • [ESTABLISHED] Rajbhandari et al., “ZeRO” (2020) §3 — the memory accounting
  • [ESTABLISHED] Kaplan et al. (2020), Appendix B — FLOP accounting
  • [ESTABLISHED] Hu et al., “LoRA” (2021)
  • Next: 11 — Number formats

↑↓ navigate↵ openesc close