Skip to content

Bahdanau additive attention

ModuleL4.2 · build · Python · Pass 4 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/seq2seq/additive.py: length_mask and AdditiveAttention (project_keys, scores, forward); and your own oracle tests in python/tests/l4-2-additive/
Contractcourse/contracts/py/tinyllm/seq2seq/additive.pyi
Testscourse/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
NeedsL0.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 byL4.3 builds masks with length_mask · L4.1 plugs AdditiveAttention into the decoder · later L6.7 the zoo’s attention rows
MilestoneMS-L4 (em_bahdanau >= 0.97 on the dates task)
Optional depthBahdanau, 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
  • 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 −∞-\infty, 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 Uks+bU k_s + b 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).
Terminal window
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 tests
ol diff L4.2 # after passing: your code against the reference

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.

SymbolMeaningType / shape
BB, SSbatch size, source length (padded)int
qqthe query: the decoder state that asksfloat32[B, Dq]
ksk_sthe key at source position ss: an encoder output, also the value readfloat32[B, S, Dk]
ℓb\ell_bthe real length of source bbint
mbsm_{bs}mask: true when s<ℓbs < \ell_bbool[B, S]
WW, UU, bb, vvparameters query.weight [A,Dq][A, Dq], key.weight [A,Dk][A, Dk], key.bias [A][A], v.weight [1,A][1, A]float32
AAd_attn, the width of the scoring networkint
ese_sthe score of position ssfloat
asa_sthe attention weight of position ssfloat in [0,1][0, 1]
ccthe context: ∑sasks\sum_s a_s k_sfloat32[B, Dk]

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 qq.

Bahdanau scores each position with a one-hidden-layer network on the pair (q,ks)(q, k_s):

es=v⊤tanh⁡(Wq+Uks+b),a=softmax⁡(e),c=∑sasks.e_s = v^\top \tanh(W q + U k_s + b), \qquad a = \operatorname{softmax}(e), \qquad c = \sum_s a_s k_s .

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 qq (the decoder learns where to look), the keys (the encoder learns what to offer), and W,U,b,vW, U, b, v. The context has the keys’ width DkD_k, not AA: 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 ss itself: permute the keys and the weights permute with them. Order has to come from the encoder.

A batch pads every source to the longest length SS. Padded positions must never be read. The rule is to set their scores to −∞-\infty before the softmax: e−∞=0e^{-\infty} = 0, so their weights are exactly 0 and the real weights still sum to 1. length_mask(lengths, S) builds mbs=(s<ℓb)m_{bs} = (s < \ell_b). 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 −∞-\infty; 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.

Uks+bU k_s + b involves only the keys, which are fixed for the whole sentence, while WqW q changes at every decoder step. Computing Uks+bU k_s + b once (project_keys) and passing it to every step as proj saves S⋅A⋅DkS \cdot A \cdot D_k multiply-adds per step; Bahdanau’s appendix points this out. The result must be identical to recomputing.

Every size 1 (Dq=Dk=A=1D_q = D_k = A = 1), W=U=v=1W = U = v = 1, b=0b = 0. Query q=0q = 0, keys k=(0,1,2)k = (0, 1, 2), and the third position is padding (mask [True, True, False]).

  1. Scores before the mask: tanh⁡(0+0)=0\tanh(0 + 0) = 0, tanh⁡(1)=0.761594\tanh(1) = 0.761594, tanh⁡(2)=0.964028\tanh(2) = 0.964028. The padded one would be the largest.
  2. After the mask: (0,0.761594,−∞)(0, 0.761594, -\infty).
  3. Softmax: e0=1e^0 = 1, e0.761594=2.141688e^{0.761594} = 2.141688, e−∞=0e^{-\infty} = 0; the sum is 3.1416883.141688, so a=(0.318300,0.681700,0)a = (0.318300, 0.681700, 0).
  4. Context: c=0.318300⋅0+0.681700⋅1+0⋅2=0.681700c = 0.318300 \cdot 0 + 0.681700 \cdot 1 + 0 \cdot 2 = 0.681700.

