Attention Mechanisms and Transformers

The Encoder-Decoder Architecture


Suppose you are translating Der Hund, den die Katze jagte, bellte into English: The dog that the cat chased barked. Notice what has to happen. German puts jagte (chased) at the end of the relative clause and bellte (barked) at the very end of the sentence; English puts both verbs in different places entirely. To emit the first English word you already need to know that Der Hund is the subject of a verb you have not yet reached.

So the two sides of this problem have opposite requirements. Reading the source, you want every word to see every other word, in both directions, all at once — there is no reason to hide the end of a sentence you already have. Writing the target, you must not see the future, because at inference time the future does not exist yet: you are generating it one token at a time.

The transformer's answer is two stacks with different rules. The encoder reads with unrestricted bidirectional attention. The decoder writes with causal attention plus a third mechanism that reaches back into the encoder's output. Getting the interaction between them right — especially the difference between how it trains and how it runs — is where most implementations go wrong.

Where the German verb reaches the English sentenceEncoder block: self-attention, then FFNEncoder output: one vector per source tokenDecoder: causalself-attention over the output so farDecoder: cross-attention onto the encoder vectorsDecoder: FFN, then a projection to the vocabulary
Cross-attention is the only place the two languages meet, which is why the verb can move position freely.

The shape of the whole thing

Text
          SOURCE                                TARGET (shifted right)             |                                          |      token embeddings                           token embeddings      + positional enc.                          + positional enc.             |                                          |   +---------v---------+                     +----------v-----------+   |  ENCODER LAYER 1  |                     |   DECODER LAYER 1    |   |  self-attn (full) |                     |  self-attn (causal)  |   |  feed-forward     |        +----------> |  cross-attn (K,V)    |   +---------+---------+        |            |  feed-forward        |             |                  |            +----------+-----------+            ...  (x N)          |                      ... (x N)             |                  |                       |   +---------v---------+        |            +----------v-----------+   |  ENCODER LAYER N  |--------+            |   DECODER LAYER N    |   +-------------------+   memory            +----------+-----------+                                                        |                                              linear -> softmax                                                        |                                             next-token probabilities

The base configuration from the original transformer paper: N=6N = 6 layers each side, dmodel=512d_{model} = 512, h=8h = 8 heads, dff=2048d_{ff} = 2048, roughly 65 million parameters. Every encoder layer has the same shape as every other; the same for decoder layers. Only the learned weights differ.

One detail that catches people: every encoder layer output does not go to the matching decoder layer. The final encoder output — often called the memory — goes to all decoder layers. There are NN separate cross-attention modules but they all read the same memory tensor.

The encoder stack

What it is for

The encoder turns a sequence of token embeddings into a sequence of contextual representations, one per input position, same length in and out. Position ii of the output is a vector representing "the token at position ii, as understood in the context of this whole sentence."

Bidirectionality is the point. Consider bank in these two sentences:

SentenceDisambiguating wordWhere it sits
The bank approved my loanloan3 tokens after
The bank of the river was muddyriver3 tokens after

Both disambiguators come after the ambiguous word. A left-to-right model must commit to a representation of bank before seeing either. The encoder has no such constraint, so its representation of bank is built with loan or river already in view.

Sub-layer 1: multi-head self-attention

Every position attends to every position, itself included. For a source of length TsrcT_{src} this produces a Tsrc×TsrcT_{src} \times T_{src} attention matrix per head. The only masking is padding — blocking the [PAD] columns of short sequences in a batch — and never a causal mask.

A concrete size check. Batch 32, Tsrc=128T_{src} = 128, 8 heads: the score tensor is 32×8×128×128=4,194,30432 \times 8 \times 128 \times 128 = 4{,}194{,}304 entries, 8 MiB in fp16. Comfortable. At Tsrc=1024T_{src} = 1024 the same tensor is 512 MiB per layer, which is why long source documents get expensive fast.

Sub-layer 2: the position-wise feed-forward network

FFN(x)=max⁡(0, xW1+b1)W2+b2\mathrm{FFN}(x) = \max(0,\ xW_1 + b_1)W_2 + b_2

Two linear layers with a nonlinearity between them, applied independently and identically to each position. "Position-wise" means exactly that: it is the same small network run at every position with no mixing between them. If you implemented it with a 1-D convolution of kernel size 1, you would get the same thing.

