Skip to content

RNN language model with stateful TBPTT

ModuleL3.6 · build · Python · Pass 4 · 4 to 5 h, plus your graded property tests (rung R4)
You buildpython/tinyllm/rnn/rnnlm.py: ElmanRNN (L3.1’s forward and backward as one autograd op), RNNLM (forward, init_state, detach_state, nll, generate), tbptt_batches, train_tbptt, save_rnnlm, load_rnnlm; and your own property tests in python/tests/l3-6-rnnlm/
Contractcourse/contracts/py/tinyllm/rnn/rnnlm.pyi
Testscourse/tests/L3.6/test_rnnlm.py (what they check: section 4); the learning test trains on course/fixtures/L3.6/corpus.txt (an original synthetic text, course/oracle/L3.6/corpus.py) against a bar in course/fixtures/ref-thresholds.tsv; your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL3.1 rnn_forward, rnn_backward · L3.2 LSTM · L3.3 GRU · L0.5 train_step · L0.3 cross_entropy · L0.4 Embedding, Linear · L0.2 ops · L0.1 Tensor, from_op, no_grad · L0.6 safetensors I/O · M07.3 xavier_uniform · M03.3 orthogonal_init · M06.3 PCG32 · tests: L1.1 CharTokenizer, M10.3 AdamW · reading: M10.4 (clipping), M11.2 (bits per character) (or --ref-deps)
Used bylater: L6.7 the zoo’s rnnlm bits-per-byte rows (joins the registry with B7)
MilestoneMS-L3 ({tinyllm} train rnnlm --cell {rnn,lstm,gru}: LSTM and GRU at the calibrated bar, RNN worse than LSTM)
Optional depthMikolov et al., “Recurrent neural network based language model” (Interspeech 2010); Zaremba, Sutskever, and Vinyals, “Recurrent Neural Network Regularization” (2014), section 4; Williams and Peng, “An efficient gradient-based algorithm for on-line training of recurrent network trajectories” (1990)
  • An RNN language model is an embedding, a recurrent layer, and a linear read-out; the three cells share one interface, and the Elman cell is L3.1’s hand-written BPTT wrapped as a single autograd op (test_elman_is_one_fused_op_over_l3_1).
  • Stateful training cuts the stream into contiguous lanes and walks them window by window: the hidden state flows from one window to the next (test_hand_example_windows, test_state_carries_across_calls).
  • The gradient stops at the window boundary: the carried state is detached, not reset (test_gradient_is_cut_at_the_window_boundary, test_state_flows_into_the_next_window).
  • Each new epoch starts from zeros, and clipping guards every step (test_epoch_wraps_and_resets_the_state, test_clip_is_applied).
  • A character LSTM trained this way for 150 steps beats the unigram entropy of the text by more than half a bit per character (test_lstm_learns_characters).
Terminal window
ol start L3.6 # stubs rnnlm.py; prints your test path and rung (R4)
ol tests L3.6 # the course tests
# write the properties of section 4 as tests in python/tests/l3-6-rnnlm/, then:
ol check L3.6 # course tests and the mutation grade of your tests
ol diff L3.6 # after passing: your code against the reference

The n-gram model (L2.1) and the NPLM (L2.2) predict the next token from a fixed window: whatever happened 20 characters ago is invisible to them. You now have recurrent layers (L3.1 to L3.3) whose state can, in principle, carry information across any distance. This module turns them into a language model and trains it the way recurrent LMs were trained in practice: on one long stream, with a state that persists across training windows. The result is the first model in the zoo (L6.7) whose context is not bounded by its architecture, and MS-L3 compares its bits per byte with the NPLM’s on the same text.

SymbolMeaningType / shape
VV, ded_e, dhd_hvocabulary size, embedding width, hidden widthint
xtx_ttoken id at position ttint
EEembedding tablefloat32[V, d_e]
hth_trecurrent state after reading xtx_t (hh and cc for an LSTM)float32[n_layers, B, d_h]
WoW_o, bob_oread-outfloat32[V, d_h], [V]
NNstream lengthint
BBnumber of lanes (batch)int
L=⌊N/B⌋L = \lfloor N / B \rfloorlane lengthint
kkwindow length (truncation)int
bpc\mathrm{bpc}bits per character: mean NLL in nats divided by ln⁡2\ln 2float

ht=cell(E[xt],ht−1),logitst=Woht+bo,L=1T∑t−log⁡softmax⁡(logitst)[xt+1].h_t = \mathrm{cell}(E[x_t], h_{t-1}), \qquad \text{logits}_t = W_o h_t + b_o, \qquad \mathcal{L} = \frac{1}{T}\sum_t -\log \operatorname{softmax}(\text{logits}_t)[x_{t+1}] .

