Personalize this lesson
Adapt explanations and teaching visuals to your background and preferred voice.
An on-call assistant sees error logged, then reproduction confirmed, then rollback requested. Reverse those events and the latest evidence changes: an error now arrives after the rollback request. Neither history proves the incident is resolved, but their order matters for the next decision. How can a model track that order without requiring a different architecture for every possible history length?
The previous lesson converted raw class scores, or logits, into a softmax probability distribution. We'll keep that classifier at the tail end, outputting decisions across three actions: ignore, rollback, and escalate. Our job here is to transform a variable-length stream of events into the single, fixed-size vector that the downstream classifier expects.
Order is part of the input
If you shuffle raw event vectors, a bag-of-events sum or average treats the history as identical because those operations are commutative. An order-aware model can distinguish the histories. This doesn't guarantee that every permutation has a different state: compressed representations can collide, and a model may deliberately ignore order that isn't useful for its task.
The tiny accumulator below isn't a trained neural network. Each incoming event updates a running state influenced by earlier events, with older evidence discounted.
Before running it, predict which order leaves the larger final state. The same three values arrive, but the recurrence discounts earlier state by 0.6, so the event in the final position exerts the strongest direct influence:
1signals = {
2 "error logged": 1.0,
3 "reproduction confirmed": 0.4,
4 "rollback requested": 0.9,
5}
6
7def running_state(events):
8 state = 0.0
9 for event in events:
10 state = 0.6 * state + signals[event]
11 return state
12
13incident_order = ["error logged", "reproduction confirmed", "rollback requested"]
14reversed_order = list(reversed(incident_order))
15
16print("incident order:", round(running_state(incident_order), 3))
17print("reversed order:", round(running_state(reversed_order), 3))
18print("same events:", sorted(incident_order) == sorted(reversed_order))1incident order: 1.5
2reversed order: 1.564
3same events: TrueThe incident order computes 0.6 * (0.6 * 1.0 + 0.4) + 0.9 = 1.5. Reversing it puts the largest event 1.0 in the final position, producing 0.6 * (0.6 * 0.9 + 0.4) + 1.0 = 1.564. A simple sum would yield 2.3 in both cases; the recurrent update makes sequence order inseparable from the final value.
One cell, reused at every step
A recurrent neural network (RNN) makes this update rule trainable. Its hidden state () is a running vector of numbers carried from one event to the next, serving as working memory. The update rule is called a recurrent cell, and each new event advances the sequence by one timestep.
Start with a scalar hidden state. Multiply the incoming event by 0.8, multiply the previous hidden state by 0.5, add them, and pass the result through hyperbolic tangent (tanh). Like sigmoid, tanh is a smooth S-shaped activation function, but its output spans and is zero-centered. With an initial state and a first event of 1.0, the weighted sum is 0.8, giving .
For the second event (0.4), combine 0.8 * 0.4 with 0.5 * 0.664. That sum is about 0.652, yielding . Run all three steps with these demonstration weights:
1import math
2
3events = [1.0, 0.4, 0.9]
4w_input, w_state, bias = 0.8, 0.5, 0.0
5state = 0.0
6
7for step, event in enumerate(events, start=1):
8 state = math.tanh(w_input * event + w_state * state + bias)
9 print(f"h_{step} = {state:.3f}")1h_1 = 0.664
2h_2 = 0.573
3h_3 = 0.764Three scalar inputs condense into one final scalar summary. When we expand the hidden state into a vector, different coordinates can track distinct patterns in parallel. The fundamental vector RNN recurrence equation is:
Here is the current input vector, is the prior hidden state, and is a bias vector. The weight matrices and project the new observation and transition the prior memory, respectively. Tanh applies element-wise to squash each hidden coordinate into .
Notice the core architectural insight: the exact same parameter matrices , , and bias execute at step 1, step 2, and step 1,000. Parameters are tied across time rather than allocated anew for every token.[1]
Why can an RNN read either 3 events or 30 events without changing its parameter count?
Answer
It reuses the exact same input and recurrent weight matrices at every timestep. Increasing sequence length increases how many times the cell executes, not the number of learned weights.
Trace a vector RNN
Represent each incident event with three binary features:
| Event | error_seen | repro_confirmed | rollback_requested |
|---|---|---|---|
| Error logged | 1 | 0 | 0 |
| Reproduction confirmed | 0 | 1 | 0 |
| Rollback requested | 0 | 0 | 1 |
These are one-hot event vectors: exactly one coordinate is 1 and all others are 0. They describe the incoming event at the current timestep, not cumulative flags.
Set up a two-dimensional hidden state () and an initial state , with biases set to zero for clean manual verification. At the first event , the matrix-vector multiplication extracts the first column of , which is . Since , the recurrent term contributes nothing. The cell computes .
The tensor dimensions dictate how the calculations scale. Let be input feature width and be hidden width:
| Object | This example | General shape | Role |
|---|---|---|---|
(3,) | One incoming event vector | ||
| , | (2,) | Hidden state and bias vectors | |
(2, 3) | Maps current event into hidden space | ||
(2, 2) | Maps prior hidden state to new hidden state |
Each row of and computes one hidden coordinate. Step through all three events in pure Python:
1import math
2
3def matvec(matrix, vector):
4 return [sum(weight * value for weight, value in zip(row, vector, strict=True)) for row in matrix]
5
6def add(left, right):
7 return [a + b for a, b in zip(left, right, strict=True)]
8
9def tanh_vec(values):
10 return [math.tanh(value) for value in values]
11
12events = [
13 [1.0, 0.0, 0.0],
14 [0.0, 1.0, 0.0],
15 [0.0, 0.0, 1.0],
16]
17W_xh = [
18 [0.8, 0.2, 0.1],
19 [0.1, 0.7, 0.4],
20]
21W_hh = [
22 [0.5, 0.1],
23 [0.0, 0.6],
24]
25
26state = [0.0, 0.0]
27for step, event in enumerate(events, start=1):
28 state = tanh_vec(add(matvec(W_xh, event), matvec(W_hh, state)))
29 print(f"h_{step} = {[round(value, 3) for value in state]}")
30
31print("final width:", len(state))1h_1 = [0.664, 0.1]
2h_2 = [0.494, 0.641]
3h_3 = [0.39, 0.655]
4final width: 2![Three unrolled copies of the same RNN cell carry hidden state left to right from h0 equal to [0, 0] to h3 equal to [0.390, 0.655]. Input events error, reproduction, and rollback enter each cell through shared input matrix Wxh. Recurrent matrix Whh carries memory across steps. Output matrix Wout maps final state h3 to the predicted rollback action.](/cdn/content-image/preparation/rnns-lstms-grus-sequence-modeling/illustrations/_generated/unrolled_rnn_dark.png?v=5459dc7ebea1)
Attach a classification head to that final hidden state: a linear projection followed by softmax. This reuses the multiclass head from the previous lesson, with weight matrix mapping the 2D summary into 3 action logits.
For row 1 (ignore), the pre-activation logit is . Rows 2 and 3 produce (rollback) and (escalate). Because is the largest logit, we know rollback will claim the highest softmax probability before running math.exp:
1import math
2
3final_state = [0.390, 0.655]
4W_out = [
5 [0.7, -0.3], # ignore
6 [-0.4, 0.8], # rollback
7 [0.1, 0.2], # escalate
8]
9labels = ["ignore", "rollback", "escalate"]
10
11logits = [
12 sum(weight * value for weight, value in zip(row, final_state, strict=True))
13 for row in W_out
14]
15shift = max(logits)
16weights = [math.exp(logit - shift) for logit in logits]
17total = sum(weights)
18probabilities = [weight / total for weight in weights]
19
20print("logits:", [round(logit, 3) for logit in logits])
21print("probabilities:", [round(probability, 3) for probability in probabilities])
22print("prediction:", labels[probabilities.index(max(probabilities))])1logits: [0.076, 0.368, 0.17]
2probabilities: [0.291, 0.389, 0.32]
3prediction: rollbackWith these hand-chosen weights and the rounded final state, the head selects rollback with 38.9% probability. This isn't evidence of a trained incident classifier's accuracy. The head receives only , rather than the full history or a separate length input; the state may still encode information about length.
Count the parameters in a trainable version with one bias vector per affine block: has 6 weights, has 4, and has 2 biases, giving 12 recurrent parameters. A head with 6 weights and 3 trainable biases adds 9, totaling 21. The manual trace used zero biases. Changing sequence length doesn't change this count:
1input_size, hidden_size, output_size = 3, 2, 3
2rnn_parameters = hidden_size * input_size + hidden_size * hidden_size + hidden_size
3head_parameters = output_size * hidden_size + output_size
4
5for timesteps in [3, 30, 300]:
6 print(f"{timesteps:>3} events -> {rnn_parameters + head_parameters} parameters")13 events -> 21 parameters
2 30 events -> 21 parameters
3300 events -> 21 parametersThe final-state interface is a compression bottleneck: useful evidence from the history must be encoded in a fixed-width vector, and some details can be lost. During training, credit assignment to distant events travels backward through the recurrent channel.
Backpropagation through time: why a distant signal can fade
The forward pass assigned its highest softmax probability to rollback. During training, a target label and cross-entropy loss measure the error of that prediction. Backpropagation through time (BPTT) applies the chain rule backward across the unrolled sequence to compute gradients for the shared weight matrices.[1]
Because acts at every timestep, its gradient accumulates local contributions across time. For a summed loss , we can write:
Here is the immediate local derivative at timestep with previous state held constant. The key term governing credit assignment over long sequences is the temporal chain:
To see why this chain causes instability, look at the one-dimensional cell we ran earlier: . The local derivative from one state to the next is:
Because , its mathematical derivative lies in for finite inputs. With recurrent weight , the local factor is at most , attaining that bound at . Trace the product in our example:
1import math
2
3events = [1.0, 0.4, 0.9]
4w_input, w_state = 0.8, 0.5
5state = 0.0
6factor = 1.0
7
8for step, event in enumerate(events, start=1):
9 state = math.tanh(w_input * event + w_state * state)
10 local = (1.0 - state * state) * w_state
11 factor *= local
12 print(f"step {step}: h={state:.3f} local factor={local:.6f} product so far={factor:.6f}")1step 1: h=0.664 local factor=0.279528 product so far=0.279528
2step 2: h=0.573 local factor=0.335820 product so far=0.093871
3step 3: h=0.764 local factor=0.207910 product so far=0.019517The temporal derivative is . If a final loss supplies derivative , the gradient that reaches is only 0.0195 * g. That's less than 2% of the learning signal surviving after just three recurrent steps.
When every step compounds a factor of , the shrinkage is exponentially steep:
1for steps in [1, 4, 8, 12]:
2 gradient_factor = 0.6 ** steps
3 print(f"{steps:>2} recurrent steps: {gradient_factor:.6f}")11 recurrent steps: 0.600000
2 4 recurrent steps: 0.129600
3 8 recurrent steps: 0.016796
412 recurrent steps: 0.002177A matrix bound, and what it actually proves
In the general multivariate case, the one-step derivative is a Jacobian matrix. Because :
Backpropagating an error gradient column vector from timestep back to timestep multiplies by transposed Jacobians in reverse chronological order:[2][1]
This expression assumes a loss only at the final state. When intermediate positions have losses, their gradient contributions must also be added. Matrix order matters: apply first, then work backward. Using the induced Euclidean norm and its submultiplicative property gives:
where , and is the maximum singular value (spectral norm) of the recurrent matrix.
If one uniform bound satisfies across the span, the temporal Jacobian contracts exponentially. This is a sufficient condition for a vanishing path. The total parameter gradient can still include strong contributions from recent positions.
If the bound exceeds one, growth is possible, but the upper bound doesn't prove explosion. Direction, successive activations, and alignment matter. For example, a matrix can have a singular value above one while its repeated powers eventually decay. Spectral-radius arguments for constant linear recurrences don't replace the Jacobian-product analysis of a nonlinear, input-dependent RNN.[1]
For a linear recurrence, take W = [[0.5, 1], [0, 0.5]]. Its largest singular value is about 1.21. Does that prove distant gradients grow without bound?
Answer
No. Here Wⁿ = [[0.5ⁿ, n·0.5ⁿ⁻¹], [0, 0.5ⁿ]], which tends to zero. Some directions can grow temporarily, but repeated multiplication eventually contracts. The bound ‖Wⁿ‖ ≤ 1.21ⁿ becomes loose; an expanding upper bound isn't evidence of an expanding actual gradient.
There is no universal 20- or 50-step float32 cutoff. For our illustrative factor, is still representable in float32. A small gradient may have little optimization effect before it underflows; whether a parameter update rounds away also depends on the parameter's scale and learning rate. In an expanding path, large gradients can instead destabilize updates or become nonfinite.

