Skip to content

Encoder-decoder with teacher forcing

ModuleL4.1 · build · Python · Pass 4 · 4 to 6 h, plus your graded tests (rung R5)
You buildpython/tinyllm/seq2seq/model.py: Seq2Seq (encode, init_state, decode_step, forward, greedy), EncoderState, DecoderState, save_seq2seq, load_seq2seq; and your own oracle tests in python/tests/l4-1-seq2seq/
Contractcourse/contracts/py/tinyllm/seq2seq/model.pyi
Testscourse/tests/L4.1/test_seq2seq.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L4.1/seq2seq_torch.npz (course/oracle/L4.1/seq2seq_torch.py), learning bars in course/fixtures/ref-thresholds.tsv; your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL3.4 bidirectional · L3.3 GRU, GRUCell · L3.2 LSTMCell · L4.2 AdditiveAttention, length_mask · L4.3 LuongAttention · L0.4 Embedding, Linear · L0.2 ops · L0.1 Tensor, no_grad · L0.6 safetensors I/O · M06.3 PCG32 · tests: L0.3 cross_entropy, M10.3 AdamW · reading: L3.1 (or --ref-deps)
Used byL4.4 beam-searches a Seq2Seq in its tests · later L6.7 the zoo’s seq2seq rows on the dates task
MilestoneMS-L4 ({tinyllm} train seq2seq, then {tinyllm} translate --beam 5)
Optional depthSutskever, Vinyals, and Le, “Sequence to Sequence Learning with Neural Networks” (2014); Cho et al., “Learning Phrase Representations using RNN Encoder-Decoder” (2014); Bengio et al., “Scheduled Sampling for Sequence Prediction with Recurrent Neural Networks” (2015)
  • An encoder reads the source, a bridge turns its two final states into the decoder’s first state, and the decoder writes the target one token at a time (test_hand_example_zero_weights, test_golden_torch).
  • Teacher forcing feeds the true previous token, so training is decode_step over the target, all steps in one graph; scheduled sampling sometimes feeds the model’s own guess instead (test_forward_is_the_step_loop, test_teacher_forcing_ratio).
  • Bahdanau attention reads with the state from before the step, Luong’s with the state after it, plus input feeding (test_bahdanau_reads_with_the_previous_state, test_luong_reads_with_the_new_state_and_feeds_input).
  • Each row’s final encoder states are read at its own length, so padding never changes a sentence (test_padding_never_changes_a_sentence).
  • With the same 150 steps, attention reverses digit strings almost perfectly and the bottleneck model does not (test_attention_learns_to_reverse, test_attention_beats_the_bottleneck).
Terminal window
ol start L4.1 # stubs model.py; prints your test path and rung (R5)
ol tests L4.1 # the course tests
# write your oracle tests in python/tests/l4-1-seq2seq/, then:
ol check L4.1 # course tests and the mutation grade of your tests
ol check L4.1 --ref-deps # if you skipped L3.2 to L3.4, L4.2, or L4.3
ol diff L4.1 # after passing: your code against the reference

Every model so far maps a sequence to the next token of the same sequence. Translation, summarization, and the dates task of MS-L4 ("3 March 2021" to "2021-03-03") map one sequence to a different one, of a different length, in a different vocabulary. The encoder-decoder is the architecture that does that, and it is the setting where attention was invented: this module builds the model with and without attention, and its learning test measures the gap. Without attention the decoder sees the source only through one vector; the bottleneck test shows what that costs on a task as simple as reversing six digits.

SymbolMeaningType / shape
x1:Sx_{1:S}, ℓb\ell_bsource ids, padded to SS, real length ℓb\ell_bint[B, S], int[B]
y0:T−1y_{0:T-1}decoder inputs tgt_in: bos then the target without its last tokenint[B, T]
dhd_hdecoder width; each encoder direction has dh/2d_h / 2int (even)
ot=[h→t;h←t]o_t = [\overrightarrow{h}_t ; \overleftarrow{h}_t]encoder output at source position tt (the keys)float32[B, S, d_h]
s0=tanh⁡(Wb[h→ℓ−1;h←0]+bb)s_0 = \tanh(W_b [\overrightarrow{h}_{\ell - 1} ; \overleftarrow{h}_0] + b_b)the bridge: the decoder’s first statefloat32[B, d_h]
sts_tdecoder state after step ttfloat32[B, d_h]
ctc_t, h~t\tilde h_tcontext and Luong’s attentional statefloat32[B, d_h]
ρ\rhoteacher_forcing, the probability of feeding the true tokenfloat in [0,1][0, 1]

The encoder is L3.4’s bidirectional over two L3.3 GRUs of width dh/2d_h/2: position tt‘s output oto_t sees the whole source, left context from the forward GRU and right context from the backward one. The bridge summarizes the source for the decoder: the forward GRU’s state after the last real token (position ℓb−1\ell_b - 1) and the backward GRU’s state after reading back to position 0, concatenated, through a Linear and a tanh. The decoder is a GRUCell (or LSTMCell, with c0=0c_0 = 0) over target embeddings, and an output Linear gives logits over the target vocabulary.