cell is L3.2’s LSTM, L3.3’s GRU, or ElmanRNN, all with the same forward(x, state) -> (out, state). ElmanRNN is the vanilla ht=tanh⁡(xtWih⊤+ht−1Whh⊤+bih+bhh)h_t = \tanh(x_t W_{ih}^\top + h_{t-1} W_{hh}^\top + b_{ih} + b_{hh}) in torch’s parameter names; its forward over a whole sequence is one autograd node built with from_op: forward calls L3.1’s rnn_forward, backward calls rnn_backward on the saved cache. That is how cuDNN runs recurrent layers (one fused kernel per layer, not one graph node per time step), and it makes your hand-written BPTT the engine of a real model. Two details: L3.1 uses Wxh=Wih⊤W_{xh} = W_{ih}^\top, so the weight gradients are transposed back; the two biases are summed in the forward, so each receives the full bias gradient.

A long text could be backpropagated through in one piece only with memory proportional to its length. Truncated BPTT runs windows of kk tokens. The key choice is what state starts each window:

  • Stateless: zeros every time. Simple, but the model never sees context from before the window, so it cannot learn dependencies longer than kk.
  • Stateful: the state the previous window ended with. The forward pass sees unbounded context; only the gradient is truncated at the window start.

Stateful training needs the state to be a value at the boundary: detach_state keeps the arrays and drops the graph, so backward on window j+1j + 1 stops there instead of running on into window jj‘s graph (whose gradients were already applied).

For the carried state to make sense, row bb of window j+1j + 1 must continue row bb of window jj. So the stream is cut into BB contiguous lanes, lane bb = stream[b L : (b + 1) L] (the tail N−BLN - BL is dropped), and window jj is columns [jk,jk+k)[jk, jk + k) of every lane as inputs, shifted by one as targets. A window needs k+1k + 1 tokens, so the windows start at p=0,k,2k,…p = 0, k, 2k, \dots while p+k+1≤Lp + k + 1 \le L. After the last window the lanes start over and the state is reset to zeros: carrying the end of a lane into its own beginning would condition text on what never precedes it.

The exploding gradients L3.1 diagnosed are real in training: one bad window can produce a gradient that throws the weights far away. train_tbptt passes clip to L0.5’s train_step, which rescales the global gradient norm to at most clip (M10.4) before every step.

nll(ids) scores a held-out stream in one lane, chunk by chunk under no_grad, carrying the state, so its result does not depend on the chunk size; the mean divided by ln⁡2\ln 2 is bits per character, and per byte with M11.2’s accumulator. generate feeds each sampled token back in, drawing from PCG32(seed) exactly as L0.5’s bigram sampler does.

The stream 0,1,…,190, 1, \dots, 19, B=2B = 2 lanes, k=3k = 3. L=10L = 10: lane 0 is 0…90 \dots 9, lane 1 is 10…1910 \dots 19. A window needs k+1=4k + 1 = 4 tokens, so it starts at p=0,3,6p = 0, 3, 6 (p=9p = 9 would need token 12 of a 10-token lane):

windowinputstargets
0[[0 1 2] [10 11 12]][[1 2 3] [11 12 13]]
1[[3 4 5] [13 14 15]][[4 5 6] [14 15 16]]
2[[6 7 8] [16 17 18]][[7 8 9] [17 18 19]]

The state after reading 0 1 2 starts the read of 3 4 5: lane 0 continues itself. Tokens 9 and 19 are only targets. With interleaved lanes (reshape(L, B).T), lane 0 would be 0,2,4,…0, 2, 4, \dots and the state would carry across text that is not contiguous. This is test_hand_example_windows.

class ElmanRNN(Module):
def __init__(self, d_in, d_h, num_layers=1, rng=None): ...
def forward(self, x, state=None, lengths=None) -> tuple[Tensor, Tensor]: ...
class RNNLM(Module):
def __init__(self, vocab, d_emb, d_h, cell, n_layers=1, rng=None): ...
def forward(self, ids, state=None) -> tuple[Tensor, Any]: ... # ids [B, T] -> logits [B, T, V]
def init_state(self, batch): ...; def detach_state(self, state): ...
def nll(self, ids, chunk=256) -> NDArray: ...; def generate(self, prefix, n, temperature, seed) -> list[int]: ...
def tbptt_batches(stream, k, batch) -> list[tuple[NDArray, NDArray]]: ...
def train_tbptt(model, stream, k, batch, opt, clip, steps) -> list[float]: ...
def save_rnnlm(model, dir, tokenizer="bytes") -> None: ...; def load_rnnlm(dir) -> RNNLM: ...
TestKINDChecksWhy it matters downstream
test_hand_example_windowsunitsection 3’s three windowsthe lanes the state flows along
test_windows_drop_the_tail_and_validateboundarythe tail is dropped; too-short lanes raiseno empty epochs
test_elman_is_one_fused_op_over_l3_1differentialoutput and every gradient equal rnn_forward and rnn_backwardyour BPTT drives a real model
test_gradcheck_rnnlmgradcheckevery parameter of a 2-layer model, Elman and LSTM cells, float64the model trains through all its layers
test_state_carries_across_callspropertytwo calls with the carried state equal one callthe definition of stateful
test_gradient_is_cut_at_the_window_boundarydifferentialstep 2’s gradients equal window 2 alone from a constant statetruncation is where you put it
test_state_flows_into_the_next_windowdifferentialstep 2’s loss equals the run from window 1’s state, not from zeroscontext crosses windows
test_epoch_wraps_and_resets_the_stateunitstep len(windows) repeats step 0 exactlyeach epoch starts clean
test_clip_is_appliedunitthe optimizer sees a gradient norm of at most clipthe guard against explosions
test_nll_does_not_depend_on_the_chunkpropertychunk 1, 7, 256 agree and equal one passevaluation of long held-out text
test_generate_is_seeded_and_greedy_is_argmaxunitsame seed, same ids; greedy feeds back the argmax{tinyllm} generate in MS-L3
test_save_load_roundtripunitconfig.json, torch key names, identical logitsthe zoo’s checkpoint contract
test_validationboundaryunknown cell, bad ids, lengths for the Elman layer, negative stepscaller bugs fail loudly
test_lstm_learns_characterslearning150 steps: held-out bpc at the reference bar and 0.5 below the unigram entropythe model learns

