GRU (torch gate order)
Overview
Section titled “Overview”| Module | L3.3 · build · Python · Pass 4 · 3 to 4 h, plus your graded tests (rung R3) |
| You build | python/tinyllm/rnn/gru.py: gru_cell, GRUCell, GRU (stacked layers, inter-layer dropout, lengths for padded batches); and your own tests in python/tests/l3-3-gru/, written first |
| Contract | course/contracts/py/tinyllm/rnn/gru.pyi |
| Tests | course/tests/L3.3/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L3.3/gru_torch.npz; your tests are graded by mutation, threshold 0.70 plus one required fault, with a red-then-green journal |
| Needs | L0.1 Tensor · L0.2 the op library · L0.4 Module and Dropout · M07.3 xavier_uniform · M03.3 orthogonal_init · M06.3 the default PCG32 stream · reading: L3.2 the LSTM (same layout, same masking), M01.3, craft.03 (or --ref-deps) |
| Used by | L3.4 two GRUs make a bidirectional layer · later: L3.6 cell='gru', L4.1 the default encoder and decoder cell |
| Milestone | MS-L3 (the GRU language model reaches the calibrated threshold) |
| Optional depth | Cho et al., “Learning Phrase Representations using RNN Encoder-Decoder for Statistical Machine Translation” (2014); Chung et al., “Empirical Evaluation of Gated Recurrent Neural Networks on Sequence Modeling” (2014); the cuDNN developer guide, RNN formulas |
Key Takeaways
Section titled “Key Takeaways”- A GRU has one state and two gates: the update gate interpolates, , so copies the state and its gradient unchanged (
test_update_gate_near_one_keeps_the_state). - The reset gate multiplies the recurrent part after its bias, : torch’s and cuDNN’s form, not Cho’s original, and weights trained one way do not load the other (
test_hand_example_gru_cell,test_matches_torch_cell). - Gate blocks are stacked with torch’s names, so torch GRUs load by name and give torch’s outputs and gradients, padded batches included (
test_matches_torch_sequence). - Three gate blocks instead of four: three quarters of the LSTM’s parameters for a similar memory, which is why
L4.1defaults to it.
How to work this chapter
Section titled “How to work this chapter”ol start L3.3 # stubs gru.py; prints your test path and rung (R3)ol tests L3.3 # the course tests# write ONE test in python/tests/l3-3-gru/, then:ol tdd red L3.3 # must FAIL against your current code# make it pass, then:ol tdd green L3.3 # must PASS with the same test files# repeat for each test; then:ol check L3.3 # course tests, the journal, the mutation grade1. Why now
Section titled “1. Why now”The LSTM (L3.2) solved vanishing gradients with two states and four gates. Cho et al. found in 2014 that one state and two gates do about as well: merge the cell and hidden state, and let a single update gate decide how much of the old state to keep, with the complement going to the new candidate. Your seq2seq models of Part 4 (L4.1) use the GRU by default because it is a quarter smaller and a little faster. It also brings the course’s most common weight-loading bug: the reset gate’s position. Papers write it one way, cuDNN and torch compute it another, and a checkpoint only works with the form it was trained in. This module builds torch’s form and proves it against torch.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| , | input and state | float32[B, D], float32[B, H] |
| , | input and recurrent weights, three gate blocks stacked | [3H, D], [3H, H] |
| , | the two biases | [3H] each |
| the input part of all three pre-activations | [B, 3H] | |
| the recurrent part, bias included | [B, 3H] | |
| reset gate, update gate, candidate | [B, H] each | |
| , , | sigmoid, tanh, elementwise product | |
| length of sequence in a padded batch | integer in |
The cell. Split and into blocks of columns, in torch’s order :
decides how much of the old state the candidate may read: makes a function of the input alone, a reset. decides how much of the old state survives: is a convex combination of and , so it stays in without an extra tanh.
Where the reset gate goes. Cho’s paper computes : reset first, then the matmul. cuDNN computes : matmul first (so the recurrent matmul for all three gates is one GEMM before any gate is known), then the reset, which also scales the recurrent bias . torch follows cuDNN. The two forms are different functions with the same parameter shapes, so nothing fails at load time; the outputs are just wrong. The worked example shows the gap: 0.635 against 0.848.
Why it does not vanish. . With the first term is the identity and the others vanish ( and ), so the state and its gradient pass through a step unchanged, the GRU’s version of the LSTM’s forget gate. The test saturates with a large bias and measures over 15 steps.
Layout, initialization, stacking, padding. Everything else follows L3.2 with three blocks instead of four: torch’s names weight_ih_l{k}, weight_hh_l{k}, bias_ih_l{k}, bias_hh_l{k} in that order; per layer, three Xavier input blocks (M07.3) then three orthogonal recurrent blocks (M03.3), both biases 0 (no gate here plays the forget gate’s role at initialization: starts at ); dropout between layers in training mode only; and the lengths mask that freezes a finished sequence’s state and zeroes its outputs.
3. Worked example by hand
Section titled “3. Worked example by hand”One unit, , . Weights: , , (only the candidate’s recurrent part reads ), .
- ; .
- ; .
- .
- .
- (the gates do not read here, so no other terms).
Cho’s form would compute ; swapping the interpolation would give .
This is test_hand_example_gru_cell.
4. The interface
Section titled “4. The interface”def gru_cell(x, h, w_ih, w_hh, b_ih, b_hh) -> Tensor # h'class GRUCell(Module): # weight_ih [3H, D], weight_hh [3H, H], bias_ih, bias_hh [3H] def __init__(self, d_in, d_h, rng=None); def forward(self, x, h=None)class GRU(Module): # weight_ih_l{k}, weight_hh_l{k}, bias_ih_l{k}, bias_hh_l{k} def __init__(self, d_in, d_h, num_layers=1, dropout=0.0, rng=None) def forward(self, x, state=None, lengths=None) -> tuple[Tensor, Tensor] # x [T, B, d_in] -> out [T, B, H], h_n [num_layers, B, H]What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_gru_cell | unit | section 3: , | you and the test agree on the gates |
test_matches_torch_cell | golden | torch.nn.GRUCell outputs and every gradient | torch weights load and run |
test_matches_torch_sequence | golden | one layer with a state, two layers, a packed padded batch | L4.1 loads and trains GRUs |
test_state_dict_names_match_torch | golden | torch’s names, order, and shapes | the safetensors key contract |
test_gradcheck_gru_cell | gradcheck | autograd through the cell against central differences | your backward is right without a torch |
test_update_gate_near_one_keeps_the_state | property | saturated : and | the reason the GRU remembers |
test_init_orthogonal_recurrence_and_zero_bias | unit | zero biases, orthogonal blocks, Xavier bound, seeding | trainable from step 0 |
test_lengths_keep_padding_out | boundary | padded outputs 0; final state equals running alone; no gradient into padding | batches of sentences (L4.1) |
test_dropout_between_layers_only | unit | dropout between layers, in training mode only | evaluation is deterministic |
test_default_state_and_bad_arguments | boundary | no state means zeros; bad shapes, lengths, sizes raise | errors at the call site |
Your graded tests (rung R3)
Section titled “Your graded tests (rung R3)”The given test:
import mathimport numpy as npfrom tinyllm.autograd.tensor import Tensorfrom tinyllm.rnn.gru import gru_cell
def test_hand_example_gru_cell(): """r = 0.5, z = 0.75, n = tanh(0.5 * 1.5): h' = 0.25 n + 0.75 * 0.5.""" w = [Tensor(a, dtype=np.float64) for a in ([[0.0]] * 3, [[0.0], [0.0], [1.0]], [0, math.log(3), 0], [0, 0, 1.0])] h2 = gru_cell(Tensor([[1.0]], dtype=np.float64), np.array([[0.5]]), *w) assert np.allclose(h2.data, [[0.25 * np.tanh(0.75) + 0.375]])Then, one at a time, red then green: the cell against your own numpy formula with random weights and (this is the test that tells the two reset placements apart); the state_dict keys; orthogonal blocks; a padded batch; a 2-layer GRU against stacking your numpy cell by hand from a nonzero ; dropout ignored by a 1-layer GRU; bad arguments and the default state. Import only tinyllm.rnn.gru and tinyllm.autograd.tensor. The required fault is pitfall 2 (Cho’s placement).
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. the gates in the order | trains from scratch; torch weights give garbage | test_hand_example_gru_cell, test_matches_torch_cell (mutant s01) |
| 2. the reset gate in Cho’s position, or applied to but not to | close to torch, never equal; a loaded checkpoint loses accuracy for no visible reason | test_hand_example_gru_cell, test_matches_torch_cell (mutants s02, s04) |
| 3. the interpolation swapped, | the update gate’s meaning inverted: a torch checkpoint forgets what it should keep | test_hand_example_gru_cell, test_update_gate_near_one_keeps_the_state (mutant s03) |
| 4. padding leaking into the state or the outputs | a short sentence’s encoding depends on its batch-mates | test_lengths_keep_padding_out (mutants s07, s08) |
| 5. a sigmoid candidate | cannot be negative; the state drifts to positive values | test_hand_example_gru_cell (mutant s05) |
| 6. stacking mistakes: dropout after the last layer, layer starting from layer ‘s state | 1-layer train and eval differ; 2-layer outputs differ from torch | test_dropout_between_layers_only, test_matches_torch_sequence (mutants s09, s11) |
7. part of the step computed on raw arrays (h.data), so carries no gradient | the forward is right and training is quietly worse: the state’s own path back through time is cut | test_gradcheck_gru_cell, test_hand_example_gru_cell (mutant s12) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L0.1 | parameters are Tensors; x[t] and h0[k] are getitem |
| Back | L0.2 | F.matmul, F.transpose, F.sigmoid, F.tanh, F.where, F.stack |
| Back | L0.4 | Module registration order is the key order; Dropout between layers |
| Back | M07.3 | xavier_uniform for the input blocks |
| Back | M03.3 | orthogonal_init for the recurrent blocks |
| Back | M06.3 | PCG32(0).substream("init") when no rng is given |
| Forward | L3.4 | bidirectional(GRU, GRU, x, lengths), the encoder of L4.1 |
| Forward | L3.6 | RNNLM(cell='gru') at MS-L3 |
| Forward | L4.1 | the default encoder and decoder cell of the seq2seq model |
If you skip this module, ol check L3.4 stops with L3.4 needs L3.3: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
gru_cell | cuDNN’s GRU (CUDNN_GRU) | the “linear before reset” form you implemented, chosen so the recurrent GEMM runs before the gates | NVIDIA cuDNN API reference, cudnnRNNMode_t |
GRU | torch.nn.GRU | packed sequences, bidirectional layers, projection sizes | torch/nn/modules/rnn.py |
| the original form | Keras GRU(reset_after=False) | Cho’s placement, kept for old checkpoints; reset_after=True is the cuDNN form | keras/layers/rnn/gru.py |
| gating | minGRU (2024) | gates that read only the input, so the recurrence becomes a parallel scan | Feng et al., “Were RNNs All We Needed?” |