Truncation and gradient clipping
Two operational techniques manage these dynamics during training:
Truncated BPTT limits how far backward the computational graph extends. Training can split a long sequence into chunks, carrying state values forward while h = h.detach() cuts the graph at each boundary. An LSTM needs both h and c detached. This reduces graph memory but also removes credit assignment across the boundary; the forward state can retain evidence the current loss can no longer train earlier steps to store.
Gradient norm clipping addresses exploding gradients. If the norm of the gradient vector exceeds a threshold , we rescale the vector:
1import math
2
3raw_x = 1.4 ** 12
4raw_y = -0.8 * raw_x
5raw_norm = math.hypot(raw_x, raw_y)
6max_norm = 5.0
7scale = min(1.0, max_norm / raw_norm)
8clipped = (raw_x * scale, raw_y * scale)
9
10print("raw norm:", round(raw_norm, 3))
11print("clipped norm:", round(math.hypot(*clipped), 3))
12print("clipped gradient:", [round(value, 3) for value in clipped])1raw norm: 72.604
2clipped norm: 5.0
3clipped gradient: [3.904, -3.123]For a finite, nonzero gradient, norm clipping rescales its magnitude while preserving direction in exact arithmetic.[1] It bounds the gradient, rather than guaranteeing a safe optimizer update. For plain SGD without momentum or weight decay, the step norm is at most ; adaptive optimizers and momentum change that relationship. With , leave it unchanged rather than evaluating . Nonfinite gradients need detection and diagnosis, not just rescaling.
Clipping doesn't restore a small or vanished signal. Gates offer a different memory path; initialization, regularization, training objectives, and dependency length can also affect gradient flow. None is a guarantee that all long-range dependencies will be learned.
LSTM: Additive cell state and gated memory
Gradient clipping can't create long-term memory. LSTMs provide a path that doesn't force the stored state through a saturating activation at every update.
Long short-term memory (LSTM) addresses this difficulty with two related state vectors:[3]
- Internal cell state (): An internal gradient superhighway or conveyor belt that updates additively.
- Working hidden state (): The exposed memory vector passed to the next layer and the output classifier.
Information flow is controlled by gates: learned, coordinate-wise affine projections squashed by the sigmoid function (). A gate value near 0 acts as a closed valve; a value near 1 leaves the valve fully open. The standard modern formulation includes the forget gate introduced by Gers et al.[4]
Suppose a single coordinate in our cell state holds 0.855 after an error event. A routine health check arrives. We want to retain that earlier error flag while recording almost nothing new from the routine check:
| Control | Value | Operational role | Calculation |
|---|---|---|---|
| Forget gate () | 0.98 | Fraction of old cell memory to keep | Retain 0.98 * 0.855 = 0.8379 |
| Input gate () | 0.02 | Fraction of new candidate to write | Write 0.02 * 0.10 = 0.0020 |
| Cell update () | Additive combination of both paths | ||
| Output gate () | 0.80 | Fraction of filtered cell to expose | Expose |
The stored value () differs from the exposed state (). Closing the output gate can hide a stored coordinate without immediately erasing it. These states are still coupled: the hidden state depends on the cell, and future gates depend on the hidden state. The values below are hand-chosen controls, not learned gates.
Step through three illustrative transitions in pure Python:
1import math
2
3cell = 0.0
4steps = [
5 ("error recorded", 0.00, 0.95, 0.90, 0.80),
6 ("routine check", 0.98, 0.02, 0.10, 0.80),
7 ("rollback verified", 0.97, 0.15, 0.50, 0.90),
8]
9
10for label, forget, write, candidate, expose in steps:
11 cell = forget * cell + write * candidate
12 hidden = expose * math.tanh(cell)
13 print(f"{label:<20} c={cell:.3f} h={hidden:.3f}")1error recorded c=0.855 h=0.555
2routine check c=0.840 h=0.549
3rollback verified c=0.890 h=0.640The full LSTM equations
In a trainable LSTM, learned affine projections compute these gate vectors from the concatenation of the prior working state and the current input (denoted ):
Here denotes the element-wise Hadamard product.
Follow the direct cell-state gradient path
In this standard LSTM without peephole connections, the cell receives as separate state inputs. Holding and fixed also holds the gates and candidate fixed. The direct cell-state Jacobian is:
The contribution from following only this direct path across several steps is:
Contrast this directly with the vanilla RNN product :
- No repeated matrix powers: There is no factor acting along this diagonal channel.
- No saturating derivative: The cell state itself doesn't pass through during the update.
- Controllable retention: Forget gates near one slow decay. If every gate on a coordinate is 0.999 for 100 transitions, its direct-path factor is about 0.905, retaining 90.5%. At 0.5, the same product is approximately .
The original LSTM's unit-weight cell self-connection was called the constant error carousel.[3] A forget-gated LSTM learns retention instead of always copying the cell.[4] The full gradient also travels through hidden states and gate dependencies; the output readout includes , which can be small. LSTMs mitigate vanishing gradients rather than eliminating them, and they can still have exploding gradients.
What does an LSTM forget gate value near 1 do?
Answer
It passes almost the entire corresponding cell-state coordinate into the next timestep. Paired with an input gate near 0, it shields that stored feature from decay and overwriting across the current step.
GRU: Unified hidden state and streamlined gating
LSTMs maintain two state vectors ( and ) and four distinct affine transformations per step. The gated recurrent unit (GRU) streamlines this architecture by collapsing the cell state and hidden state into a single vector , and replacing the separate forget and input gates with a single update gate ().[5]
Instead of independent keep and write decisions, the update gate forms a convex combination:
For an old state of 0.80 and a candidate of -0.20, retaining 95% () automatically sets the write weight to 5%: 0.95 * 0.80 + 0.05 * (-0.20) = 0.750. The two mixture fractions always sum to 1.
The second gate is the reset gate (), which decides how much of the prior hidden state is visible when computing the proposed candidate :
When , the candidate calculation ignores past history and acts like a standard feedforward projection of . When , the unit ignores the candidate and carries the past state forward unchanged.
Test the convex blend across three update gate settings:
1old_state = 0.80
2candidate = -0.20
3
4for keep_gate in [0.95, 0.50, 0.05]:
5 new_state = keep_gate * old_state + (1.0 - keep_gate) * candidate
6 print(f"z={keep_gate:.2f} -> h={new_state:.3f}")1z=0.95 -> h=0.750
2z=0.50 -> h=0.300
3z=0.05 -> h=-0.150
Reset gate placement: Cho et al. versus PyTorch
Here explicitly means the fraction kept from the old state, matching Cho et al. and PyTorch. Some descriptions give the complementary fraction the name "update gate," so compare equations rather than names. In Cho et al.'s candidate, the recurrent portion is : reset scales the old state before its recurrent matrix multiplication.[5] is the recurrent block of the concatenated matrix above.
PyTorch's nn.GRU applies the reset gate after the linear transformation:[6]
PyTorch documents this deliberate change as an efficiency choice: the recurrent projections can be computed together before applying reset. It isn't a promise about a particular GPU kernel count.
Matrix multiplication and coordinate-wise scaling generally don't commute: . Mixing coordinates can therefore change the output, although equal gates or particular inputs can make the two agree:
1import numpy as np
2
3old = np.array([0.8, -0.2])
4reset = np.array([0.1, 0.9])
5recurrent = np.array([[1.0, 2.0], [3.0, 4.0]])
6
7before = recurrent @ (reset * old)
8after = reset * (recurrent @ old)
9print("reset before matrix:", np.round(before, 3))
10print("reset after matrix: ", np.round(after, 3))
11print("candidate before: ", np.round(np.tanh(before), 3))
12print("candidate after: ", np.round(np.tanh(after), 3))
13assert not np.allclose(before, after)1reset before matrix: [-0.28 -0.48]
2reset after matrix: [0.04 1.44]
3candidate before: [-0.273 -0.446]
4candidate after: [0.04 0.894]The two candidates differ significantly (-0.273 versus 0.040). When porting weights between frameworks, always check the exact candidate formulation.
Parameter auditing and tensor geometry
Every gate and candidate vector in an RNN cell emits values. With concatenated input of width , each affine transformation requires a weight matrix of shape and a bias vector of shape .
Counting parameters per cell reveals the structural scaling:
| Cell | Affine blocks | General parameter formula | Our example () |
|---|---|---|---|
| Vanilla RNN | 1 () | ||
| GRU | 3 (, , candidate) | ||
| LSTM | 4 (, , candidate, ) |
For these one-layer cells with equal widths and one bias per block, GRU has three quarters as many parameters as LSTM. Its three affine blocks also require less dense projection work than four equal-sized blocks. Wall-clock speed still depends on batching, kernels, hardware, and other operations.
Why does PyTorch's nn.GRU report 42 parameters instead of 36? In production implementations, PyTorch stores separate input-hidden () and hidden-hidden () bias vectors for each gate:[6]
The extra biases don't change hidden width. They aren't all merely redundant storage: PyTorch's candidate recurrent bias is multiplied by reset, , so it generally can't be merged into a fixed input bias. Porting a textbook GRU requires checking both reset placement and bias placement.
For a one-layer, unidirectional nn.GRU(input_size=3, hidden_size=2, batch_first=True), an input batch tensor of shape (B, T, 3) outputs a tensor of shape (B, T, 2) and a final hidden state of shape (1, B, 2). Here counts independent sequences and counts timesteps. The leading dimension in the final hidden state represents , not sequence position.
Why can the textbook parameter formula and framework parameter count differ for the same GRU?
Answer
Our textbook formula counts one bias per block. PyTorch stores input-hidden and hidden-hidden biases for each block, adding parameters per block. In the candidate, the hidden-hidden bias sits inside the reset multiplication, so checking the count alone doesn't establish equivalent dynamics.
Sequence padding and batch masking
In production workloads, sequences within a batch rarely have the same length. Frameworks pad shorter sequences with dummy values (typically zeros) into a uniform rectangular tensor.
A critical failure mode occurs when developers run an unmasked recurrent loop over padded rows. In an RNN, padding is still processed as an event. Even an all-zero input row triggers another state update:
This generally differs from , though a fixed-point state can remain unchanged. A zero row doesn't universally rotate or shrink memory; its effect depends on the recurrence. Verify the change for our chosen weights:
1import math
2
3def matvec(matrix, vector):
4 return [sum(weight * value for weight, value in zip(row, vector, strict=True)) for row in matrix]
5
6def add(left, right):
7 return [a + b for a, b in zip(left, right, strict=True)]
8
9def tanh_vec(values):
10 return [math.tanh(value) for value in values]
11
12W_xh = [
13 [0.8, 0.2, 0.1],
14 [0.1, 0.7, 0.4],
15]
16W_hh = [
17 [0.5, 0.1],
18 [0.0, 0.6],
19]
20
21def run(events):
22 state = [0.0, 0.0]
23 for event in events:
24 state = tanh_vec(add(matvec(W_xh, event), matvec(W_hh, state)))
25 return [round(value, 3) for value in state]
26
27two_events = [
28 [1.0, 0.0, 0.0],
29 [0.0, 1.0, 0.0],
30]
31print("two real events: ", run(two_events))
32print("plus a zero pad row: ", run(two_events + [[0.0, 0.0, 0.0]]))
33print("plus a junk pad row: ", run(two_events + [[9.0, -9.0, 9.0]]))1two real events: [0.494, 0.641]
2plus a zero pad row: [0.302, 0.367]
3plus a junk pad row: [1.0, -0.889]The state shifts from [0.494, 0.641] to [0.302, 0.367] on a zero row, and to [1.0, -0.889] on junk values. A classifier reading the post-padding state can therefore change its prediction because of filler rather than real incident data.
PyTorch's pack_padded_sequence packages the valid prefixes and their active batch sizes so the recurrent layer skips padding on CPU or GPU.[7] Lengths must describe right-padded sequences; a tensor of lengths must be on CPU. enforce_sorted=False lets the function sort internally. Each sequence needs a positive length no greater than the available timesteps.
1import torch
2from torch import nn
3from torch.nn.utils.rnn import pack_padded_sequence
4
5torch.manual_seed(7)
6gru = nn.GRU(input_size=3, hidden_size=2, batch_first=True)
7lengths = torch.tensor([3, 2])
8base = torch.tensor([
9 [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]],
10 [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 0.0]],
11])
12changed_padding = base.clone()
13changed_padding[1, 2] = torch.tensor([9.0, -9.0, 9.0])
14
15def packed_final(batch):
16 packed = pack_padded_sequence(
17 batch, lengths.cpu(), batch_first=True, enforce_sorted=False
18 )
19 _, hidden = gru(packed)
20 return hidden[-1]
21
22def unpacked_final(batch):
23 _, hidden = gru(batch)
24 return hidden[-1]
25
26packed_changed = not torch.allclose(packed_final(base), packed_final(changed_padding))
27unpacked_changed = not torch.allclose(unpacked_final(base), unpacked_final(changed_padding))
28
29print("batch shape:", tuple(base.shape))
30print("packed final shape:", tuple(packed_final(base).shape))
31print("packed state changed by padding:", packed_changed)
32print("unpacked state changed by padding:", unpacked_changed)
33output, hidden = gru(base)
34parameter_count = sum(parameter.numel() for parameter in gru.parameters())
35print("unpacked output shape:", tuple(output.shape))
36print("unpacked hidden shape:", tuple(hidden.shape))
37print("GRU parameters:", parameter_count)
38assert not packed_changed and unpacked_changed
39assert parameter_count == 42
40assert output.shape == (2, 3, 2) and hidden.shape == (1, 2, 2)1batch shape: (2, 3, 3)
2packed final shape: (2, 2)
3packed state changed by padding: False
4unpacked state changed by padding: True
5unpacked output shape: (2, 3, 2)
6unpacked hidden shape: (1, 2, 2)
7GRU parameters: 42Packing ensures the second sequence halts at . Its final representation remains completely immune to whatever values occupy the padded slots.
Why can padding corrupt a final-state classifier even if every padded row uses zeros?
Answer
Without masking or packing, the recurrent cell treats the zero row as another valid input event. Applying the recurrent transition and bias can alter the state, so the classifier evaluates a post-padding vector instead of the true final event's state. A fixed point could happen to remain unchanged; zero padding doesn't guarantee that.
Sequence compression and the encoder-decoder bottleneck
The vector RNN we built is an encoder: it takes a variable-length sequence of observations and condenses them into a single fixed-width hidden vector .
In 2014, Cho et al.[5] and Sutskever et al.[8] used recurrent encoders and decoders to model a target sequence conditioned on a source sequence. Their interfaces differ: Cho et al. condition the decoder on a fixed context throughout, while Sutskever et al. use the encoder's state to initialize the decoder. Both communicate the source through a fixed-size representation.
| History consumed | Final context vector | Width |
|---|---|---|
| Error | [0.664, 0.100] | 2 |
| Error, reproduction | [0.494, 0.641] | 2 |
| Error, reproduction, rollback | [0.390, 0.655] | 2 |
While effective for short sentences, this setup encounters a strict information bottleneck: every nuance, entity, and causal dependency across a 100-word paragraph must fit through that single final hidden vector. The decoder has no direct access to earlier tokens; it can only peer through the keyhole of .
The sequential bottleneck, and how attention changes it
Standard nonlinear recurrent cells have a sequential execution bottleneck across timesteps.
For a standard nonlinear RNN, LSTM, or GRU, calculating the new state requires the previous state:
Here for the vanilla RNN or GRU, and for an LSTM.
The state transition at a position must wait for the previous state. A 4,096-token sequence therefore has 4,096 dependent transitions, not necessarily 4,096 GPU kernel launches. Input projections can be computed across positions in advance, batches run in parallel, and backend kernels may fuse work or use persistent algorithms.[9]
Accelerators benefit from large matrix multiplications. With a small batch or hidden width, recurrent transitions may expose too little work to use them efficiently. With a batch, these projections are matrix-matrix operations, not just matrix-vector operations. Utilization and memory traffic depend on shapes and kernels; recurrent weights needn't be fetched from HBM anew for every step.
Self-attention breaks the serial chain
A Transformer replaces these recurrent transitions with self-attention and position-wise transformations.[10] For tokens already present during training:
- Parallel projections: Queries, keys, and values can be projected across the sequence together (often with a combined projection):
- Direct routing: A position can attend directly to any position permitted by the mask: A causal mask has for future positions and zero for allowed positions. Position 4,096 can receive a weighted value from Position 1 within one layer, instead of 4,095 recurrent relays. That is a constant graph path length, not constant computation time.
Attention also needs information about position to distinguish ordering beyond what the mask supplies. The original Transformer adds positional encodings. During autoregressive generation, future tokens aren't available yet, so generating new tokens still has a sequential dependency.
| Property | Recurrent layer (RNN / LSTM / GRU) | Self-attention layer (Transformer) |
|---|---|---|
| Sequential operations during training | (strictly serial across timesteps) | (parallel tensor operations across all tokens) |
| Interaction path length between distant tokens | (information must survive transitions) | (direct pairwise dot product) |
| Compute at comparable input/hidden width | , including projections | |
| Parallel work during training | Across batch items and projection coordinates; transitions remain dependent | Across supplied positions; actual utilization depends on shapes and kernels |
| State available for the next update | Fixed-width , plus for LSTM | Full attention reads explicit representations of all allowed positions |
Fixed-width recurrent state doesn't imply constant training memory. Full BPTT saves intermediate activations across the sequence; truncation or recomputation changes that cost. The table's parallel-operation counts describe a fixed-depth layer's dependency structure, not GPU kernel counts or measured latency.
Full attention's quadratic pairwise work can be expensive for long sequences. FlashAttention avoids materializing the full score matrix in GPU HBM, while still computing full attention rather than making the pairwise arithmetic linear.[11]
Some recurrent architectures support a different execution strategy. Mamba's selective state-space recurrence has input-dependent coefficients that can be prepared before its scan, enabling an associative parallel scan during training.[12] Its recurrent state size and per-token update cost are constant with respect to sequence length for a fixed model. This doesn't make arbitrary nonlinear LSTM or GRU recurrences associative, and a parallel scan still has scheduling and reduction costs.
Diagnosing recurrent failure modes
When debugging sequence models in production, verify these three critical failure points:
In the scalar RNN, start a second, independent incident containing only event 0.4. What happens if you mistakenly reuse the first incident's final state instead of resetting to zero?
Answer
Resetting gives tanh(0.8 * 0.4) ≈ 0.310. Reusing the previous final state (about 0.764) gives tanh(0.32 + 0.5 * 0.764) ≈ 0.606. Detachment preserves values, so it doesn't prevent cross-incident contamination. Reset each independent sequence to its intended initial state (zero in this example, or a learned initialization in another model).
In the LSTM routine-check step, change only the output gate from 0.80 to 0.10. Which state changes, and which stays fixed at this step?
Answer
The cell state remains exactly 0.8399: its additive update depends only on the forget gate, input gate, and candidate state. The exposed working state drops to 0.10 * tanh(0.8399) ≈ 0.069. Future cell updates may shift downstream because future gates receive , but the current cell state update is completely decoupled from the output gate.
In the reset-placement example, replace the recurrent matrix with the 2-by-2 identity matrix. Why do both candidates now agree, and why is that a weak test of GRU compatibility?
Answer
Both paths reduce to reset * old because the identity matrix does not mix coordinates across dimensions and there are no biases. A diagonal matrix always commutes with coordinate-wise scaling. Always test compatibility with a non-diagonal recurrent matrix that mixes coordinates.
When diagnosing training instability, evaluate gradient norms independently from loss values. A model can show stable forward losses while its early-timestep gradients have completely vanished.