One decoder step from token yy and state ss, by attention type (attention.query_from):

AttentionStep
nones′=cell(E[y],s)s' = \mathrm{cell}(E[y], s), logits =Wos′+bo= W_o s' + b_o
"previous" (Bahdanau, L4.2)c,a=att(s,o)c, a = \mathrm{att}(s, o); s′=cell([E[y];c],s)s' = \mathrm{cell}([E[y] ; c], s); logits =Wo[s′;c]+bo= W_o [s' ; c] + b_o
"current" (Luong, L4.3)s′=cell([E[y];h~],s)s' = \mathrm{cell}([E[y] ; \tilde h], s); c,a=att(s′,o)c, a = \mathrm{att}(s', o); h~′=tanh⁡(Wc[c;s′])\tilde h' = \tanh(W_c [c ; s']); logits =Woh~′+bo= W_o \tilde h' + b_o

The parameters are registered in the order of the contract (src_emb, enc_fwd, enc_bwd, bridge, tgt_emb, cell, attention, out), and that order is the safetensors key list of a checkpoint.

Training maximizes ∑tlog⁡p(yt+1∣y≤t,x)\sum_t \log p(y_{t+1} \mid y_{\le t}, x). Teacher forcing feeds the true yty_t at step tt regardless of what the model predicted, so all TT losses come from one forward pass and the gradient at step tt does not depend on the model’s own mistakes. At inference the model must feed its own guesses, a mismatch called exposure bias. Scheduled sampling (Bengio et al. 2015) trains with ρ<1\rho < 1: before each step after the first, one uniform uu is drawn for the batch; if u<ρu < \rho the true token is fed, otherwise the argmax of the previous step’s logits, with no gradient through the choice. ρ=1\rho = 1 draws nothing, so a seeded run is unchanged by adding the option.

Targets are shifted by one: tgt_in = bos y1…yT−1y_1 \dots y_{T-1}, and the loss compares step tt‘s logits with yt+1y_{t+1} (eos at the end). Padding in the target is ignored with ignore_index (L0.3).

Attention is a module passed to the constructor; the decoder reads its query_from attribute to choose the order of section 2.2. The encoder computes attention.project_keys(o) once per sentence and stores it in the state, so the keys are projected once, not at every step (L4.2 section 2.4). decode_step returns the attention weights too, which is what an alignment plot draws.

greedy decodes all rows at once under no_grad, feeding back each row’s argmax and stopping a row at its first eos. L4.4 replaces it with beam search over decode_step.

Take a GRU encoder-decoder with every GRU weight and bias 0. One GRU step gives r=z=σ(0)=12r = z = \sigma(0) = \tfrac12 and n=tanh⁡(0+r⋅0)=0n = \tanh(0 + r \cdot 0) = 0, so

h′=(1−z) n+z h=12h.h' = (1 - z)\, n + z\, h = \tfrac12 h .

From h0=0h_0 = 0 the encoder outputs 0 everywhere, whatever the source tokens. Let bridge.weight be 0 and bridge.bias =(0.5,−0.5,1,0)= (0.5, -0.5, 1, 0): s0=tanh⁡(0.5,−0.5,1,0)=(0.462117,−0.462117,0.761594,0)s_0 = \tanh(0.5, -0.5, 1, 0) = (0.462117, -0.462117, 0.761594, 0). Each decoder step halves it: s1=(0.231059,−0.231059,0.380797,0)s_1 = (0.231059, -0.231059, 0.380797, 0), s2=(0.115529,−0.115529,0.190399,0)s_2 = (0.115529, -0.115529, 0.190399, 0). With out.weight the first two rows of the identity and no bias, logits1=(0.231059,−0.231059)_1 = (0.231059, -0.231059) and logits2=(0.115529,−0.115529)_2 = (0.115529, -0.115529), whatever the target tokens. The example fixes the data path: encoder, bridge, decoder cell, output layer; dropping the bridge gives zero logits. This is test_hand_example_zero_weights.