Rung R4 grades properties: write them as tests that hold for any seed. The ones this module rests on: two calls with a carried state equal one call; nll does not depend on the chunk; with an optimizer that only records gradients (a small class with zero_grad and step), step 2’s gradients equal those of window 2 from a constant state, the step len(windows) loss equals step 0’s, and every recorded gradient norm is at most clip. Add the section 3 windows, the Elman layer against rnn_forward, greedy generation, and a save and load. Import only contract modules.

PitfallSymptomCaught by
1. the carried state not detachedbackward runs into the previous window’s graph again; gradients double-counttest_gradient_is_cut_at_the_window_boundary (mutant s06)
2. “detaching” by resetting to zeros every windowstateless training: no context beyond kktest_state_flows_into_the_next_window (mutant s07)
3. interleaved lanesthe state carries across non-contiguous texttest_hand_example_windows (mutant s01)
4. no clippingone exploding window ruins the runtest_clip_is_applied (mutant s09)
bias gradient to bias_ih onlybias_hh never trainstest_elman_is_one_fused_op_over_l3_1 (mutant s03)
weight gradients not transposed backwrong updates (or a crash when de≠dhd_e \ne d_h)test_elman_is_one_fused_op_over_l3_1 (mutant s04)
forward ignores its stateno statefulness at alltest_state_carries_across_calls (mutant s05)
no reset at the epoch startthe first window is conditioned on the lane’s endtest_epoch_wraps_and_resets_the_state (mutant s08)
nll restarting the state per chunkthe score depends on the chunk sizetest_nll_does_not_depend_on_the_chunk (mutant s10)
generate feeding back the wrong tokengreedy text differs from the argmax looptest_generate_is_seeded_and_greedy_is_argmax (mutant s11)
tokenizer ignored by save_rnnlmthe zoo loads the wrong tokenizertest_save_load_roundtrip (mutant s12)
DirectionModuleHow it uses this
BackL3.1rnn_forward and rnn_backward are the Elman layer’s forward and backward
BackL3.2LSTM, cell = "lstm"
BackL3.3GRU, cell = "gru"
BackL0.5train_step runs each window (zero_grad, backward, clip, step)
BackL0.3cross_entropy over the window
BackL0.4Embedding and Linear
BackL0.2F.transpose, F.stack, F.reshape
BackL0.1Tensor, from_op, no_grad
BackL0.6save_safetensors and load_safetensors for the model directory
BackM07.3xavier_uniform for the Elman input weights
BackM03.3orthogonal_init for the Elman recurrent weights
BackM06.3PCG32 for initialization and sampling
BackL1.1CharTokenizer encodes the corpus in the learning test
BackM10.3AdamW trains the learning test
ForwardL6.7the zoo trains rnnlm checkpoints and reports their bits per byte next to the n-gram and the transformer
ForwardL3.5ELMo’s biLM is two of these language models, one per direction
Your pieceProduction equivalentWhat it addsWhere to look
train_tbpttPyTorch’s word language model examplethe same batchify lanes and repackage_hidden detachpytorch/examples, word_language_model/main.py
ElmanRNN fused opcuDNN RNN kernelsone kernel per layer, weights packed for both directionstorch/nn/modules/rnn.py, _VF.rnn_tanh
LSTM LMAWD-LSTMweight-dropped hidden matrices, variational dropout, averaged SGDMerity, Keskar, and Socher (2017), salesforce/awd-lstm-lm
truncationRWKV, Mambarecurrent models trained in parallel over the whole sequence, run as RNNs at inferencePeng et al. 2023; Gu and Dao 2023