skip to content

GPU Training Fundamentals

You will learn the hardware realities of training: why batch size is bounded by GPU memory, what mixed precision and loss scaling buy you, and when gradient accumulation or data parallelism is the right fix. Interviewers use these to separate practitioners who have trained real models from paper readers.

on this pageshow

explore

questions

page 1 of 2

What does activation checkpointing trade away to cut a training step's activation memory?

level: juniorimportance: must knowfreq 58%

answer

  1. one resource bought with another
  2. keep the boundaries, drop the middle
  3. rebuild what was dropped in backward
  4. about one extra forward per step

basics

~20 s

Activation checkpointing trades compute for memory. Instead of keeping every layer's forward outputs alive until the backward pass, it keeps only a few segment boundaries and recomputes the rest on demand, costing roughly one extra forward pass per step.

solid answer

~50 s

By default every intermediate tensor produced in the forward pass has to stay resident until the backward pass reaches it, so peak activation memory grows with depth and with batch size, sequence length or resolution. Activation checkpointing cuts the network into segments, keeps only the tensor entering each segment, and discards the segment's interior. When the backward pass arrives, that segment is run forward again from its saved input to rebuild the interior values, the local gradients are formed, and the rebuilt values are freed immediately. Only one segment's interior is live at a time. The price is roughly one extra forward pass over the checkpointed region per step — about a third more compute, since a backward pass costs roughly twice a forward one. The gradients are mathematically unchanged as long as recomputation is deterministic.

go deeper

for a junior

Be ready to say in one breath what is traded for what: memory is saved, extra compute is spent, and the model's mathematics is unchanged. Also be ready to distinguish it from saving model checkpoints to disk.

for a middle

Explain the mechanics — segment boundaries retained, interiors discarded and rebuilt during the backward pass — and produce the cost estimate of roughly one extra forward pass, about a third more compute per step.

for a senior

Show you know when it is the wrong lever: a compute-bound step with headroom to spare, or a model that already fails at batch size one. Be ready to say how you measured the memory profile before turning it on.

for a principal

Own the throughput argument: recompute costs a fixed fraction of step time, and the case for it is that the batch size it unlocks earns more efficiency back than it spends. Frame it against buying more devices or shrinking the model.

