Skip to content

LSTM (torch gate order)

ModuleL3.2 · build · Python · Pass 4 · 4 to 5 h, plus your graded tests (rung R3)
You buildpython/tinyllm/rnn/lstm.py: lstm_cell, LSTMCell, LSTM (stacked layers, inter-layer dropout, lengths for padded batches); and your own tests in python/tests/l3-2-lstm/, written first
Contractcourse/contracts/py/tinyllm/rnn/lstm.pyi
Testscourse/tests/L3.2/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L3.2/lstm_torch.npz; your tests are graded by mutation, threshold 0.70 plus one required fault, with a red-then-green journal
NeedsL0.1 Tensor · L0.2 the ops the cell is written in · L0.4 Module, load_state_dict, Dropout · M07.3 xavier_uniform · M03.3 orthogonal_init · M06.3 the default PCG32 stream · reading: M01.3 sigmoid and tanh, L3.1 BPTT, craft.03 (or --ref-deps)
Used byL3.4 two LSTMs make a bidirectional layer · later: L3.6 cell='lstm', L4.1 encoder option
MilestoneMS-L3 (the LSTM language model must beat the vanilla RNN)
Optional depthHochreiter and Schmidhuber, “Long Short-Term Memory” (1997); Gers, Schmidhuber, Cummins, “Learning to Forget” (2000); Jozefowicz, Zaremba, Sutskever, “An Empirical Exploration of Recurrent Network Architectures” (2015); Olah, Understanding LSTM Networks (2015)
  • An LSTM keeps a second state, the cell ctc_t, updated additively: ct=ft⊙ct−1+it⊙gtc_t = f_t \odot c_{t-1} + i_t \odot g_t. Along it the gradient is multiplied only by the forget gate, so with ff near 1 it survives many steps (test_cell_state_is_a_gradient_highway).
  • The four gate blocks are stacked in torch’s order i,f,g,oi, f, g, o, with two biases, so torch weights load by name and give torch’s outputs and gradients (test_hand_example_lstm_cell, test_matches_torch_cell, test_matches_torch_sequence).
  • Initialization matters at step 0: a forget bias of 1 makes the cell remember by default, and orthogonal recurrent blocks keep the state’s scale (test_init_forget_bias_and_orthogonal_recurrence).
  • In a padded batch each sequence stops at its own length: its outputs past it are 0 and its final state is the state after its last real token (test_lengths_keep_padding_out).
Terminal window
ol start L3.2 # stubs lstm.py; prints your test path and rung (R3)
ol tests L3.2 # the course tests
# write ONE test in python/tests/l3-2-lstm/, then:
ol tdd red L3.2 # must FAIL against your current code
# make it pass, then:
ol tdd green L3.2 # must PASS with the same test files
# repeat for each test; then:
ol check L3.2 # course tests, the journal, the mutation grade

L3.1 showed the vanilla RNN’s problem in numbers: the gradient into a state 30 steps back is the late gradient times Whh⊤W_{hh}^\top thirty times, and with ρ(Whh)=0.8\rho(W_{hh}) = 0.8 that is a factor of about 10−310^{-3}. Your RNN language model (L3.6) will learn which character comes next but not that a quote opened a line ago must close. The LSTM fixes this with one structural change, a memory cell updated by addition instead of by a squashing matmul, plus three gates that decide what to write, what to keep, and what to show. It is also the first module whose weights must match a framework’s layout exactly: torch-trained LSTMs are everywhere, and the tests load torch’s weights by name. With autograd from Part 0 you write only the forward; the backward through time comes for free.

SymbolMeaningType / shape
xtx_tinput at step ttfloat32[B, D]
hth_t, ctc_thidden state (what the next layer sees) and cell state (the memory)float32[B, H] each
WihW_{ih}, WhhW_{hh}input and recurrent weights, four gate blocks stacked[4H, D], [4H, H]
bihb_{ih}, bhhb_{hh}the two biases (torch keeps both)[4H] each
ztz_tall four pre-activations, xtWih⊤+bih+ht−1Whh⊤+bhhx_t W_{ih}^\top + b_{ih} + h_{t-1} W_{hh}^\top + b_{hh}[B, 4H]
it,ft,gt,oti_t, f_t, g_t, o_tinput gate, forget gate, candidate, output gate[B, H] each
σ(v)=1/(1+e−v)\sigma(v) = 1/(1 + e^{-v})the logistic sigmoid, values in (0,1)(0, 1)function
⊙\odotelementwise product
ℓb\ell_bthe length of sequence bb in a padded batchinteger in 1..T1..T

