Bahdanau additive attention
Overview
Section titled “Overview”| Module | L4.2 · build · Python · Pass 4 · 3 to 4 h, plus your graded tests (rung R5) |
| You build | python/tinyllm/seq2seq/additive.py: length_mask and AdditiveAttention (project_keys, scores, forward); and your own oracle tests in python/tests/l4-2-additive/ |
| Contract | course/contracts/py/tinyllm/seq2seq/additive.pyi |
| Tests | course/tests/L4.2/test_additive.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L4.2/additive_torch.npz (course/oracle/L4.2/additive_torch.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required |
| Needs | L0.4 Linear, Module · L0.2 F.tanh, F.masked_fill, F.softmax, F.matmul · L0.1 Tensor · M06.3 PCG32 (default initialization) · reading: M09.2 (the masked softmax), S-M08 (its VJP by hand) (or --ref-deps) |
| Used by | L4.3 builds masks with length_mask · L4.1 plugs AdditiveAttention into the decoder · later L6.7 the zoo’s attention rows |
| Milestone | MS-L4 (em_bahdanau >= 0.97 on the dates task) |
| Optional depth | Bahdanau, Cho, and Bengio, “Neural Machine Translation by Jointly Learning to Align and Translate” (ICLR 2015), section 3.1 and appendix A.1.2; Graves, “Generating Sequences With Recurrent Neural Networks” (2013), section 5 |
Key Takeaways
Section titled “Key Takeaways”- Attention is a weighted average of the encoder outputs whose weights the decoder computes: a score per source position, a softmax, a sum (
test_hand_example_weights_and_context). - Padding is masked on the scores, to , before the softmax: padded weights are exactly 0, the rest sum to 1, and padding content never matters (
test_masked_weights_are_exactly_zero_and_rows_sum_to_one,test_padding_content_never_matters). - An empty source reads nothing: weights and context 0, not NaN (
test_fully_masked_row_reads_nothing). - The key projection does not depend on the decoder step, so it is computed once per sentence (
test_precomputed_keys_equal_recomputed). - Attention by itself ignores order: permute the source and the weights permute with it (
test_permuting_the_source_permutes_the_weights).
How to work this chapter
Section titled “How to work this chapter”ol start L4.2 # stubs additive.py; prints your test path and rung (R5)ol tests L4.2 # the course tests# write your oracle tests in python/tests/l4-2-additive/, then:ol check L4.2 # course tests and the mutation grade of your testsol diff L4.2 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your encoder-decoder (L4.1, next in this pass) squeezes the whole source sentence into one vector, the decoder’s first state. On short inputs that works; on a 30-character date string the decoder has forgotten the start of the input by the time it writes the end, and exact match drops sharply with length. Bahdanau’s fix, the first attention mechanism in neural translation, lets every decoder step look back at all encoder outputs and choose which to read. It is the idea the 2017 transformer (L5.1) keeps while dropping the recurrence, so this is where the core operation of every model in the rest of the course is born.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| , | batch size, source length (padded) | int |
| the query: the decoder state that asks | float32[B, Dq] | |
| the key at source position : an encoder output, also the value read | float32[B, S, Dk] | |
| the real length of source | int | |
| mask: true when | bool[B, S] | |
| , , , | parameters query.weight , key.weight , key.bias , v.weight | float32 |
d_attn, the width of the scoring network | int | |
| the score of position | float | |
| the attention weight of position | float in | |
| the context: | float32[B, Dk] |
2.1 Reading by weighted average
Section titled “2.1 Reading by weighted average”A decoder that must write the next target word needs some of the source, not all of it equally. A weighted average of the encoder outputs, with weights that sum to 1, reads exactly that: a weight near 1 on one position copies that position; spread weights blend several. The weights must depend on what the decoder is doing now, so they are computed from the decoder state .
2.2 Scores, weights, context
Section titled “2.2 Scores, weights, context”Bahdanau scores each position with a one-hidden-layer network on the pair :
The softmax (M09.2) turns arbitrary real scores into positive weights that sum to 1. Every operation is an op-library call (L0.2), so backward reaches (the decoder learns where to look), the keys (the encoder learns what to offer), and . The context has the keys’ width , not : the scoring network decides how much of each key to take, the keys themselves are what is read. Nothing in the formula looks at the position itself: permute the keys and the weights permute with them. Order has to come from the encoder.
2.3 Masking padding
Section titled “2.3 Masking padding”A batch pads every source to the longest length . Padded positions must never be read. The rule is to set their scores to before the softmax: , so their weights are exactly 0 and the real weights still sum to 1. length_mask(lengths, S) builds . Masking after the softmax (multiplying the weights by the mask) leaves the real weights summing to less than 1 and lets padding influence the normalization. A row with no real position at all (an empty source) is all ; M09.2 defines its softmax as zeros, so the context is 0 and no gradient flows back, where a naive softmax would give $0/0 = $ NaN.
2.4 Precomputing the keys
Section titled “2.4 Precomputing the keys” involves only the keys, which are fixed for the whole sentence, while changes at every decoder step. Computing once (project_keys) and passing it to every step as proj saves multiply-adds per step; Bahdanau’s appendix points this out. The result must be identical to recomputing.
3. Worked example by hand
Section titled “3. Worked example by hand”Every size 1 (), , . Query , keys , and the third position is padding (mask [True, True, False]).
- Scores before the mask: , , . The padded one would be the largest.
- After the mask: .
- Softmax: , , ; the sum is , so .
- Context: .
The padded key, the biggest value in the row, contributes nothing. This is test_hand_example_weights_and_context.
4. The interface
Section titled “4. The interface”def length_mask(lengths: ArrayLike, S: int) -> NDArray: ... # bool [B, S]class AdditiveAttention(Module): query_from = "previous" # read by L4.1's decoder def __init__(self, d_query: int, d_key: int, d_attn: int, rng=None) -> None: ... def project_keys(self, keys: Tensor) -> Tensor: ... # [B, S, A] def scores(self, query: Tensor, keys: Tensor, proj=None) -> Tensor: ... # [B, S], unmasked def forward(self, query, keys, mask, proj=None) -> tuple[Tensor, Tensor]: ... # context, weightsWhat the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_weights_and_context | unit | section 3 to float32 tolerance | you and the test agree on the formula |
test_length_mask | unit | s < lengths[b], bad lengths raise | the only place lengths enter attention |
test_golden_torch | golden | context, weights, and gradients of the query, keys, and all four parameters against torch, lengths 5, 3, 1 | your encoder and decoder train through this |
test_gradcheck_every_input_and_parameter | gradcheck | float64 central differences with a padded row | both sides of the model learn |
test_masked_weights_are_exactly_zero_and_rows_sum_to_one | property | exact zeros, positive real weights, rows sum to 1 | padding is never read |
test_padding_content_never_matters | property | garbage in padded keys changes nothing bitwise; their gradient is 0 | a batch padded differently translates the same |
test_fully_masked_row_reads_nothing | boundary | zeros, no NaN, no gradient | empty inputs in a batch |
test_precomputed_keys_equal_recomputed | property | proj passed or recomputed, bitwise equal | L4.1 passes it at every step |
test_permuting_the_source_permutes_the_weights | property | weights permute, context unchanged | order comes from the encoder, not from attention |
test_parameter_names_and_shapes | unit | keys, shapes, order, query_from, seeded init | the safetensors keys of seq2seq checkpoints |
test_validation | boundary | wrong widths, missing source axis, wrong or float mask, wrong proj | wiring bugs fail loudly |
Your graded tests (rung R5)
Section titled “Your graded tests (rung R5)”The oracle for your tests is the formula itself, written out in numpy with the module’s own weights read from state_dict(): compute , mask, softmax, and the context in float64, and compare with forward. Add the section 3 numbers, a finite-difference check of the gradients of the query and the keys (your own, in float64), the masking properties, and the parameter names. Import only tinyllm.seq2seq.additive, tinyllm.autograd.tensor, and tinyllm.autograd.functional (as import tinyllm.autograd.functional as F). ol check L4.2 requires 0.80 with every pitfall fault killed.
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. masking after the softmax | real weights no longer sum to 1; padding still shapes the normalization | test_masked_weights_are_exactly_zero_and_rows_sum_to_one, test_hand_example_weights_and_context (mutant s01) |
| 2. masking with instead of | an empty source gets uniform weights and reads padding | test_fully_masked_row_reads_nothing (mutant s02) |
| 3. softmax over the wrong axis | weights sum to 1 over the batch, not the source | test_masked_weights_are_exactly_zero_and_rows_sum_to_one, test_golden_torch (mutant s03) |
| 4. reading detached keys | the encoder gets no gradient through attention | test_gradcheck_every_input_and_parameter (mutant s04) |
| mask polarity inverted | only padding is read | test_padding_content_never_matters (mutant s05) |
no tanh | a linear score: the golden values disagree | test_golden_torch (mutant s06) |
<= in length_mask | one padding position is read per row | test_length_mask (mutant s07) |
| caching instead of | the query is added after the nonlinearity | test_hand_example_weights_and_context (mutant s08) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L0.4 | Linear holds , , , |
| Back | L0.2 | masked_fill, softmax, tanh, matmul give the backward for free |
| Back | L0.1 | Tensor, the type of queries, keys, and weights |
| Back | M06.3 | PCG32 initializes the layers when no rng is given |
| Forward | L4.3 | Luong attention keeps the same masked read and builds its masks with length_mask |
| Forward | L4.1 | the decoder calls attention(s_{t-1}, keys, mask, proj) before each step |
| Forward | L5.1 | scaled dot-product attention is this read with a cheaper score |
| Forward | L6.7 | the zoo’s attention variants of the seq2seq rows |
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
AdditiveAttention | torch.nn.MultiheadAttention, F.scaled_dot_product_attention | dot-product scores (no hidden layer), many heads, fused kernels | torch/nn/functional.py (multi_head_attention_forward) |
| masking to | attn_mask and key_padding_mask in torch | boolean and additive masks, broadcast over heads | PyTorch SDPA documentation |
project_keys | the KV cache | keys and values projected once and reused at every step of generation | L8.2 in this course; vLLM’s paged KV cache |
| additive scores | location-sensitive attention (Tacotron 2) | adds the previous weights as a feature so speech alignment moves forward | Chorowski et al. 2015 |