Why does an RNN's hidden-state chain block parallel training across time steps?
answer
- critical path, not total work
- step t needs step t-1
- only the batch dimension is wide
- small matmuls starve the matrix units
- sequence length sets serial depth
basics
~20 sEach hidden state is computed from the previous one, so step t cannot start until step t-1 finishes. Training a length-T sequence is T dependent matrix multiplies in a row, and no amount of hardware shortens that chain.
solid answer
~50 sA recurrent layer computes `h_t = f(W_h h_(t-1) + W_x x_t + b)`, so the state at position t is an input to position t+1. That data dependency makes the forward pass a chain of T operations that must run in order, and backpropagation through time is the same chain walked in reverse. The only things you can spread across hardware are the batch and the units inside a single step, and each step's matrix multiply is small, so the matrix units sit mostly idle while the run is bound by launch overhead and memory traffic. Doubling the sequence length doubles the critical path no matter how many accelerators you add. A layer that computes every position directly from the inputs has no such chain: all T positions collapse into one large matrix multiply, the hardware saturates, and the same corpus at the same token budget finishes in a fraction of the wall-clock time.
go deeper
Be ready to state the recurrence relation and say in one sentence why step t waits for step t-1. Knowing that the batch is what keeps the hardware busy is enough at this level.
An interviewer expects the mechanics: the critical path grows with sequence length, backpropagation through time reverses the same chain, per-step matrix multiplies are too small to saturate matrix units, and only the batch and within-step dimensions are wide.
Show you have watched the profile: low utilisation, memory bound by stored activations across T steps, batch and length competing for the same budget, and adding accelerators failing to move wall-clock. Quantify the trade you would expect from a parallel-over-positions rewrite.
Own the framing that this is a hardware-fit argument, not a modelling one. Be prepared to argue when the wall-clock win justifies retraining an established recurrent pipeline, and what the replacement costs in arithmetic and in properties recurrence supplied for free.
## The dependency, stated precisely A recurrent layer is defined by a single equation applied over and over with the same weights: `h_t = f(W_h h_(t-1) + W_x x_t + b)`. Here `x_t` is the input at position t, `h_t` is the hidden state after position t, `W_h` and `W_x` are weight matrices shared across all positions, and `f` is an elementwise nonlinearity. Gated cells such as LSTM and GRU add more terms, but the shape is identical: the output at position t is a function of the state at position t-1. That is a *data dependency*, not a code-organisation problem. You cannot compute `h_5` without `h_4`, and you cannot compute `h_4` without `h_3`. For a sequence of length T, the forward pass therefore contains T operations that must execute one after another. Backpropagation through time (BPTT) unrolls the same chain and walks it backwards, so the backward pass is another T serial operations. The *critical path* through the computation grows linearly with sequence length. ## What is still parallel, and why it is not enough Three things do parallelize inside a recurrent trainer: - **The batch.** Independent sequences share no hidden state, so B sequences can advance through step t together. This is the only large source of parallelism available. - **Within a step.** The matrix multiply `W_h h_(t-1)` is itself parallel across rows and columns. - **Across stacked layers,** partially: layer 2 at position t needs layer 1 at position t, so a wavefront can overlap layers, but the depth of a stack is small compared with T. The problem is that these are small. At each step you multiply a `B x d` activation matrix by a `d x d` weight matrix, where d is the hidden size — a few hundred to a couple of thousand. Modern accelerators are built for very large matrix multiplies; a stack of small ones is dominated by fixed per-operation overhead and by moving weights and activations through memory rather than by arithmetic. Utilisation is low, and the standard fix — enlarge the batch — has limits: memory holds the saved activations for *all* T steps for the backward pass, so batch size and sequence length trade against each other, and very large batches change optimisation behaviour. The decisive point is the critical path. Adding accelerators shortens the time to process a batch, but it cannot make step 4 start before step 3 ends. A multi-day translation run on eight accelerators is slow not because the arithmetic is large but because the arithmetic is *strung out*: the machine spends most of its time waiting on the previous step. ## What dropping recurrence buys Replace the recurrent layer with one that computes every position's output directly from the layer's inputs — every position's result depends on the inputs, never on another position's *output*. Now the whole sequence is one operation. A batch of B sequences of length T becomes a `(B*T) x d` matrix, multiplied once. The critical path per layer no longer depends on T at all; it depends only on the depth of the network. Sequence length turns from a serial dimension into another parallel dimension, alongside the batch. The consequence is wall-clock, and it is large. On the same corpus with the same number of training tokens, a trainer that is parallel over positions saturates the matrix units and finishes an epoch in a small fraction of the time a recurrent unroll takes on the same hardware. That single fact — not a claim about better modelling — is what motivated moving away from recurrence for large-scale sequence training. ## The distinction candidates get wrong Removing recurrence does not necessarily reduce total arithmetic. A parallel-over-positions layer often performs *more* floating-point work than a recurrent unroll on the same sequence, because each position is computed against many others rather than against a single carried state. What changes is that the work is issued all at once instead of in a chain, so it maps onto hardware that is designed for exactly that shape. Wall-clock time and total FLOPs are different quantities, and this whole trade is about the first. ## How to talk about it in an interview Say the equation, name the dependency, name the critical path, and name what stays parallel (the batch). Then give the consequence: low utilisation, wall-clock scaling with T, and no rescue from more hardware. Finish with the trade — the parallel replacement pays in arithmetic and in what it loses (ordering has to be supplied, and per-step cost stops being bounded), which shows you understand it as an engineering exchange rather than a strict upgrade.
- If the batch dimension is parallel, why not simply train with a far larger batch?Because backpropagation through time keeps the saved activations for every one of the T steps, so memory is roughly proportional to batch size times sequence length; batch and length compete for the same budget. Very large batches also change the optimisation problem — fewer updates per epoch, and a learning rate that has to be retuned. Batching widens the machine but never shortens the chain.
- Does removing recurrence reduce the total arithmetic of a training run?Usually not, and often the opposite. A parallel-over-positions layer computes each position against many others rather than against one carried state, so floating-point work can rise. What falls is wall-clock time, because the work is issued as a few large matrix multiplies instead of thousands of dependent small ones. Confusing FLOPs with elapsed time is the classic error here.
- Where does the backward pass fit into this picture?It is the same chain reversed. Backpropagation through time must visit position T before T-1, so the backward critical path is also length T, and it needs the stored forward activations for every step, making memory grow with sequence length. Truncating the unroll caps both, but it shortens the gradient horizon rather than making the forward pass parallel.
A hundred cooks cannot shorten a recipe whose every step needs the previous step's pot. Extra hands only help if you are cooking a hundred separate dinners at once.
saying these in an interview costs you the question
- Claims accelerators already parallelize an RNN over time steps
- Says the problem is parameter count, not the data dependency
- Thinks a larger batch removes the sequential chain
- Confuses total floating-point work with wall-clock time
- Believes truncated backpropagation makes the forward pass parallel