class EncoderState(NamedTuple): keys; mask; proj; init
class DecoderState(NamedTuple): h; c; feed; keys; mask; proj
class Seq2Seq(Module):
def __init__(self, src_vocab, tgt_vocab, d_emb, d_h, cell="gru", attention=None, rng=None): ...
def encode(self, src, src_lens) -> EncoderState: ...
def init_state(self, enc: EncoderState) -> DecoderState: ...
def decode_step(self, y_prev, state) -> tuple[Tensor, DecoderState, Optional[Tensor]]: ...
def forward(self, src, src_lens, tgt_in, teacher_forcing=1.0, rng=None) -> Tensor: ... # [B, T, Vt]
def greedy(self, src, src_lens, bos, eos, max_len) -> list[list[int]]: ...
def save_seq2seq(model, dir) -> None: ...
def load_seq2seq(dir) -> Seq2Seq: ...
TestKINDChecksWhy it matters downstream
test_hand_example_zero_weightsunitsection 3: logits s0/2s_0/2, s0/4s_0/4the data path through the bridge
test_golden_torchgoldenlogits and every gradient against torch’s bidirectional nn.GRU, GRUCell, LSTMCell, four attention configurationsyour model computes what the papers describe
test_forward_is_the_step_loopdifferentialforward equals decode_step over tgt_inbeam search and the zoo drive the same steps
test_teacher_forcing_ratiounitρ=0\rho = 0 feeds the model’s argmax, one draw per step; ρ=1\rho = 1 draws nothingscheduled sampling, reproducibly
test_bahdanau_reads_with_the_previous_stateunitthe query is s0s_0, then s1s_1Bahdanau’s order
test_luong_reads_with_the_new_state_and_feeds_inputunitthe query is the new state; h~\tilde h reaches the next input; logits read h~\tilde hLuong’s order and input feeding
test_padding_never_changes_a_sentencepropertygarbage padding (even out-of-vocabulary ids) changes nothing; a row alone equals it in a batchbatches of mixed lengths
test_gradient_reaches_the_encoderunitevery parameter gets a nonzero gradientthe encoder learns from the decoder’s loss
test_greedy_matches_step_loop_and_stops_at_eosunitgreedy equals its own step loop; eos ends a rowthe decoder MS-L4 runs
test_save_load_roundtripunitconfig, key order, identical logits for every attentionthe zoo’s checkpoint contract
test_validationboundaryodd dhd_h, unknown cell, non-attention module, bad lengths and idscaller bugs fail loudly
test_attention_learns_to_reverselearning150 steps of AdamW on reversal: exact match at the reference within 3 sdthe model trains
test_attention_beats_the_bottlenecklearningsame budget without attention: at least 0.2 lower exact matchthe lesson of this part

The oracles for your tests are the model’s own pieces recombined by hand: run enc_fwd and enc_bwd on one unpadded sentence at a time and check init against tanh⁡(Wb[… ])\tanh(W_b [\dots]); recompute one decoder step from cell, attention, and out and compare with decode_step; check that forward equals the step loop and that free running feeds back the argmax. Add the zero-weights example, padding with out-of-vocabulary ids, a greedy decode that must stop at once (a huge eos bias), and a save and load. Import only contract modules.

PitfallSymptomCaught by
1. the decoder starts from zeros (bridge state dropped)the source reaches the decoder only through attention, or not at alltest_hand_example_zero_weights (mutant s01)
2. teacher forcing fed tgt_in[:, t - 1] at step ttthe model is asked to copy its input; training loss falls, translation failstest_forward_is_the_step_loop (mutant s02)
3. final encoder states read at the padded end or the wrong positionshort sentences are summarized from paddingtest_padding_never_changes_a_sentence (mutant s06), test_gradient_reaches_the_encoder (mutant s07)
4. attention in the wrong place: Luong reading with the old state, Bahdanau’s query not st−1s_{t-1}, or its context left out of the outputa different model from the paper’s; golden values disagreetest_luong_reads_with_the_new_state_and_feeds_input (mutant s04), test_bahdanau_reads_with_the_previous_state (mutant s09), test_golden_torch (mutant s03)
5. no input feedingthe next step never sees h~\tilde htest_luong_reads_with_the_new_state_and_feeds_input (mutant s05)
teacher forcing when u≥ρu \ge \rhoρ=0\rho = 0 still feeds the truthtest_teacher_forcing_ratio (mutant s08)
DirectionModuleHow it uses this
BackL3.4bidirectional runs the encoder both ways inside each length
BackL3.3the encoder GRUs and the default decoder GRUCell
BackL3.2LSTMCell, the other decoder cell
BackL4.2AdditiveAttention and length_mask
BackL4.3LuongAttention and attentional
BackL0.4Embedding and Linear
BackL0.2the op library
BackL0.1Tensor and no_grad
BackL0.6save_safetensors and load_safetensors
BackM06.3PCG32 for the default initialization
BackL0.3cross_entropy with ignore_index, in the tests
BackM10.3AdamW, in the tests
ForwardL4.4beam search over decode_step, gathering DecoderState by parent
ForwardL4.5exact match, BLEU, and chrF grade its outputs
ForwardL6.7the zoo trains and evaluates seq2seq checkpoints on the dates task
ForwardL5.3the 2017 transformer keeps the encoder-decoder and replaces both RNNs with attention
Your pieceProduction equivalentWhat it addsWhere to look
Seq2SeqOpenNMT-py, fairseq LSTM modelsmulti-layer stacks, dropout, copy attention, batched beam searchOpenNMT-py onmt/models/model.py; fairseq models/lstm.py
the bridgebridge in OpenNMTa learned map from encoder to decoder state per layerOpenNMT-py onmt/encoders/rnn_encoder.py
teacher forcing ratioscheduled sampling, professor forcingcurricula for ρ\rho, adversarial matching of free-running dynamicsBengio et al. 2015; Lamb et al. 2016
encoder-decoder RNNT5, BARTthe same split with transformer blocks and span corruptionL6.4; Hugging Face T5ForConditionalGeneration