skip to content

When one embedding id appears many times in a batch, what gradient does its row receive?

level: middleimportance: should knowfreq 44%

answer

  1. gather forward, what comes back?
  2. plus-equals, never equals
  3. untouched rows get exactly zero
  4. one row, forty vectors summed
  5. sparse rows versus a dense buffer

basics

~20 s

The sum of the upstream gradient vectors from every occurrence. The lookup gathers rows, so the backward scatter-adds into them. Overwriting instead of accumulating would apply only the last occurrence and silently discard the rest.

solid answer

~50 s

An embedding lookup is a gather: the forward reads row `id` out of the table for each position in the batch. Its backward is the transpose of a gather, which is a scatter-add. Every row that no position in the batch referenced receives exactly zero, and every referenced row receives the sum of the upstream gradient vectors of all positions that referenced it — so a popular id appearing forty times gets forty vectors added together. The accumulation must be `+=`, never `=`: assignment keeps only whichever occurrence was processed last, which is a silent bug because loss still falls. The gradient is therefore structurally sparse, and on a two-million-row table the interesting question becomes whether to materialise a dense gradient buffer the size of the whole table each step or keep only the touched rows.

code

python · 17 lines
python
# Embedding backward: the forward gathers rows, the backward scatters and ADDS.
n_ids, width = 5, 2
grad_table = [[0.0] * width for _ in range(n_ids)]

batch_ids = [3, 1, 3, 3]                 # id 3 appears three times
upstream = [[1.0, 2.0],                  # gradient arriving at each batch position
            [0.5, 0.5],
            [2.0, 0.0],
            [-0.5, 1.0]]

for pos, row_id in enumerate(batch_ids):
    for c in range(width):
        grad_table[row_id][c] += upstream[pos][c]   # += , never =

print(grad_table[3])   # [2.5, 3.0]  -> 1.0+2.0-0.5 and 2.0+0.0+1.0, summed
print(grad_table[1])   # [0.5, 0.5]  -> the single occurrence
print(grad_table[0])   # [0.0, 0.0]  -> an id absent from the batch stays zero

go deeper

for a junior

Recall that an embedding lookup is just a row copy, so its gradient goes only to the rows the batch actually used, and that a repeated id accumulates rather than replaces.

for a middle

Explain the gather-then-scatter-add structure, why the chain rule produces a sum over occurrences, and why an assignment instead of an accumulation is a silent bug that a duplicate-id gradient check catches.

for a senior

Show the production judgment: sparse versus dense gradient representation on a multi-million-row table, what dense weight decay does to rows the batch never saw, and the frequency imbalance that makes head rows move far faster than tail rows.

for a principal

Own the systems tradeoff — gradient buffer size and cross-device reduction cost against optimizer-state staleness and reproducibility — and decide when a sharded table with restricted updates is worth the extra machinery.

