← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 6 · Chapters 11–12

Autograd & the Mathematics of Distributed Training

Data parallelism is not an approximation. Running a model on eight shards and averaging the gradients produces bit-for-bit the same mathematical object as running the full batch on one device. That claim needs a proof, and the proof is short — but its assumptions are exactly where real bugs live.

Reverse mode is a chain of vector-Jacobian products

Take a network as a composition $\mathcal L=f_n\circ\cdots\circ f_1$ ending in a scalar loss. Write $x_k$ for the output of layer $k$ and $J_k=\partial x_k/\partial x_{k-1}$ for its Jacobian. The chain rule gives

$$\frac{\partial\mathcal L}{\partial x_0}=J_1^\top J_2^\top\cdots J_n^\top\frac{\partial\mathcal L}{\partial x_n}.$$

The order in which you multiply that product decides everything. Associating from the right multiplies a matrix by a vector at each step; associating from the left multiplies matrix by matrix. For a scalar loss the right-to-left order is dramatically cheaper, and it is what "reverse mode" means.

The only operation autograd ever performs is the vector-Jacobian product $\bar x\mapsto J^\top\bar x$. Crucially, no layer ever materializes $J$. A linear layer's VJP is $\bar Y W^\top$; a ReLU's is an elementwise mask. This is why the backward pass costs a small constant times the forward pass rather than $O(\text{parameters}^2)$.

Two structural facts about the graph follow, and between them they explain every collective operation in the rest of this series.

The two rules of split graphs

1

Fan-out becomes a sum. If a value $v$ is consumed by $k$ downstream nodes $u_1,\dots,u_k$, the multivariate chain rule gives

$$\bar v=\sum_{j=1}^{k}\Big(\frac{\partial u_j}{\partial v}\Big)^{\!\top}\bar u_j.$$

Copying a tensor forward means adding gradients backward. A weight used by every example in a batch is the largest fan-out in the graph.

2

Concatenation becomes a split. If $u=[v_1;v_2]$ then $\bar v_1$ and $\bar v_2$ are just the corresponding slices of $\bar u$. No arithmetic, only routing.

Memorize the pair, not the special cases. Replication forward $\Rightarrow$ summation backward. Splitting forward $\Rightarrow$ concatenation backward. Every distributed strategy later in this series is one of these two statements dressed up in a collective operation.

Why averaging gradients is exact

The training loss over a batch $\mathcal B$ is an average of per-example losses:

$$\mathcal L(\theta)=\frac{1}{|\mathcal B|}\sum_{b\in\mathcal B}\ell(x_b;\theta).$$

Split $\mathcal B$ into $P$ disjoint shards $\mathcal B_1,\dots,\mathcal B_P$ of equal size $|\mathcal B|/P$, one per rank, and let rank $p$ compute its own local mean

$$g_p=\nabla_\theta\Big[\frac{P}{|\mathcal B|}\sum_{b\in\mathcal B_p}\ell(x_b;\theta)\Big].$$

Because differentiation is linear and the shards partition the batch,

$$\frac{1}{P}\sum_{p=1}^{P}g_p =\frac{1}{P}\sum_{p=1}^{P}\frac{P}{|\mathcal B|}\sum_{b\in\mathcal B_p}\nabla_\theta\ell(x_b;\theta) =\frac{1}{|\mathcal B|}\sum_{b\in\mathcal B}\nabla_\theta\ell(x_b;\theta) =\nabla_\theta\mathcal L(\theta).$$
The result. A mean all-reduce of per-rank gradients yields the exact full-batch gradient. Data parallelism changes where arithmetic happens, not what is computed. There is no approximation to trade off, no accuracy penalty to budget for.

Notice this is rule 1 in disguise. The parameters $\theta$ are replicated to every rank, which is a fan-out of width $P$, so the backward pass must sum across that fan-out. The all-reduce is the sum the chain rule demands.

The replicas start identical and stay identical, because every rank applies the same averaged gradient. That invariant is what makes the whole scheme work; if it is ever violated, the ranks silently diverge into different models.

A worked example you can check by hand

Scepticism is healthy, so take the smallest possible case: least squares with one parameter, $\ell(x,y;w)=\tfrac12(wx-y)^2$, so $\partial\ell/\partial w=(wx-y)x$. Use four examples and $w=1$.

RankExamples $(x,y)$Local mean gradient
0$(1,0)$, $(2,1)$$\tfrac12\big[(1)(1)+(1)(2)\big]=1.5$
1$(3,1)$, $(4,2)$$\tfrac12\big[(2)(3)+(2)(4)\big]=7.0$
Mean of the two local gradients$\tfrac12(1.5+7.0)=4.25$
Full batch on one device$\tfrac14\big[1+2+6+8\big]=4.25$