The dimensions matter. W1W_1 maps 512→2048512 \to 2048 and W2W_2 maps 2048→5122048 \to 512 — it expands four-fold and comes back. Count the parameters:

ComponentShapeParameters
W1W_1512×2048512 \times 20481,048,576
b1b_1204820482,048
W2W_22048×5122048 \times 5121,048,576
b2b_2512512512
FFN total2,099,712
Multi-head attention total4 matrices of 512×512512 \times 5121,048,576

The feed-forward network holds twice as many parameters as the attention block. Two thirds of a transformer's parameters are in the FFNs. This surprises people who think of transformers as "attention models".

What is it doing? Attention mixes information across positions but is, per position, a linear combination of value vectors. Without the FFN the whole stack would collapse into something close to a single linear map plus softmax mixing. The FFN supplies the per-position nonlinear processing — the place where a mixed representation gets transformed into something new. A useful mental model: attention decides what to look at, the FFN decides what to think about it.

Encoder in code

Python
import torchimport torch.nn as nnclass FeedForward(nn.Module):    def __init__(self, d_model, d_ff, dropout=0.1):        super().__init__()        self.net = nn.Sequential(            nn.Linear(d_model, d_ff),            nn.ReLU(),            nn.Dropout(dropout),            nn.Linear(d_ff, d_model),        )    def forward(self, x):        return self.net(x)class EncoderLayer(nn.Module):    def __init__(self, d_model, num_heads, d_ff, dropout=0.1):        super().__init__()        self.self_attn = MultiHeadAttention(d_model, num_heads, dropout)        self.ff = FeedForward(d_model, d_ff, dropout)        self.norm1 = nn.LayerNorm(d_model)        self.norm2 = nn.LayerNorm(d_model)        self.drop = nn.Dropout(dropout)    def forward(self, x, src_mask=None):        # residual + norm around each sub-layer        a, attn = self.self_attn(x, x, x, mask=src_mask)        x = self.norm1(x + self.drop(a))        f = self.ff(x)        x = self.norm2(x + self.drop(f))        return x, attnclass Encoder(nn.Module):    def __init__(self, num_layers, d_model, num_heads, d_ff, dropout=0.1):        super().__init__()        self.layers = nn.ModuleList(            EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)        )    def forward(self, x, src_mask=None):        maps = []        for layer in self.layers:            x, attn = layer(x, src_mask)            maps.append(attn)        return x, maps          # x is the "memory"

Note that self_attn(x, x, x) — the same tensor three times — is what makes it self-attention.

The decoder stack

Three sub-layers, not two

Sub-layerQuery fromKey/Value fromMaskJob
1. Masked self-attentiondecoderdecodercausal + paddingMake the partial output coherent with itself
2. Cross-attentiondecoderencoder memorysource padding onlyPull in the relevant part of the source
3. Feed-forward———Per-position nonlinear processing

Sub-layer 1: causal self-attention

Position ii may attend only to positions ≤i\le i. This is enforced by adding −∞-\infty to the scores of future positions before softmax.

Work it through for a four-token target at position 3 (0-indexed position 2). Scaled scores [2.0, 1.0, 3.0, 5.0][2.0,\ 1.0,\ 3.0,\ 5.0], with position 4 in the future:

pos 1pos 2pos 3pos 4sum
Score, unmasked2.01.03.05.0
ese^{s}7.3892.71820.086148.413178.606
Weight, unmasked0.0410.0150.1120.8311.000
Score, masked2.01.03.0−∞-\infty
ese^{s}7.3892.71820.086030.193
Weight, masked0.2450.0900.66501.000

Without the mask, 83% of the attention at this position goes to the token the model is being asked to predict. Training loss would drop near zero and generation would produce noise, because at inference that column is empty. Note also that the masked weights renormalise to sum to 1 over the legal positions — which is why you must mask the scores rather than zeroing the weights afterwards.

A causal mask is not a regulariser or a heuristic. It is the constraint that makes training conditions match inference conditions; break it and the two diverge completely.

Sub-layer 2: cross-attention

This is where the two stacks meet, and its asymmetry is the thing to internalise:

Q=decoder state⋅WQ,K=memory⋅WK,V=memory⋅WVQ = \text{decoder state} \cdot W^Q, \qquad K = \text{memory} \cdot W^K, \qquad V = \text{memory} \cdot W^V

The attention matrix is Ttgt×TsrcT_{tgt} \times T_{src} — rectangular, not square. Row ii says: while producing target token ii, how much did I draw on each source token?

Take the German example. Source tokens: Der Hund , den die Katze jagte , bellte. When the decoder emits barked, a trained model's cross-attention row typically looks like:

SourceDerHund,dendieKatzejagte,bellte
Weight0.040.110.010.010.010.020.030.010.76

Sum: 0.04+0.11+0.01+0.01+0.01+0.02+0.03+0.01+0.76=1.000.04+0.11+0.01+0.01+0.01+0.02+0.03+0.01+0.76 = 1.00. The mass sits on bellte, with a secondary lobe on Hund — the subject it agrees with. Nobody supplied word alignments; this is learned from sentence pairs alone. This is also why cross-attention maps are the single most useful diagnostic when a translation model misbehaves.

No causal mask here. The whole source exists from the start, so hiding parts of it would only remove information. The only mask is source padding.

Decoder in code

Python
class DecoderLayer(nn.Module):    def __init__(self, d_model, num_heads, d_ff, dropout=0.1):        super().__init__()        self.self_attn  = MultiHeadAttention(d_model, num_heads, dropout)        self.cross_attn = MultiHeadAttention(d_model, num_heads, dropout)        self.ff = FeedForward(d_model, d_ff, dropout)        self.norm1 = nn.LayerNorm(d_model)        self.norm2 = nn.LayerNorm(d_model)        self.norm3 = nn.LayerNorm(d_model)        self.drop  = nn.Dropout(dropout)    def forward(self, x, memory, tgt_mask=None, src_mask=None):        a, self_map = self.self_attn(x, x, x, mask=tgt_mask)        x = self.norm1(x + self.drop(a))        # Q from the decoder, K and V from the encoder memory        c, cross_map = self.cross_attn(x, memory, memory, mask=src_mask)        x = self.norm2(x + self.drop(c))        f = self.ff(x)        x = self.norm3(x + self.drop(f))        return x, self_map, cross_map

The argument order (x, memory, memory) in the cross-attention call is the whole architecture in one line. Swap it to (memory, x, x) and you get a shape error if the lengths differ — or, if source and target happen to be the same length, no error at all and a model that quietly learns the wrong thing.

Training and inference are different programs

Teacher forcing

During training the correct target is known, so the whole target sequence is fed in at once and every position is predicted in parallel. The input is the target shifted right by one, with a start token prepended.

Position12345
Decoder input<sos>Thedogbarked.
Target (label)Thedogbarked.<eos>

Position 3 receives dog as input, can attend to positions 1–3 only, and must predict barked. Because the causal mask prevents it from seeing positions 4 and 5, all five predictions can be computed in a single forward pass without any of them cheating. This is the entire reason transformers train faster than recurrent models: teacher forcing plus causal masking turns a sequential problem into a parallel one.

Python
decoder_input = target[:, :-1]     # drop the final tokenlabels        = target[:, 1:]      # drop the start tokenlogits = model(source, decoder_input)                     # (B, T-1, vocab)loss = F.cross_entropy(    logits.reshape(-1, logits.size(-1)),    labels.reshape(-1),    ignore_index=PAD_ID,)

The off-by-one here is the single most common bug in sequence-to-sequence code. If you forget the shift and feed the full target, position ii sees token ii and is asked to predict token ii — the model learns the identity function, training loss drops to near zero within an epoch, and generation emits the start token forever.

Autoregressive generation

At inference there is no target. You generate one token, append it, and run again:

Python
@torch.no_grad()def greedy_decode(model, source, src_mask, max_len, sos_id, eos_id):    memory, _ = model.encode(source, src_mask)        # encoder runs ONCE    ys = torch.full((source.size(0), 1), sos_id, device=source.device)    for _ in range(max_len - 1):        tgt_mask = causal_mask(ys.size(1), ys.device)        out = model.decode(ys, memory, tgt_mask, src_mask)        logits = model.generator(out[:, -1])          # only the last position matters        next_tok = logits.argmax(-1, keepdim=True)        ys = torch.cat([ys, next_tok], dim=1)        if (next_tok == eos_id).all():            break    return ys
Training (teacher forcing)Inference (autoregressive)
Decoder forward passes for a length-TT output1TT
Encoder forward passes11 (reuse the memory — recomputing it per step is a common waste)
Decoder input at step iiGround-truth token i−1i-1Model's own previous output
ErrorsDo not propagateCompound — a wrong token becomes the next step's input

