skip to content

Data Parallelism

Every device holds a full copy of the model, works on its own slice of the batch, and gradients are averaged before the update. Interviewers ask where the synchronisation point is and what it costs.

on this pageshow

questions

4

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

level: middleimportance: must knowfreq 70%

answer

  1. Model copied, data divided
  2. Gradients travel, weights do not
  3. One reduction before the step
  4. Effective batch equals replicas times micro-batch
  5. Averaged gradient is the same everywhere

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.

solid answer

~50 s

Data parallelism replicates the whole model on each device - weights, gradients and optimizer state - and splits the batch. With eight devices at a micro-batch of 32, each device runs forward and backward on its own 32 samples and produces its own local gradient; the effective batch is 256. Before the optimizer step, an all-reduce averages those gradients element-wise, so every device ends the reduction holding the same averaged gradient. When the micro-batches are equal in size and the loss is a mean over samples, that average is exactly the gradient of the mean loss over all 256 samples. Each device then applies the same update rule to the same starting weights, so the replicas stay identical without ever exchanging weights. Only gradients cross the interconnect, once per step, and the volume depends on the parameter count - not on the batch size.

go deeper

for a junior

Be ready to state the shape: one full copy of the model per device, the batch split between them, gradients averaged before the update. Knowing effective batch equals per-device batch times device count is the piece most often asked.

for a middle

You are expected to explain the mechanics - why averaging equal-sized per-replica gradients reproduces the full-batch gradient, that an all-reduce leaves every rank with the same result, and why no weight exchange is needed after initialization.

for a senior

Show you have operated this: disjoint data shards per replica, ragged-shard weighting, overlapping the reduction with the backward pass, and how you would detect replicas that have silently diverged.

for a principal

Own the framing that data parallelism trades interconnect bandwidth for throughput and buys nothing else. Be able to say when that trade stops paying and what you would move to instead, in cost terms rather than tooling terms.

## The contract Synchronous data parallelism is one of the simplest scaling ideas in deep learning, and its whole value comes from a single invariant: **at the start of every step, all replicas hold identical weights**. Everything about the design exists to preserve that invariant while letting each device do independent work. The pieces: - **Replicated:** the parameters, the gradient buffers, and the optimizer state (for a moment-based optimizer, one or two extra buffers the size of the parameters). Every device holds its own full copy of all of it. - **Split:** the batch. A global batch is partitioned into per-device micro-batches, and each device sees samples no other device sees. - **Communicated:** the gradients, and only the gradients. Once per step, after the backward pass and before the optimizer step. ## One step, concretely Eight devices, micro-batch 32 each. Effective batch = 8 x 32 = 256. 1. Each device pulls its own 32 samples from the input stream and runs the forward pass through its own copy of the weights. Activations stay local - they are never communicated. 2. Each device runs the backward pass and ends up with a gradient `g_i` for i = 1..8. These eight gradients are all different, because they came from different data. 3. An **all-reduce** with a mean (or sum-then-scale) reduction combines them: every device ends up holding `g = (1/8) * sum_i g_i`. All-reduce is defined as reduce-then-broadcast - the point is that *every* rank gets the *same* result, not just one designated rank. 4. Each device applies the optimizer update using `g`. Same gradient, same optimizer state, same starting weights, same update rule - so the weights after the step are identical everywhere. ## Why averaging is the right reduction If the per-sample loss is `L(x)` and the training objective is the mean over the batch, then the gradient of the mean over 256 samples is the mean of the 256 per-sample gradients. Grouping them into eight equal blocks and averaging the eight block-means gives the same number, because a mean of equal-sized means is the mean of the whole. So the averaged gradient is *mathematically the same* as if one device had run all 256 samples at once (up to floating-point reassociation). The equal-size condition matters. If one replica gets a ragged short shard - say 11 samples while the rest have 32 - the unweighted average of per-replica means over-weights those 11 samples relative to a true mean over all samples. Fixes are to weight each replica's contribution by its sample count, to drop the remainder, or to pad and mask. In practice the distortion is small, but it is the kind of thing that quietly changes results between a 1-device and an 8-device run. ## What is not communicated - **Weights are not averaged.** A candidate who says the devices average their weights each step has described a different (and worse) algorithm. Weights are broadcast exactly once, at initialization, so the replicas start from the same point; after that, identical inputs to an identical update rule keep them together. - **Activations are not communicated.** That is precisely what distinguishes this from splitting the model itself across devices. - **Non-parameter buffers** - normalization running statistics, step counters, dropout masks - are not part of the gradient reduction. Different dropout masks per replica are harmless and even desirable; unsynchronized normalization statistics are not always harmless. ## Cost, and where it can be hidden For a bandwidth-optimal ring reduction, each device sends and receives about `2 * P * (N-1)/N` bytes per step for `P` bytes of gradient - roughly `2P`, and essentially flat in the number of devices `N`. The number of latency-bound stages does grow with `N`. Note what is absent from that expression: the batch size. Doubling the micro-batch doubles the compute per step and leaves the communication unchanged, which is the single most useful lever for making the reduction cheap in relative terms. Much of the cost can be overlapped. The backward pass produces gradients last-layer-first, so the last layer's gradient can begin reducing while the earlier layers are still computing. Grouping several small tensors into one message amortizes per-message overhead. Done well, only the tail of the reduction is exposed. ## Classic bug If every replica seeds its data shuffling identically, all eight see the *same* 32 samples. The all-reduce still runs, the loss curve still looks plausible, and you are paying for eight devices to compute one micro-batch's gradient eight times. Each replica must be handed a disjoint shard of the data.

  • If one replica ends up with a short final shard, what does a plain gradient average actually compute?
    It computes the unweighted mean of the per-replica means, which over-weights the samples on the short replica relative to a true mean over all samples. With 32 samples on seven replicas and 11 on the eighth, those 11 samples carry the same total influence as any full shard. Weight each replica's gradient by its sample count, drop the remainder, or pad and mask.
  • Do the replicas ever need to exchange weights?
    Once, at initialization - either broadcast from one rank or produced from an identical seed - so everyone starts at the same point. After that, no. Identical starting weights plus the identical averaged gradient plus the identical update rule keep them together. If replicas drift, that is a bug: different initialization, an unsynchronized buffer, or state being updated outside the shared step.
  • Why can most of the gradient communication be hidden behind computation?
    The backward pass produces gradients in reverse layer order, so the last layers' gradients are finished while the early layers are still being computed. Their reduction can start immediately and run concurrently with the rest of the backward. Bundling many small tensors into fewer, larger messages amortizes per-message latency. Only the reduction of the earliest layers is unavoidably exposed.

Eight accountants each total a different stack of receipts, then pool the totals into one agreed figure before every one of them writes that identical number into their own identical ledger.

saying these in an interview costs you the question

  • Says every device trains on the same data and results are averaged
  • Claims the weights themselves are averaged across devices each step
  • Describes data parallelism as splitting the model across devices
  • Thinks communication volume grows with the batch size
  • Believes replicas drift apart and get resynchronized occasionally
  • Cannot say whether activations cross the interconnect

context

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

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

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