## What actually fills memory during a forward pass When a network runs forward, each layer produces an output tensor, and reverse-mode automatic differentiation needs many of those intermediate values a second time to form local gradients on the way back. The backward pass walks the graph in reverse, so the values produced earliest are needed last. By default they all sit in device memory from the moment they are produced until the backward pass consumes them. Peak activation memory therefore grows with depth and with everything that multiplies the size of a single intermediate: batch size, sequence length, image resolution, channel width. This is the term that usually pushes a training job over the device limit, because unlike the weights it scales with the batch. ## The idea Activation checkpointing — also called gradient checkpointing or rematerialization — refuses to hold that whole stack. The network is cut into segments. Only the tensor entering each segment is retained; everything computed inside a segment is released as soon as the segment's output has been produced. When the backward pass reaches a segment, the segment is executed forward a second time from its retained input, the interior values are reconstructed, the gradients for that segment are computed from them, and the reconstructions are dropped again before moving to the next segment. At any instant only the small set of boundary tensors plus one segment's interior is alive. The technique is trading a resource that is elastic (compute time) for one that is a hard ceiling (device memory). Running slower is a decision; running out of memory is a crash. ## Counting the cost A good accounting rule: if a forward pass costs F, the backward pass costs roughly 2F, because at most layers it forms both a gradient with respect to the layer's input and a gradient with respect to its parameters. An ordinary step is therefore about 3F. Checkpointing the whole network adds one more forward, giving about 4F — roughly a third more compute per step. This is why the technique is usually quoted as costing about thirty percent of step time when everything is checkpointed; checkpointing only part of the network costs proportionally less. Wall-clock time often moves by less than the compute ratio, for a practical reason: teams rarely enable checkpointing and keep the same batch size. The freed memory is spent on a larger batch, which usually improves device efficiency, so some of the recompute cost is earned back as better utilization. ## What it does not help Checkpointing touches the activation term only. Parameters, their gradient buffers and optimizer state are untouched, because none of them is a recomputable intermediate — they persist across the whole step by definition. If a model cannot fit with a batch of one and no activations at all, checkpointing cannot rescue it; that situation needs a different lever such as sharding or reducing the model itself. It also does not reduce the number of parameters, change the learning rate that is appropriate, or alter what the optimizer does. ## Correctness Recomputation reproduces the same values from the same inputs, so the gradients handed to the optimizer are the ones an un-checkpointed run would produce. The exception is any layer whose forward pass is not a pure function of its input — a layer that draws random values, or one that mutates internal state as a side effect of running forward. Those need explicit care so the second forward reproduces the first; otherwise the recomputed values do not match the ones that produced the loss. ## A naming trap worth pre-empting Activation checkpointing has nothing to do with saving model checkpoints to disk for fault tolerance or for resuming a run. The word is shared, the mechanism is not: one is about which intermediate tensors stay in device memory inside a single step, the other is about periodically persisting parameters. Interviewers ask this deliberately, because a candidate who has only read the flag name often conflates the two. ## When to reach for it Use it when activations dominate the memory profile and the model is deep enough that a segment interior is much smaller than the whole stack — long-sequence encoders, high-resolution vision models, anything where the useful batch size is being squeezed to one or two. Skip it when the step is already compute-bound and memory is comfortable, because you would be paying a third of your throughput for room you do not need.

  • Does activation checkpointing change the gradients the optimizer receives?
    Not if recomputation is deterministic — the rebuilt intermediates are the same values the first forward produced, so the gradients match an un-checkpointed run. The exception is a layer whose forward pass is not a pure function of its input, such as one drawing random values or updating internal state, which must be made to reproduce itself on the second forward.
  • Why is the overhead roughly a third rather than double the step time?
    Because the extra work is one forward pass, not a whole step. A backward pass costs about twice a forward pass, so an ordinary step is about three forward-equivalents; adding one recompute makes four. That is a ratio of about 4/3, and it shrinks further if only part of the network is checkpointed.
  • Can activation checkpointing help a model that will not fit even at batch size one?
    Usually not. At batch size one the activation term is already small and the remaining pressure comes from parameters, gradients and optimizer state, none of which checkpointing removes. That case calls for a different lever — sharding state across devices, a smaller model, or lower-precision storage of the persistent tensors.

Like not photographing every step of a recipe: you keep a photo at the start of each stage, and if you need to inspect a middle step you just cook that stage again from its saved starting point.

saying these in an interview costs you the question

  • Confuses it with saving model checkpoints to disk
  • Claims it reduces parameter or optimizer-state memory
  • Says the memory saving is free of compute cost
  • Believes it changes the gradients or the loss by design
  • Thinks the activations are written to disk rather than recomputed

context

open as a page

A training run's GPU sits at 30% utilisation while the host CPU is pinned decoding image tiles — what limits step time?

level: juniorimportance: must knowfreq 60%

basics

~20 s

The input pipeline, not the model. The GPU idles waiting for batches while the host decodes and resizes images. Fixes are overlapping data preparation with compute, adding host-side parallelism, and pre-processing tiles into a cheaper stored form.

open as a page

Why can a gradient that is nonzero in 32-bit floats round to exactly zero in 16-bit?

level: juniorimportance: must knowfreq 64%

basics

~20 s

A 16-bit float's 5 exponent bits bottom out near 6e-5 for normal values and near 6e-8 once subnormals run out. A gradient smaller than that has no representation, so it stores as exactly zero and that weight stops moving.

open as a page

What is mixed-precision training, and where do its speed and memory wins come from?

level: juniorimportance: must knowfreq 76%

basics

~20 s

Mixed-precision training runs the forward and backward passes in a 16-bit float while keeping a single-precision copy of the weights. Wins come from moving half the bytes and from matrix units that multiply 16-bit inputs at much higher throughput.

open as a page

Why does halving the training batch size cut device memory when the weights are unchanged?

level: juniorimportance: must knowfreq 78%

basics

~20 s

Weights, gradients and optimizer state are sized by the parameter count, not by batch size. Activations, the per-layer outputs held for the backward pass, exist once per sample, so halving the batch roughly halves that share.

open as a page

What is arithmetic intensity, and how does it decide whether a GPU operation is compute- or bandwidth-bound?

level: middleimportance: must knowfreq 58%

basics

~20 s

