Bidirectional RNN with length-aware reversal
Overview
Section titled “Overview”| Module | L3.4 · build · Python · Pass 4 · 2 to 3 h, plus your graded property tests (rung R4) |
| You build | python/tinyllm/rnn/bi.py: reversal_index, reverse_padded, bidirectional; and your own property tests in python/tests/l3-4-bi/ |
| Contract | course/contracts/py/tinyllm/rnn/bi.pyi |
| Tests | course/tests/L3.4/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L3.4/bi_torch.npz; your tests are graded by mutation, threshold 0.80 with every semantic fault required |
| Needs | L3.2 LSTM and L3.3 GRU (the two directions) · L0.1 Tensor (the gather is getitem) · L0.2 F.concat · reading: L0.4 Module, craft.04 property tests (or --ref-deps) |
| Used by | later: L4.1 the seq2seq encoder reads each source sentence both ways |
| Milestone | MS-L3 |
| Optional depth | Schuster and Paliwal, “Bidirectional Recurrent Neural Networks” (1997); Graves and Schmidhuber, “Framewise phoneme classification with bidirectional LSTM” (2005); the torch.nn.utils.rnn documentation |
Key Takeaways
Section titled “Key Takeaways”- A bidirectional layer is two recurrent modules: one reads left to right, the other right to left, and their outputs are concatenated, forward half first, so position sees the whole sentence (
test_hand_example_running_sums,test_matches_torch_bidirectional). - With padding, “right to left” means from each sequence’s own last real token: reverse inside each length with the index and leave the padding where it is; a plain
x[::-1]makes the backward RNN read padding first (test_hand_example_reversal,test_padding_never_reaches_either_direction). - The reversal is an involution, so the same gather puts the backward outputs back in time order, and on a Tensor its gradient is the same permutation (
test_reversal_is_an_involution_that_keeps_padding). - The batched computation equals running every sequence alone, unpadded, which is the property that makes it trustworthy (
test_matches_per_sequence_loop).
How to work this chapter
Section titled “How to work this chapter”ol start L3.4 # stubs bi.py; prints your test path and rung (R4)ol tests L3.4 # the course tests# write the properties of section 4 as tests in python/tests/l3-4-bi/, then:ol check L3.4 # course tests and the mutation grade of your testsol mutate L3.4 # the full grade, cached by your test files' hash1. Why now
Section titled “1. Why now”Your LSTM and GRU read a sentence left to right, so the state at word 3 knows words 0 to 3 and nothing after. For a language model that is the point: it must not see the future. For an encoder it is a handicap: when L4.1 encodes a source sentence to translate it, the representation of “bank” should already know whether “river” or “account” follows. Reading the sentence in both directions and concatenating gives every position both contexts. The idea is one line; the bug is in the batching. Sentences in a batch have different lengths and are padded at the end, and the backward direction must start at each sentence’s last real word, not at the padding. This module builds the length-aware reversal once, proves it against torch’s packed sequences and against a per-sentence loop, and gives L4.1 its encoder.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| padded length, batch size | integers | |
| a padded batch, time first | [T, B, D] | |
| the real length of sequence ; steps are padding | integer in | |
| the reversal index: for , else | integer in | |
| the length-aware reversal, | [T, B, ...] | |
| , | forward and backward states at position | [B, H_f], [B, H_b] |
Reversal inside each length. For sequence the real tokens are . Reading them backward means position takes token for ; positions past the length keep their padding. That map is a permutation of that mirrors the first positions and fixes the rest, so applying it twice gives the identity: . One integer index array of shape holds all permutations, and the gather x[idx, arange(B)] applies them at once; on a Tensor it is L0.1’s getitem, whose backward scatters each gradient back to where its value came from, which is the same permutation again.
Why not x[::-1]. Flipping the padded batch reverses every sequence over all steps: a sequence of length 2 in a batch of length 5 becomes three padding steps followed by its two real tokens. The backward RNN would start from padding, and even though L3.2’s and L3.3’s lengths mask makes them ignore steps past the length, after a full flip the padding is no longer past the length, it is at the front. Its state would be garbage before the first real token.
The layer. With fwd and bwd single-direction modules that honor lengths (output 0 and carry the state past each length):
is in reversed time (its row is the backward state after reading ), so it is reversed back before the concatenation; padded positions are 0 in both halves because fixes them. The forward half comes first, as in torch’s bidirectional=True, whose *_l0 weights are the forward module and *_l0_reverse weights the backward one. Each direction’s final state is in out: the forward one at (its last real step) and the backward one at (it ends on the first token).
3. Worked example by hand
Section titled “3. Worked example by hand”, three sequences with lengths : “abc”, “d”, “ef” (a dot is padding).
| seq 0 () | seq 1 () | seq 2 () | |
|---|---|---|---|
| 0 | a | d | e |
| 1 | b | . | f |
| 2 | c | . | . |
, (a one-token sequence is its own reverse, and both padding steps stay), . As a array, reversal_index([3, 1, 2], 3) , and the reversed batch reads “cba”, “d..”, “fe.”. x[::-1] would read “cba”, “..d”, “.fe”.
A bidirectional layer you can compute. Take as both modules a running sum: output is the sum of the inputs up to , 0 past the length. For the sequence the forward half is the prefix sums . The backward module reads , outputs , and reversing that back gives : the suffix sums, so position holds the sum from to the end. For a second sequence padded with two 99s, both halves are : the 99s never enter.
These are test_hand_example_reversal and test_hand_example_running_sums.
4. The interface
Section titled “4. The interface”def reversal_index(lengths, T: int) -> NDArray # int64 [T, B]def reverse_padded(x, lengths) # Tensor or ndarray [T, B, ...], same kind backdef bidirectional(fwd: Module, bwd: Module, x: Tensor, lengths=None) -> Tensor # [T, B, H_f + H_b]What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_reversal | unit | section 3: the index array and “cba”, “d..”, “fe.” | you and the test agree on the reversal |
test_hand_example_running_sums | unit | prefix sums forward, suffix sums backward, padding 0, forward half first | you and the test agree on the layer |
test_reversal_is_an_involution_that_keeps_padding | property | ; padding fixed; the Tensor gradient is the same permutation | the backward outputs come back in order; gradients reach the right tokens |
test_no_lengths_is_a_plain_flip | unit | without padding, x[::-1] | the unpadded special case |
test_matches_torch_bidirectional | golden | bidirectional=True GRU and LSTM over packed batches: outputs and every gradient | L4.1 loads torch encoders |
test_matches_per_sequence_loop | differential | the batch equals each sequence run alone, both directions | batching changes nothing |
test_padding_never_reaches_either_direction | boundary | garbage padding changes no output, gets no gradient; final states at and 0 | encoder states for the decoder (L4.1) |
test_bad_lengths | boundary | length 0, longer than , wrong count, floats | errors at the batcher |
Your graded tests (rung R4)
Section titled “Your graded tests (rung R4)”At rung R4 the course gives you properties in prose, and you write them as property tests (Hypothesis, craft.04) plus the examples they need. Every semantic fault is required, so each property matters:
- Involution. For random and lengths in ,
reverse_padded(reverse_padded(x, n), n)equalsx. - Mirror and fix. For each , the first rows of
reverse_padded(x, n)[:, b]arex[:ℓ_b, b]reversed and the rest equalx[ℓ_b:, b]. - Plain flip. With
lengths=None, the result isx[::-1]. - Per-sequence equivalence. With two
GRUs,bidirectionalequals running each sequence alone: the forward half fromfwd(x[:ℓ_b, b]), the backward half frombwdon the reversed sequence, reversed back; padded rows are 0. - Padding is inert. Changing the values in the padding changes no output.
- Bad lengths raise.
Import only tinyllm.rnn.bi, tinyllm.rnn.gru (or lstm), and tinyllm.autograd.tensor.
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
1. reversing the padded batch with x[::-1], or mirroring the padding too | the backward direction reads padding before the sentence; short sentences get worse encodings than long ones | test_hand_example_reversal, test_reversal_is_an_involution_that_keeps_padding (mutants s01, s04) |
| 2. not reversing the backward outputs back | position ‘s backward half describes position | test_hand_example_running_sums, test_matches_torch_bidirectional (mutant s02) |
| 3. concatenating backward first | torch encoders load but their output halves are swapped for the decoder | test_hand_example_running_sums, test_matches_torch_bidirectional (mutant s03) |
4. running a direction without lengths | padded outputs are no longer 0, and the backward state at the end of a short sequence has read padding | test_padding_never_reaches_either_direction, test_matches_per_sequence_loop (mutants s05, s06) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L3.2 | LSTM with lengths is one direction |
| Back | L3.3 | GRU with lengths is one direction |
| Back | L0.1 | the reversal is a Tensor getitem with two integer index arrays |
| Back | L0.2 | F.concat joins the halves |
| Forward | L4.1 | the seq2seq encoder: a bidirectional GRU over the source sentence, whose final states initialize the decoder |
| Forward | L4.2 | Bahdanau attention reads the encoder’s per-position outputs, both halves |
If you skip this module, ol check L4.1 stops with L4.1 needs L3.4: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
reverse_padded | torch.nn.utils.rnn.pack_padded_sequence | no padded step is computed at all: sequences are sorted by length and the batch shrinks as they end | torch/nn/utils/rnn.py |
bidirectional | torch.nn.LSTM(bidirectional=True) and cuDNN | both directions in one call, sharing the input GEMM | aten/src/ATen/native/RNN.cpp |
| a bidirectional encoder | BERT (L6.2) | bidirectional context by attention instead of recurrence, every position at once | transformers/models/bert/modeling_bert.py |
| reversal index | Keras go_backwards with masking | the same mask-aware reversal for masked sequences | keras/layers/rnn/bidirectional.py |