That last row is exposure bias: the model is trained on perfect prefixes and deployed on its own imperfect ones. A model with 95% per-token accuracy under teacher forcing has, naively, a 0.9520≈0.360.95^{20} \approx 0.36 chance of producing a 20-token sequence with no errors, and each error shifts the input distribution further from anything it saw in training. Beam search mitigates this by keeping several hypotheses alive rather than committing to the argmax at every step.

The TT decoder passes are also where inference cost lives. Each pass recomputes keys and values for every previous position, which is O(T2)O(T^2) work over the whole generation. A KV cache stores them instead: at step ii you compute K and V for the new token only and concatenate onto the cache, reducing the total to O(T)O(T) recomputation. Every production inference stack does this.

Encoder and decoder side by side

Encoder layerDecoder layer
Sub-layers23
Self-attention maskingPadding only — fully bidirectionalCausal + padding
Attends to the other stackNoYes, via cross-attention
Output at position ii depends onAll source positionsTarget positions ≤i\le i, and all source positions
Parameters per layer (d=512d=512, dff=2048d_{ff}=2048)≈3.15\approx 3.15 M≈4.20\approx 4.20 M
Can be run in parallel over positions at inferenceYesNo — one step per token

You do not always need both stacks

FamilyStructureAttention patternNatural tasksExamples
Encoder-decoderBoth stacksBidirectional over input, causal over output, cross between themInput and output are different sequences: translation, summarisation, speech recognitionOriginal transformer, T5, BART
Encoder-onlyEncoder plus a task headFully bidirectionalInput needs understanding, output is a label or a span: classification, NER, retrievalBERT, RoBERTa, DeBERTa
Decoder-onlyDecoder without cross-attentionCausal throughoutContinue a sequence: language modelling, chat, code completionGPT family, LLaMA, Mistral

Decoder-only models are worth a second look, because they have become dominant and the reason is not obvious. If you delete cross-attention from a decoder layer you get exactly an encoder layer with a causal mask. So a decoder-only model handles source-to-target tasks by concatenating them into one sequence:

Text
Translate to English: Der Hund bellte => The dog barked

The causal mask means the source tokens can only see earlier source tokens — strictly less context than a bidirectional encoder would give them. That is a genuine disadvantage on understanding tasks. What decoder-only models buy in exchange is uniformity: one stack, one objective (next-token prediction), any text corpus as training data, no need for paired input-output examples. At scale that trade has generally been worth it.

What this means when you build something

Pick the architecture from the shape of the task before anything else. If your output is a label, an encoder-only model is smaller, faster, and bidirectional — do not reach for a generative model to do classification. If your output is free text conditioned on a specific input document, an encoder-decoder gives the input full bidirectional context and keeps the two sequences cleanly separated. If your task is open-ended continuation, decoder-only.

Then check these four things, in this order, because each has a characteristic and misleading symptom:

SymptomCauseFix
Training loss collapses to near zero in one epoch; generation emits one token repeatedlyDecoder input not shifted right — the model is copying its inputFeed target[:, :-1], score against target[:, 1:]
Training loss excellent, generation incoherentCausal mask missing or applied after softmaxMask scores with −∞-\infty before softmax; verify output at ii is unchanged when token j>ij > i changes
Model ignores the source entirely and produces fluent but unrelated outputCross-attention arguments swapped, or memory not passed to every decoder layerConfirm cross_attn(x, memory, memory); plot a cross-attention map and check it is not uniform
Inference is far slower than expectedEncoder re-run every decode step, or no KV cacheEncode once outside the loop; cache decoder keys and values

And keep the parameter arithmetic in mind when you size a model. Doubling dffd_{ff} from 2048 to 4096 adds about 2.1 M parameters per layer — more than doubling the head count would, which adds none at all. If you need capacity, the feed-forward width is usually where it should go.