The cell. One matmul per weight computes every gate at once, and the result is split into four blocks of HH columns, in torch’s order:

i=σ(z[0:H]),f=σ(z[H:2H]),g=tanh⁡(z[2H:3H]),o=σ(z[3H:4H])i = \sigma(z_{[0:H]}),\quad f = \sigma(z_{[H:2H]}),\quad g = \tanh(z_{[2H:3H]}),\quad o = \sigma(z_{[3H:4H]}) ct=f⊙ct−1+i⊙g,ht=o⊙tanh⁡(ct).c_t = f \odot c_{t-1} + i \odot g, \qquad h_t = o \odot \tanh(c_t).

The gates are sigmoids because they are fractions: how much of the old memory to keep (ff), how much of the candidate to write (ii), how much of the memory to show (oo). The candidate gg is a tanh because it is content, centered at 0 and able to subtract. The output goes through one more tanh so hh stays in (−1,1)(-1, 1) however large cc grows.

Why it does not vanish. Differentiate the cell update: ∂ct/∂ct−1=diag⁡(ft)\partial c_t/\partial c_{t-1} = \operatorname{diag}(f_t), plus terms through the gates’ dependence on ht−1h_{t-1}. Along the cell path the gradient from cTc_T to c0c_0 is ∏tft\prod_t f_t, elementwise, with no weight matrix and no tanh slope in it. If the network wants to remember, it sets f≈1f \approx 1 and the gradient flows back almost unchanged (Hochreiter’s “constant error carousel”). With Whh=0W_{hh} = 0 the cell path is the only path, and the test measures exactly ∏tft\prod_t f_t.

torch’s layout. torch.nn.LSTMCell and torch.nn.LSTM store weight_ih [4H,D][4H, D] and weight_hh [4H,H][4H, H] (one row per gate unit, so the products use the transpose, like L0.4’s Linear), and two biases bias_ih and bias_hh. The two biases only ever appear as a sum, so a single bias is mathematically enough, but the checkpoint has both keys and both receive the same gradient. A stacked LSTM names its parameters per layer: weight_ih_l0, weight_hh_l0, bias_ih_l0, bias_hh_l0, weight_ih_l1, ..., registered in that order (L0.4: assignment order is state_dict order). Layer k>0k > 0 reads layer k−1k - 1‘s output sequence, so its weight_ih is [4H,H][4H, H], and each layer starts from its own slice of the initial state.

Initialization. The contract draws, per layer, each of the four [H,D][H, D] input blocks with xavier_uniform (M07.3, gain 1) and each of the four [H,H][H, H] recurrent blocks with orthogonal_init (M03.3): an orthogonal matrix has every singular value 1, so the recurrence neither inflates nor shrinks the state at step 0 (the L3.1 analysis with ρ=1\rho = 1). The forget-gate block of bihb_{ih} is 1 and every other bias is 0, so ff starts near σ(1)=0.73\sigma(1) = 0.73: the cell remembers by default (Gers 2000, Jozefowicz 2015). torch’s own default draws everything from U(−1/H,1/H)U(-1/\sqrt{H}, 1/\sqrt{H}); the tests that compare with torch load torch’s weights, so the initializations need not match.

Stacking and dropout. torch applies dropout to each layer’s output sequence except the last layer’s, and only in training mode; the dropout here is L0.4’s Dropout, so eval() turns it off.

Padded batches. Batching sequences of different lengths pads them to TT steps. torch’s pack_padded_sequence makes the RNN skip padding; the equivalent here is a mask: at a step t≥ℓbt \ge \ell_b, sequence bb keeps its (h,c)(h, c) unchanged (F.where(mask, new, old)) and outputs 0. Then the returned final state (hn,cn)(h_n, c_n) is each sequence’s state after its own last real step, and nothing the padding contains reaches any output or receives any gradient.