Arithmetic intensity is the floating-point operations an op performs per byte it moves to and from device memory. Compare it with the device's peak-FLOP-rate divided by peak bandwidth: below that ridge point the op is bandwidth-bound, above it compute-bound.

open as a page

In data-parallel training across eight devices, what is replicated, what is split, and what is communicated each step?

level: middleimportance: must knowfreq 70%

basics

~20 s

Every device holds a full copy of the weights and processes its own slice of the batch. Local gradients are then averaged across all devices with a single all-reduce, so every replica applies the identical update and stays in sync.

open as a page

In gradient accumulation, why is each micro-batch's loss divided by the number of micro-batches?

level: middleimportance: must knowfreq 68%

basics

~20 s

Because each micro-batch's loss is usually already averaged over its own samples, so adding N of them gives N times the true full-batch gradient. The extra 1/N factor restores the mean over all samples and keeps the step size honest.

open as a page

Which parts of a mixed-precision training step must stay in single precision, and why?

level: middleimportance: must knowfreq 70%

basics

~20 s

A mixed-precision step keeps the weights the optimizer updates, large reductions, the loss and normalization statistics in 32-bit. Those are the places where a tiny quantity is added to or accumulated with a much larger one, which 16-bit arithmetic destroys.

open as a page

In pipeline-parallel training, what is the pipeline bubble and how do micro-batches shrink it?

level: middleimportance: must knowfreq 58%

basics

~20 s

The bubble is the idle time while the pipeline fills and drains: with S stages, S-1 stage-slots are wasted at each end. Splitting the batch into M micro-batches makes the bubble fraction (S-1)/(M+S-1), so more micro-batches shrink it.

open as a page

Your training job reports 22 GB reserved but 9 GB allocated, then fails to allocate — why?

level: middleimportance: must knowfreq 60%

basics

~20 s

Reserved is what the caching allocator holds from the driver; allocated is what live tensors use. The gap is cached blocks of the wrong sizes, and a new tensor needs one contiguous block none of them can supply.

open as a page

How much device memory does an Adam-trained 1.5B-parameter model's state need before any activations?

level: middleimportance: must knowfreq 66%

basics

~20 s

About 24 GB. Adam training in 32-bit floats costs 16 bytes per parameter - weight, gradient, and two moment estimates at 4 bytes each - so 1.5 billion parameters need roughly 24 GB before any activation exists.

open as a page

What does gradient accumulation let you do when only a few samples fit in device memory?

level: juniorimportance: should knowfreq 55%

basics

~20 s

Gradient accumulation runs several small micro-batches, adding their gradients into one buffer, and updates the weights only after the whole group. You get the update of a large batch while only one micro-batch's activations sit in memory at a time.

open as a page

How do you choose activation-checkpoint segment boundaries in a deep network?

level: middleimportance: should knowfreq 42%

basics

~20 s

Place boundaries so retained memory balances: with n uniform layers in k segments it scales like k + n/k, smallest near k equal to the square root of n. When layers differ, rank candidates by bytes freed per recomputed operation.

open as a page

Does data-parallel replication let you train a model that does not fit on one device?

level: middleimportance: should knowfreq 55%

basics

~20 s

No. Data parallelism puts a full copy of the weights, gradients and optimizer state on every device, so the model's own memory cost is unchanged. Only the activations shrink, because each replica sees a smaller micro-batch.

open as a page

In 16-bit floats, why can adding a 1e-7 update to a weight of 1.0 change nothing?

level: middleimportance: should knowfreq 46%

basics

~20 s

A 16-bit float keeps about 11 significand bits, so the next value above 1.0 is roughly 1.001. A 1e-7 update is far below half that gap, so the sum rounds back to 1.0 and the weight never moves.

open as a page

Why does half-precision training multiply the loss by a large scale factor before backpropagation?

level: middleimportance: should knowfreq 58%

basics

~20 s

Half-precision gradients can be too small to represent and collapse to zero. Multiplying the loss by a large constant scales every gradient up by that factor through the chain rule; the scale is divided back out before the optimizer step.

open as a page

How does tensor parallelism split a transformer feed-forward block's two matmuls across devices?

level: middleimportance: should knowfreq 48%

basics

~20 s

