skip to content

How large should a training batch get before extra parallel compute stops paying off?

level: principalimportance: nice to knowfreq 30%

answer

  1. returns flatten past a crossover point
  2. steps to target bottom out
  3. measure gradient disagreement across sub-batches
  4. the crossover rises as training proceeds

basics

~10 s

Up to a critical batch size, doubling the batch roughly halves the steps to a target loss, so parallelism becomes wall-clock savings. Past it, the step count flattens while compute per step keeps doubling.

solid answer

~50 s

There is a batch size beyond which more examples per step stop reducing the number of steps. An empirical model that fits many runs is `S ≈ S_min * (1 + B_crit / B)` for steps to a target loss, with examples processed `E ≈ E_min * (1 + B / B_crit)`. Far below `B_crit`, steps fall almost in proportion to the batch: near-perfect parallel speedup at nearly constant total compute. Far above it, steps flatten at `S_min` while examples grow linearly, so you buy wall-clock time with a rising compute bill. `B_crit` can be estimated from the gradient noise scale, roughly the trace of the per-example gradient covariance over the squared norm of the true gradient. It typically **grows during a run** as the gradient signal weakens, which argues for raising the batch as training proceeds.

go deeper

for a junior

Know that bigger batches help only up to a point, and that past it you are paying more compute per step for almost no reduction in how many steps the run needs.

for a middle

Be able to state the shape of the tradeoff — steps to target falling roughly in proportion to the batch below a crossover and flattening above it — and what the crossover means for total examples processed.

for a senior

Show that you would measure rather than guess: estimate the gradient noise scale from disagreement between sub-batch gradients, and recognise that the crossover moves upward as the run progresses.

for a principal

Own the resource argument. State plainly whether the team is time-bound or compute-bound, and defend the batch size as a conversion rate between the two rather than as a hardware default.

### The question behind the question Scaling the learning rate correctly tells you *how* to move to a bigger batch. It does not tell you *whether* you should. That is a separate, and more senior, decision: at some point extra examples per step stop paying for themselves, and you are spending hardware for a smaller and smaller reduction in the number of steps. ### The tradeoff curve An empirical relationship that fits a wide range of training runs is, for the number of optimizer steps `S` needed to reach a target loss at batch size `B`: `S ≈ S_min * (1 + B_crit / B)` and correspondingly for the total examples processed, `E = S * B`: `E ≈ E_min * (1 + B / B_crit)` Read the two limits. - **`B` much smaller than `B_crit`.** Then `S ≈ S_min * B_crit / B`: the step count is inversely proportional to the batch. Doubling the batch halves the steps, and `E` stays near `E_min`, so you get near-perfect parallel speedup at essentially constant total compute. This is the free-lunch regime, and it is where most under-parallelised runs live. - **`B` much larger than `B_crit`.** Then `S ≈ S_min`: the step count has bottomed out, and doubling the batch doubles `E`. You are paying twice the compute for a negligible reduction in wall-clock time. `B_crit` is the crossover, where the run costs roughly twice `S_min` steps and twice `E_min` examples. It is the natural definition of 'as big as it is worth going', and it makes the tradeoff explicit: below it you are wasting time, above it you are wasting compute. ### Estimating it without a sweep A full sweep over batch sizes is expensive. A cheaper estimate comes from the **gradient noise scale**: roughly the trace of the per-example gradient covariance divided by the squared norm of the true gradient. Intuitively it is the ratio of how much the individual per-example gradients disagree with each other to how strong the common signal among them is. When disagreement dominates, averaging more examples per step buys a lot; when the signal dominates, a small batch already recovers most of the direction and more examples are largely redundant. You can approximate it during a run by computing gradients on two or more disjoint sub-batches and comparing their norms against the norm of their average — the discrepancy is exactly what the covariance trace measures. That is a cheap instrument to add to a training loop and it gives you a number to argue from. ### It moves during training The important operational fact: `B_crit` is not a constant of the problem. Early in training the gradient is large and coherent, so the noise scale is small and a modest batch already captures the direction. Later, as the loss flattens, the common signal shrinks while per-example disagreement does not, so the noise scale rises — often by an order of magnitude or more over a long run. This is the principled argument for **increasing the batch size as training proceeds**: match the batch to a moving target rather than paying for the late-training batch size from step one. ### Making the decision The call is a straight resource tradeoff and it should be argued explicitly: - **Time-bound** — a deadline, an experiment loop that must turn around inside a day, a shared cluster whose idle capacity is free to you. Then pushing past `B_crit` is defensible: you are converting compute you were not otherwise going to use into wall-clock time. - **Compute-bound** — a fixed budget, a metered bill, many experiments competing for the same hardware. Then sitting well above `B_crit` is the most expensive mistake available, because the marginal examples are almost entirely wasted, and the same compute spent on more runs at a batch near `B_crit` buys strictly more information. And note what does **not** help above `B_crit`: raising the learning rate further. Once the step count has bottomed out at `S_min`, the run is limited by how many sequential updates the optimization genuinely requires, and no step size manufactures updates you did not take. The remaining levers are a longer budget, a smaller batch, a change in the form of the update, or accepting the target loss you can reach. ### What a strong answer sounds like Name the regime you are in, say how you would measure it rather than guess, state which resource you are actually short of, and be explicit that the answer differs at the start and end of the same run. A candidate who says 'as large as fits' has answered a memory question, not this one.

  • Why does the critical batch size grow as training proceeds?
    It tracks the ratio of per-example gradient disagreement to the strength of the shared signal. Early on the gradient is large and coherent, so a modest batch already recovers the direction. As the loss flattens the common signal shrinks while the disagreement does not, so more examples per step are needed before the averaged direction is trustworthy. That is the principled case for raising the batch size over the course of a run.
  • You are sitting well above the critical batch size and the run is still slow. What do you change?
    Not the learning rate — once the step count has bottomed out, the run is short of sequential updates and no step size creates them. The real options are a longer epoch budget, a smaller batch so the same compute funds more updates or more parallel experiments, a different form of update such as a layer-wise adaptive scheme, or accepting a less ambitious target loss.
  • How would you measure the gradient noise scale inside an existing training run?
    Compute the gradient on two or more disjoint sub-batches at the same weights and compare the norms of the individual sub-batch gradients against the norm of their average. Their disagreement estimates the trace of the per-example gradient covariance while the average estimates the true gradient's norm, and the ratio gives the noise scale. Sample it periodically rather than every step; it is cheap enough to log throughout.

Adding cashiers to a shop cuts the queue quickly at first, then stops helping once every customer already has one — past that point each extra cashier costs the same and saves nobody any time.

saying these in an interview costs you the question

  • Says the batch should be as large as memory allows
  • Thinks doubling the batch always halves time to target
  • Treats the critical batch size as fixed for a model
  • Raises the learning rate to fix a flattened step count
  • Confuses total compute with wall-clock time

context