Identical, as promised. What matters is why: both shards had the same number of examples, so a plain mean of means equalled the mean of the whole. Break that condition and the identity breaks with it.

Where exactness quietly breaks

The proof needed three assumptions: the loss is a sum over examples, the shards are equal in weight, and every rank starts from the same $\theta$. Real training violates all three by accident.

SituationWhat goes wrongFix
Uneven token counts per rankLanguage-model loss is a mean over tokens, not sequences. If ranks hold different numbers of real (non-padding) tokens, a plain gradient mean weights every rank equally and therefore weights short-sequence tokens more heavily.Scale each rank's loss by its local token count, all-reduce the gradient with sum, and divide by the all-reduced global token count.
Batch normalizationStatistics are computed within a shard, so the function itself depends on how the batch was split. The gradient is genuinely different, not just rearranged.Use a batch-independent norm (LayerNorm, RMSNorm) — which is what transformers do — or synchronize the statistics explicitly.
Global gradient clippingThe clip coefficient depends on the norm of the whole gradient, which no single rank possesses.All-reduce the sum of squared norms first, then apply the same coefficient on every rank.
A rank drops out of the collectiveWith a variable number of batches per rank, one rank finishes early and the others block forever, or worse, proceed with a stale gradient.Pad the epoch to a common length, or use a join-style protocol.
Non-deterministic reduction orderFloating-point addition is not associative, so different ranks can produce marginally different sums and drift apart over thousands of steps.Rely on the collective to return a bitwise-identical result to all ranks, which correct all-reduce implementations guarantee.
The token-count bug is the common one. It produces a model that trains, converges, and is subtly worse than it should be, with no error message. If sequences are packed to a fixed length it cannot occur; if you use padding and variable-length batches, check it explicitly.

Gradient accumulation is the same identity in time

The proof above split the batch across space. Splitting it across time works for exactly the same reason: run $k$ micro-batches sequentially, add their gradients into the same buffer, then step. That yields the gradient of the concatenated batch, because addition is addition.

$$\nabla\mathcal L_{\text{total}}=\sum_{j=1}^{k}\nabla\mathcal L_j\quad\text{(with each }\mathcal L_j\text{ scaled by }1/k\text{ for a mean loss).}$$

Two practical consequences:

Scale the loss, not the gradient

Divide each micro-batch loss by $k$ before calling backward. Doing it afterwards means the accumulated buffer briefly holds $k$ times the intended magnitude, which interacts badly with FP16 gradient scaling.

Suppress the all-reduce until the last micro-batch

DDP syncs gradients during every backward pass by default. Across $k$ micro-batches that is $k$ times the necessary communication. Frameworks expose a "no sync" context for exactly this; forgetting it is a pure throughput loss with no correctness effect.

The duality table

Once a tensor lives on many devices, its forward placement determines its backward collective. This table is worth internalizing now, because tensor, context and expert parallelism are each just one row of it applied in a different place.

Forward operationBackward operationReasoning
identity (value already replicated)all-reduce (sum)Replication is a fan-out; rule 1 says sum the incoming gradients.
all-reduce (sum of partial results)identityEach rank contributed one addend, so each receives the incoming gradient unchanged.
all-gather (assemble shards)reduce-scatterConcatenation forward routes gradients backward, and the summation over replicas rides along.
reduce-scatterall-gatherThe exact inverse of the row above.
all-to-allall-to-all (transposed)A permutation of data; its adjoint is the inverse permutation.
split / scatterall-gatherSlicing forward, concatenating backward.
Conjugate pairs. Megatron-style tensor parallelism names the first two rows $f$ and $\bar f$ and treats them as a matched pair of autograd functions: one is the identity forward and an all-reduce backward, the other is the mirror image. Writing them as custom autograd functions is all it takes to make a sharded layer differentiate correctly.

What a distributed graph actually looks like

Putting it together, a data-parallel step is a single computation graph whose root fans out to $P$ ranks and whose leaves converge back through a sum:

$\theta$ replicatedfan-out of width $P$
→
$P$ local lossesdisjoint data shards
→
sum of gradientsthe all-reduce

The framework never builds this graph explicitly — each rank builds only its own local slice — but the mathematics is the global graph, and the collective is the piece of the chain rule that crosses device boundaries.

Takeaway

Replication forward means summation backward. That one sentence, plus the linearity of differentiation, is the complete justification for data parallelism. Everything harder in this series comes from applying the same rule to tensors that are sharded rather than replicated.

Check yourself