Split the first weight matrix by columns so each device computes its own slice of the hidden activations, then split the second by rows so each device produces a partial sum of the output. One all-reduce adds the partials.

open as a page

Why can activation checkpointing around a stochastic layer corrupt a training run?

level: seniorimportance: should knowfreq 34%

basics

~20 s

The recomputed forward draws fresh random numbers, so the rebuilt activations differ from the ones that produced the loss and the gradients belong to a different sampled network. Fix it by capturing and restoring the generator state per segment.

open as a page

A speech encoder's normalization, activation and residual-add chain dominates step time — what does fusing it into one pass remove?

level: seniorimportance: should knowfreq 40%

basics

~20 s

Fusion removes the trips to device memory between the stages. Unfused, the activation tensor is written and re-read after each operation; fused, it is read once and written once. The arithmetic is identical — only the memory traffic shrinks.

open as a page

How do batch-normalization statistics behave under data parallelism when each replica holds only four samples?

level: seniorimportance: should knowfreq 38%

basics

~20 s

Each replica normalizes with the mean and variance of its own four samples, not of the full effective batch. Gradient averaging cannot fix this, because the statistics are used in the forward pass. Synchronizing them across replicas restores single-device behaviour.

open as a page

Why must global-norm gradient clipping run after the last micro-batch of an accumulation group?

level: seniorimportance: should knowfreq 45%

basics

~20 s

Clipping inside the loop rescales each micro-batch gradient separately, so the sum of clipped pieces is not the clipped full-batch gradient. Clip once after the final backward pass and before the step, so the threshold sees the real update.

open as a page

Your half-precision run needs a tuned loss scaler; why does bf16 let you delete it?

level: seniorimportance: should knowfreq 46%

basics

~20 s

bf16 carries the same exponent range as 32-bit float, so gradients that underflow in fp16 stay representable and no loss scaler is needed. The price is fewer mantissa bits, so 32-bit master weights and accumulation still matter.

open as a page

Your weights fit on one device but Adam's moments do not — how do you shard the optimizer state?

level: seniorimportance: should knowfreq 40%

basics

~20 s

Partition parameters across replicas so each device keeps only its slice of Adam's moments and master weights, updates that slice, then all-gathers the updated weights. State memory falls by the replica count at no extra communication.

open as a page

A training run OOMs at step 8,000 with no code change — leak or peak, and how do you tell?

level: seniorimportance: should knowfreq 52%

basics

~20 s

Log memory per step. Steady growth with a constant increment means retention, usually a logged metric still attached to its graph. Flat memory with one tall spike means a peak: the longest sample, or the evaluation pass.

open as a page

How do you estimate whether a 3D segmentation run on 256-cubed volumes fits in 40 GB?

level: seniorimportance: should knowfreq 45%

basics

~20 s

Budget two piles. Model state, parameter count times bytes per parameter, is often under a gigabyte here. Activations dominate: one 256-cubed feature map with 32 channels in half precision is about 1 GB, and several stay resident.

open as a page

Sixteen data-parallel replicas deliver six times one device's throughput - do you buy more devices?

level: principalimportance: should knowfreq 44%

basics

~20 s

Not before you know where the missing throughput went. At 37 percent scaling efficiency each added device returns about a third of a device, and a step bound by gradient communication only gets worse with more replicas.

open as a page

Step time doubled after offloading stored activations to host memory — what do you check?

level: seniorimportance: nice to knowfreq 24%

basics

~20 s

Check whether the bytes moved each step exceed what the host link can carry in the time available. Offloading pays only when a round trip is faster than recomputing the values and overlaps with compute.

open as a page

Why can a global gradient norm in 16-bit come out infinite when every gradient is finite?

level: seniorimportance: nice to knowfreq 31%

basics

~20 s

A global norm sums squares across every parameter. Each square is finite, but the running total crosses the 16-bit ceiling of 65504 and saturates to infinity. The failure lives in the accumulator, not in any individual gradient.

open as a page

Why do batch-statistics normalization layers not see the full effective batch under gradient accumulation?

level: seniorimportance: nice to knowfreq 32%

basics

~20 s

Each micro-batch is normalized with statistics computed from its own samples alone, because normalization happens in the forward pass, before anything is accumulated. Accumulation enlarges the update, not the sample set a forward pass can see.

open as a page

showing 1–30 of 32