Skip to content

Luong attention and input feeding

ModuleL4.3 · build · Python · Pass 4 · 2 to 3 h, plus your graded tests (rung R5)
You buildpython/tinyllm/seq2seq/luong.py: LuongAttention with the dot, general, and concat scores (project_keys, scores, forward) and attentional, the state input feeding passes on; and your own oracle tests in python/tests/l4-3-luong/
Contractcourse/contracts/py/tinyllm/seq2seq/luong.pyi
Testscourse/tests/L4.3/test_luong.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L4.3/luong_torch.npz (course/oracle/L4.3/luong_torch.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL4.2 length_mask (and the masked read you wrote there) · L0.4 Linear, Module · L0.2 ops · L0.1 Tensor · M06.3 PCG32 · reading: M09.2 (or --ref-deps)
Used byL4.1 plugs it into the decoder with input feeding · L4.4 decodes a Luong seq2seq in its tests · later L6.7 the zoo’s attention rows
MilestoneMS-L4 (em_luong >= 0.97 on the dates task)
Optional depthLuong, Pham, and Manning, “Effective Approaches to Attention-based Neural Machine Translation” (EMNLP 2015), sections 3.1 and 3.3; Britz et al., “Massive Exploration of Neural Machine Translation Architectures” (2017), on additive vs multiplicative scores
  • Luong keeps Bahdanau’s masked read (L4.2) and changes two things: the score function, and when the decoder asks, after its recurrent step (test_hand_example_dot).
  • dot is h⋅ksh \cdot k_s, general the bilinear h⋅(Waks)h \cdot (W_a k_s) (with Wa=IW_a = I it is dot), concat a hidden layer on [h;ks][h ; k_s] (test_general_with_identity_is_dot, test_general_is_not_symmetric, test_concat_projection_splits_the_weight).
  • The attentional state h~=tanh⁡(Wc[c;h])\tilde h = \tanh(W_c [c ; h]), context first, is what the output layer reads and what input feeding passes to the next step (test_hand_example_attentional_order).
  • The plain dot score is not scaled here; scaling by 1/d1/\sqrt{d} is the transformer’s change in L5.1 (test_dot_is_not_scaled).
Terminal window
ol start L4.3 # stubs luong.py; prints your test path and rung (R5)
ol tests L4.3 # the course tests
# write your oracle tests in python/tests/l4-3-luong/, then:
ol check L4.3 # course tests and the mutation grade of your tests
ol diff L4.3 # after passing: your code against the reference

Bahdanau’s attention (L4.2) works, but every decoder step runs a hidden layer over every source position, and its query is the state from before the step, so the decoder chooses where to look before it has seen the token it was just given. Luong, Pham, and Manning simplified both: score with a dot product, and attend with the state just computed. They also noticed that the decoder forgets where it looked, and fixed that by feeding the attentional state back into the next step’s input. The dot score is the one the transformer keeps, so this module is the last step before L5.1.

SymbolMeaningType / shape
ddwidth of both the decoder state and the keysint
hhthe query: the decoder state after this stepfloat32[B, d]
ksk_sthe encoder output at source position ssfloat32[B, S, d]
WaW_ascore_proj.weight: [d,d][d, d] for general, [d,2d][d, 2d] for concatfloat32
vvv.weight, concat onlyfloat32[1, d]
WcW_ccombine.weightfloat32[d, 2d]
ese_s, asa_s, ccscore, weight, context as in L4.2
h~\tilde hthe attentional state tanh⁡(Wc[c;h])\tanh(W_c [c ; h])float32[B, d]

Bahdanau: read with st−1s_{t-1}, feed the context into the cell. Luong: run the cell first to get hth_t, then read with hth_t, then combine:

ht=cell([yt−1;h~t−1],ht−1),ct=∑sasks,h~t=tanh⁡(Wc[ct;ht]),logitst=Woh~t+bo.h_t = \mathrm{cell}([y_{t-1} ; \tilde h_{t-1}], h_{t-1}), \quad c_t = \sum_s a_s k_s, \quad \tilde h_t = \tanh(W_c [c_t ; h_t]), \quad \text{logits}_t = W_o \tilde h_t + b_o .

The decoder (L4.1) does the cell and the output layer; this module provides the scores, the read, and h~\tilde h. The class attribute query_from = "current" tells the decoder which order to use.

Scoreese_sParameters
doth⋅ksh \cdot k_snone
generalh⋅(Waks)h \cdot (W_a k_s)WaW_a [d,d][d, d]
concatv⋅tanh⁡(Wa[h;ks])v \cdot \tanh(W_a [h ; k_s])WaW_a [d,2d][d, 2d], vv

dot needs hh and ksk_s in the same space; general learns a bilinear form between them (with Wa=IW_a = I it is dot, and WaW_a and Wa⊤W_a^\top give different scores, so weights load only one way round); concat is Bahdanau’s network with one weight matrix on the concatenation. Masking, the softmax, and the context are exactly L4.2’s: padded scores −∞-\infty, exact zero weights, an empty row reads nothing.

For concat, Wa[h;ks]=Whh+WkksW_a [h ; k_s] = W_h h + W_k k_s where WhW_h is the first dd columns of WaW_a and WkW_k the last dd. So WkksW_k k_s is computed once per sentence (project_keys), and each step adds WhhW_h h. For general, project_keys is WaksW_a k_s; for dot, the keys themselves.

Without it, the decoder at step tt does not know what it attended to at step t−1t - 1; Luong’s input feeding concatenates h~t−1\tilde h_{t-1} to the next input embedding (h~0=0\tilde h_0 = 0). The model can then avoid translating the same source word twice, a coverage signal for free. It also makes the network deeper in time: the gradient of step tt reaches the attention of step t−1t - 1.

d=2d = 2, the dot score, h=(1,0)h = (1, 0), keys k1=(1,0)k_1 = (1, 0), k2=(0,1)k_2 = (0, 1), k3=(2,0)k_3 = (2, 0), the third padding.

  1. Scores: h⋅k=(1,0,2)h \cdot k = (1, 0, 2); after the mask (1,0,−∞)(1, 0, -\infty).
  2. Weights: e1/(e1+1)=0.731059e^1 / (e^1 + 1) = 0.731059 and 1/(e+1)=0.2689411 / (e + 1) = 0.268941, then 0.
  3. Context: c=0.731059(1,0)+0.268941(0,1)=(0.731059,0.268941)c = 0.731059 (1, 0) + 0.268941 (0, 1) = (0.731059, 0.268941).
  4. With Wc=[I  I]W_c = [I \; I] (so Wc[c;h]=c+hW_c [c ; h] = c + h): h~=tanh⁡(1.731059,0.268941)=(0.939181,0.262640)\tilde h = \tanh(1.731059, 0.268941) = (0.939181, 0.262640).
  5. With Wc=[I  0]W_c = [I \; 0] instead, h~=tanh⁡(c)=(tanh⁡0.731059,tanh⁡0.268941)\tilde h = \tanh(c) = (\tanh 0.731059, \tanh 0.268941): the context half is the first dd columns.

These are test_hand_example_dot and test_hand_example_attentional_order.

class LuongAttention(Module):
query_from = "current"
def __init__(self, d: int, score: Literal["dot", "general", "concat"], rng=None) -> None: ...
def project_keys(self, keys: Tensor) -> Tensor: ... # [B, S, d]
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
def attentional(self, query: Tensor, context: Tensor) -> Tensor: ... # tanh(W_c [c ; h])
TestKINDChecksWhy it matters downstream
test_hand_example_dotunitsection 3 steps 1 to 4you and the test agree on the read
test_hand_example_attentional_orderunitWc=[I  0]W_c = [I \; 0] gives tanh⁡(c)\tanh(c)context first, as Luong trained it
test_golden_torchgoldeneach score: weights, context, h~\tilde h, and every gradient against torchthe decoder trains through this
test_gradcheck_every_input_and_parametergradcheckfloat64 central differences, each score, a padded rowevery parameter learns
test_general_with_identity_is_dotpropertygeneral with Wa=IW_a = I equals dotthe relation between the scores
test_general_is_not_symmetricunitWaW_a, not Wa⊤W_a^\topcheckpoints load the right way round
test_concat_projection_splits_the_weightpropertyproject_keys and the score equal v⋅tanh⁡(Wa[h;k])v \cdot \tanh(W_a [h ; k])the per-sentence cache is exact
test_masked_weights_are_exactly_zeropropertyexact zeros, no gradient, an empty row reads nothingpadding is never read
test_dot_is_not_scaledunitscores exactly (1,0,2)(1, 0, 2)the scaled score belongs to L5.1
test_parameter_names_and_shapesuniteach score’s keys and shapes, combine last, query_fromthe safetensors keys of Luong checkpoints
test_validationboundaryunknown score, wrong widths, wrong maskwiring bugs fail loudly

As in L4.2, the oracle is the formula in numpy with the module’s own weights from state_dict(): all three scores, the masked softmax, the context, and h~\tilde h in float64. Add the section 3 numbers, your own finite-difference gradient check of the keys, an empty row, and the parameter names. Import only contract modules (tinyllm.seq2seq.luong, tinyllm.seq2seq.additive for length_mask, tinyllm.autograd.tensor, tinyllm.autograd.functional).

PitfallSymptomCaught by
1. Wc[h;c]W_c [h ; c] instead of Wc[c;h]W_c [c ; h]Luong-trained weights decode garbagetest_hand_example_attentional_order (mutant s01)
2. general with Wa⊤W_a^\topthe bilinear form is transposedtest_general_is_not_symmetric (mutant s02)
3. masking after the softmaxreal weights do not sum to 1test_masked_weights_are_exactly_zero (mutant s04)
4. scaling the dot score by 1/d1/\sqrt{d}every weight differs from Luong’stest_dot_is_not_scaled (mutant s07)
5. concat precomputing with the query half of WaW_athe cached keys are multiplied by the wrong columnstest_concat_projection_splits_the_weight (mutant s03)
masking with −109-10^9an empty row reads padding uniformlytest_masked_weights_are_exactly_zero (mutant s05)
h~\tilde h without the tanhthe attentional state is unboundedtest_hand_example_dot (mutant s06)
DirectionModuleHow it uses this
BackL4.2the masked read and length_mask
BackL0.4Linear holds WaW_a, vv, WcW_c
BackL0.2the op library gives the backward
BackL0.1Tensor
BackM06.3PCG32 for the default initialization
ForwardL4.1the decoder steps first, attends with the new state, and feeds h~\tilde h into the next input
ForwardL4.4its tests beam-search a Luong seq2seq, whose state carries feed
ForwardL5.1scaled dot-product attention: Luong’s dot divided by d\sqrt{d}, over many heads
ForwardL6.7the zoo’s Luong rows
Your pieceProduction equivalentWhat it addsWhere to look
dot scoreF.scaled_dot_product_attentionthe 1/d1/\sqrt{d} scale, heads, causal masks, fused kernelsL5.1; FlashAttention
input feedingthe attentional decoder of OpenNMTthe same input_feed option, with coverage and copy attentionOpenNMT-py onmt/decoders/decoder.py (InputFeedRNNDecoder)
general scorebilinear attention in readers and parsersthe biaffine scorer of dependency parsersDozat and Manning, “Deep Biaffine Attention” (2017)
local attentionLuong’s local-p attentionreads a window around a predicted position: cost per step independent of SSLuong et al. 2015, section 3.2; sliding-window attention in L7.7