One unit (D=H=B=1D = H = B = 1), x=1x = 1, ht−1=0h_{t-1} = 0, ct−1=0.5c_{t-1} = 0.5. Weights: Wih=[0,0,ln⁡2,0]⊤W_{ih} = [0, 0, \ln 2, 0]^\top (only the candidate reads xx), Whh=0W_{hh} = 0, bih=[0,ln⁡3,0,ln⁡3]b_{ih} = [0, \ln 3, 0, \ln 3], bhh=0b_{hh} = 0. So z=[0,ln⁡3,ln⁡2,ln⁡3]z = [0, \ln 3, \ln 2, \ln 3].

  1. i=σ(0)=0.5i = \sigma(0) = 0.5; f=σ(ln⁡3)=3/4=0.75f = \sigma(\ln 3) = 3/4 = 0.75; g=tanh⁡(ln⁡2)=(4−1)/(4+1)=0.6g = \tanh(\ln 2) = (4 - 1)/(4 + 1) = 0.6; o=σ(ln⁡3)=0.75o = \sigma(\ln 3) = 0.75.
  2. ct=0.75×0.5+0.5×0.6=0.375+0.3=0.675c_t = 0.75 \times 0.5 + 0.5 \times 0.6 = 0.375 + 0.3 = 0.675.
  3. ht=0.75×tanh⁡(0.675)=0.75×0.588259=0.441194h_t = 0.75 \times \tanh(0.675) = 0.75 \times 0.588259 = 0.441194.
  4. Backward from ∂L/∂ht=1\partial L/\partial h_t = 1: ∂ht/∂ct=o (1−tanh⁡2ct)=0.75×0.653952=0.490463\partial h_t/\partial c_t = o\,(1 - \tanh^2 c_t) = 0.75 \times 0.653952 = 0.490463, and through the cell path ∂ct/∂ct−1=f=0.75\partial c_t/\partial c_{t-1} = f = 0.75, so ∂L/∂ct−1=0.367847\partial L/\partial c_{t-1} = 0.367847.

With the gate order i,f,o,gi, f, o, g the same weights give g=tanh⁡(ln⁡3)=0.8g = \tanh(\ln 3) = 0.8 and o=σ(ln⁡2)=2/3o = \sigma(\ln 2) = 2/3, so ct=0.775c_t = 0.775 and ht=0.433h_t = 0.433: one swapped block, every number different.

This is test_hand_example_lstm_cell.

python/tinyllm/rnn/lstm.py
def lstm_cell(x, h, c, w_ih, w_hh, b_ih, b_hh) -> tuple[Tensor, Tensor] # (h', c')
class LSTMCell(Module): # weight_ih [4H, D], weight_hh [4H, H], bias_ih, bias_hh [4H]
def __init__(self, d_in, d_h, rng=None); def forward(self, x, state=None)
class LSTM(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, tuple[Tensor, Tensor]]
# x [T, B, d_in] -> out [T, B, H], (h_n, c_n) [num_layers, B, H]
TestKINDChecksWhy it matters downstream
test_hand_example_lstm_cellunitsection 3: c′=0.675c' = 0.675, h′=0.441194h' = 0.441194, ∂L/∂c=0.367847\partial L/\partial c = 0.367847you and the test agree on the gates
test_matches_torch_cellgoldentorch.nn.LSTMCell outputs and every gradienttorch weights load and run
test_matches_torch_sequencegoldenone layer with a state, two layers, a packed padded batch: outputs, (hn,cn)(h_n, c_n), all gradientsL3.6 and L4.1 load torch checkpoints
test_state_dict_names_match_torchgoldentorch’s parameter names, order, and shapesthe safetensors key contract (L0.6)
test_gradcheck_lstm_cellgradcheckautograd through the cell against central differences, float64your backward is right without a torch
test_cell_state_is_a_gradient_highwaypropertywith Whh=0W_{hh} = 0, ∂cT/∂c0=∏tft\partial c_T/\partial c_0 = \prod_t f_tthe reason the LSTM remembers
test_init_forget_bias_and_orthogonal_recurrenceunitforget bias 1, others 0; orthogonal WhhW_{hh} blocks; Xavier bound; seedingtrainable from step 0 (L3.6)
test_lengths_keep_padding_outboundarypadded outputs 0; final state equals running the sequence alone; padding gets no gradientbatches of sentences (L4.1)
test_dropout_between_layers_onlyunitdropout between layers, in training mode onlyevaluation is deterministic
test_default_state_and_bad_argumentsboundaryno state means zeros; bad shapes, lengths, sizes raiseerrors at the call site

The given test:

python/tests/l3-2-lstm/test_lstm.py
import math
import numpy as np
from tinyllm.autograd.tensor import Tensor
from tinyllm.rnn.lstm import lstm_cell
def test_hand_example_lstm_cell():
"""i = 0.5, f = 0.75, g = 0.6, o = 0.75: c' = 0.675, h' = 0.75 tanh(0.675)."""
w = [Tensor(a, dtype=np.float64) for a in ([[0.0], [0.0], [math.log(2)], [0.0]], [[0.0]] * 4,
[0, math.log(3), 0, math.log(3)], [0.0] * 4)]
h2, c2 = lstm_cell(Tensor([[1.0]], dtype=np.float64), np.zeros((1, 1)), np.array([[0.5]]), *w)
assert np.allclose(c2.data, [[0.675]])
assert np.allclose(h2.data, [[0.75 * np.tanh(0.675)]])

Then, one at a time, red then green: the cell against your own numpy formula with random weights and both biases nonzero; the state_dict keys of a 2-layer LSTM; the forget bias and orthogonal WhhW_{hh} blocks; a padded batch (zeros past each length, the final state of a short sequence equal to running it alone); a 2-layer LSTM against stacking your numpy cell by hand from a nonzero (h0,c0)(h_0, c_0); dropout ignored by a 1-layer LSTM; bad arguments and the default state. Import only tinyllm.rnn.lstm, tinyllm.autograd.tensor, and tinyllm.autograd.functional (as import tinyllm.autograd.functional as F). The required fault is pitfall 1.

PitfallSymptomCaught by
1. the gates in another order (i,f,o,gi, f, o, g is common in papers and other frameworks)trains fine from scratch; torch weights load and produce garbagetest_hand_example_lstm_cell, test_matches_torch_cell (mutant s01)
2. the wrong squashing: gg through a sigmoid, h=o⊙ch = o \odot c without the tanh, the input gate dropped, or the old cell squashed (f⊙tanh⁡cf \odot \tanh c)the cell can only add, hh grows without bound, or the gradient highway gains a tanh slope per step and vanishes againtest_hand_example_lstm_cell, test_cell_state_is_a_gradient_highway (mutants s02, s03, s04, s15)
3. forget bias 0, or the +1+1 on the wrong blockthe cell forgets half its memory per step at the start; long dependencies learned late or nevertest_init_forget_bias_and_orthogonal_recurrence (mutants s05, s06)
4. padding leaking into the state or the outputsthe final state of a short sentence depends on how long its batch-mates aretest_lengths_keep_padding_out (mutants s07, s08)
5. one bias instead of twothe model computes the same function but a torch checkpoint has an extra key, and loading it drops half the biastest_matches_torch_cell (mutant s09)
6. stacking mistakes: dropout after the last layer, layer kk starting from layer k−1k - 1‘s final statetraining and eval outputs differ in a 1-layer model; 2-layer outputs differ from torchtest_dropout_between_layers_only, test_matches_torch_sequence (mutants s10, s11)
DirectionModuleHow it uses this
BackL0.1parameters are Tensors; x[t] and h0[k] are getitem
BackL0.2F.matmul, F.transpose, F.sigmoid, F.tanh, F.where, F.stack
BackL0.4Module registration order is the key order; Dropout between layers
BackM07.3xavier_uniform for the input blocks
BackM03.3orthogonal_init for the recurrent blocks
BackM06.3PCG32(0).substream("init") when no rng is given
ForwardL3.4bidirectional(LSTM, LSTM, x, lengths)
ForwardL3.6RNNLM(cell='lstm'), trained with stateful TBPTT; must beat the vanilla RNN at MS-L3
ForwardL4.1the seq2seq encoder’s LSTM option

If you skip this module, ol check L3.4 stops with L3.4 needs L3.2: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
lstm_cellcuDNN’s fused LSTMthe input projection for all TT steps as one GEMM, the four gates fused into one kernelaten/src/ATen/native/cudnn/RNN.cpp
lengths maskingtorch.nn.utils.rnn.pack_padded_sequencesorts by length and shrinks the batch as sequences end, so no padded step is computedtorch/nn/utils/rnn.py
forget bias 1Keras unit_forget_bias=Truethe same default, on by defaultkeras/layers/rnn/lstm.py
the LSTM itselfxLSTM (2024)exponential gating and a matrix memory, parallelizable over timeNX-AI/xlstm