### The forward is a gather, so the backward is a scatter-add An embedding table is a matrix `E` with one row per id and `d` columns. The forward pass for a batch of ids `[i_1, ..., i_B]` copies out rows: `out[b] = E[i_b]`. There is no arithmetic, only indexing, so the local gradient of `out[b]` with respect to `E[i_b]` is the identity and with respect to every other row is zero. The transpose of a gather is a scatter-add. Concretely, for upstream gradient `g` with one row per batch position: ``` dL/dE[r] = sum over all b with i_b == r of g[b] ``` Two facts fall out immediately. Any row `r` that no position referenced gets exactly zero — on a two-million-id user/item table with a batch of a few thousand positions, the overwhelming majority of rows have a structurally zero gradient. And any row referenced more than once gets a **sum**, not a replacement and not an average. ### Why the accumulation must be plus-equals Suppose one very popular item id appears 40 times in a batch. The loss contains 40 terms that each read that same row, so the chain rule sums 40 partial derivatives into it. If a hand-written backward writes `dE[r] = g[b]` inside the loop over batch positions rather than `dE[r] += g[b]`, the row keeps only the last occurrence's gradient and 39 of 40 are discarded. This is a genuinely dangerous bug because it is silent. Training does not crash, gradient magnitudes stay plausible, and loss still decreases — the model just learns the head of the id distribution far more slowly than it should, and precisely the ids that matter most are hurt the worst. A gradient check on a batch containing a duplicated id catches it instantly; a gradient check on a batch of distinct ids never will. ### Structural sparsity and what to do with it Because only the batch's unique ids have nonzero rows, the true gradient is sparse: at most (number of unique ids in the batch) times `d` nonzero values, against a table of two million times `d`. Materialising it densely allocates and zeroes a buffer the size of the entire table on every step, and the reduction across devices moves that whole buffer, which for a large table dwarfs the rest of the model. Keeping only the touched rows — an index list plus a compact value block — is orders of magnitude cheaper, at the cost of a more complicated update path. The complication lands on the optimizer. Momentum, second-moment estimates and weight decay are naturally *dense* operations: they touch every parameter every step regardless of whether that parameter had gradient. Applied densely to a two-million-row table this means every unseen row is decayed toward zero every step, so a rare id shrinks purely from not being observed — a form of regularization nobody asked for and one that interacts badly with long-tail catalogues. Restricting the update to touched rows fixes that, but then the optimizer state for a row is stale by however many steps have passed since it was last seen, so a momentum buffer or an Adam second moment reflects an old part of training. Neither choice is free; the point is to know which one you are making. ### Frequency imbalance The sum has a second consequence: a head id that appears 40 times per batch accumulates roughly 40 times the gradient magnitude of a tail id that appears once. With a shared learning rate, head rows move much faster and grow larger norms, tail rows barely move. Teams sometimes divide a row's gradient by its occurrence count in the batch to equalise this. That is a deliberate reweighting of the objective, not the gradient of the loss you wrote down, and it should be described as such — it is closer to changing the loss to a per-id mean than to fixing a bug. ### Determinism When many positions scatter into the same row concurrently, the additions happen in a nondeterministic order. Floating-point addition is not associative, so two runs with identical data and seeds can differ in the last bits of a popular row's gradient, and that difference amplifies over training. If bitwise reproducibility matters — for a debugging bisect or a regulated audit trail — the accumulation has to be forced into a deterministic order, usually by sorting the indices first, which costs throughput. ### A note on reserved rows Tables usually reserve a row for padding or for an unknown id. That row is referenced constantly, so its scatter-add accumulates an enormous gradient and it drifts into becoming a very well-trained, very meaningless vector. The usual remedy is to force its gradient to zero so it stays fixed at its initial value.

  • What breaks if the backward assigns into the row instead of accumulating?
    Only the last occurrence's gradient survives, so a row referenced forty times is trained on one fortieth of its evidence. The bug is silent — no crash, plausible magnitudes, loss still falling — and it hurts the most frequent ids hardest. A gradient check on a batch with a duplicated id exposes it immediately.
  • Why is an embedding gradient usually kept in a sparse form?
    Only the batch's unique ids have nonzero rows. On a two-million-row table a dense buffer allocates, zeroes and communicates the entire table every step to carry a few thousand nonzero rows. Storing an index list plus a compact value block is far cheaper in memory and in cross-device reduction traffic.
  • How does applying weight decay densely interact with a sparse embedding gradient?
    It shrinks every row on every step, including the millions never referenced, so rare ids decay toward zero purely from not being observed. Restricting decay to touched rows removes that, but it changes the regularizer into something that depends on id frequency, and it leaves optimizer state stale between a row's appearances.
  • Why might two identical training runs produce slightly different embedding gradients?
    Many batch positions scatter into the same popular row concurrently, and floating-point addition is not associative, so a nondeterministic accumulation order changes the low-order bits. Forcing determinism usually means sorting indices before accumulating, which costs throughput.

saying these in an interview costs you the question

  • Says the row keeps the last occurrence's gradient, overwriting earlier ones
  • Thinks every row of the table receives a nonzero gradient each step
  • Claims averaging over occurrences is the true gradient of the loss
  • Ignores that a dense gradient buffer costs the whole table per step
  • Assumes concurrent accumulation into one row is bitwise deterministic
  • Forgets that dense weight decay shrinks rows the batch never touched

context