The padded key, the biggest value in the row, contributes nothing. This is test_hand_example_weights_and_context.

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, weights
TestKINDChecksWhy it matters downstream
test_hand_example_weights_and_contextunitsection 3 to float32 toleranceyou and the test agree on the formula
test_length_maskunits < lengths[b], bad lengths raisethe only place lengths enter attention
test_golden_torchgoldencontext, weights, and gradients of the query, keys, and all four parameters against torch, lengths 5, 3, 1your encoder and decoder train through this
test_gradcheck_every_input_and_parametergradcheckfloat64 central differences with a padded rowboth sides of the model learn
test_masked_weights_are_exactly_zero_and_rows_sum_to_onepropertyexact zeros, positive real weights, rows sum to 1padding is never read
test_padding_content_never_matterspropertygarbage in padded keys changes nothing bitwise; their gradient is 0a batch padded differently translates the same
test_fully_masked_row_reads_nothingboundaryzeros, no NaN, no gradientempty inputs in a batch
test_precomputed_keys_equal_recomputedpropertyproj passed or recomputed, bitwise equalL4.1 passes it at every step
test_permuting_the_source_permutes_the_weightspropertyweights permute, context unchangedorder comes from the encoder, not from attention
test_parameter_names_and_shapesunitkeys, shapes, order, query_from, seeded initthe safetensors keys of seq2seq checkpoints
test_validationboundarywrong widths, missing source axis, wrong or float mask, wrong projwiring bugs fail loudly

The oracle for your tests is the formula itself, written out in numpy with the module’s own weights read from state_dict(): compute ee, 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.

PitfallSymptomCaught by
1. masking after the softmaxreal weights no longer sum to 1; padding still shapes the normalizationtest_masked_weights_are_exactly_zero_and_rows_sum_to_one, test_hand_example_weights_and_context (mutant s01)
2. masking with −109-10^9 instead of −∞-\inftyan empty source gets uniform weights and reads paddingtest_fully_masked_row_reads_nothing (mutant s02)
3. softmax over the wrong axisweights sum to 1 over the batch, not the sourcetest_masked_weights_are_exactly_zero_and_rows_sum_to_one, test_golden_torch (mutant s03)
4. reading detached keysthe encoder gets no gradient through attentiontest_gradcheck_every_input_and_parameter (mutant s04)
mask polarity invertedonly padding is readtest_padding_content_never_matters (mutant s05)
no tanha linear score: the golden values disagreetest_golden_torch (mutant s06)
<= in length_maskone padding position is read per rowtest_length_mask (mutant s07)
caching tanh⁡(Uk+b)\tanh(U k + b) instead of Uk+bU k + bthe query is added after the nonlinearitytest_hand_example_weights_and_context (mutant s08)
DirectionModuleHow it uses this
BackL0.4Linear holds WW, UU, bb, vv
BackL0.2masked_fill, softmax, tanh, matmul give the backward for free
BackL0.1Tensor, the type of queries, keys, and weights
BackM06.3PCG32 initializes the layers when no rng is given
ForwardL4.3Luong attention keeps the same masked read and builds its masks with length_mask
ForwardL4.1the decoder calls attention(s_{t-1}, keys, mask, proj) before each step
ForwardL5.1scaled dot-product attention is this read with a cheaper score
ForwardL6.7the zoo’s attention variants of the seq2seq rows
Your pieceProduction equivalentWhat it addsWhere to look
AdditiveAttentiontorch.nn.MultiheadAttention, F.scaled_dot_product_attentiondot-product scores (no hidden layer), many heads, fused kernelstorch/nn/functional.py (multi_head_attention_forward)
masking to −∞-\inftyattn_mask and key_padding_mask in torchboolean and additive masks, broadcast over headsPyTorch SDPA documentation
project_keysthe KV cachekeys and values projected once and reused at every step of generationL8.2 in this course; vLLM’s paged KV cache
additive scoreslocation-sensitive attention (Tacotron 2)adds the previous weights as a feature so speech alignment moves forwardChorowski et al. 2015