🔗 Deep Learning · Lecture 23 of 38

The Attention Mechanism: Learning Where to Look

Attention lets a model compute a weighted focus over all input positions for each output. We derive Bahdanau and Luong attention, generalise to queries, keys and values, and see why it became the foundation of transformers.

When a human translator writes each word of a translation, they glance back at the relevant words of the source sentence rather than recalling a compressed memory of the whole thing. In 2014 Dzmitry Bahdanau, Kyunghyun Cho and Yoshua Bengio gave neural networks the same ability. Their attention mechanism removed the seq2seq bottleneck, dramatically improved translation of long sentences — and planted the idea that grew into the Transformer.

The idea#

Keep all encoder hidden states $\mathbf{h}_1, \dots, \mathbf{h}_S$. At each decoder step $t$, compute a separate context vector as a weighted average of them, with weights reflecting how relevant each source position is to the current output:

$$ \mathbf{c}_t = \sum_{j=1}^{S}\alpha_{tj}\,\mathbf{h}_j, \qquad \alpha_{tj} = \frac{\exp(e_{tj})}{\sum_{k=1}^{S}\exp(e_{tk})} $$

The alignment scores $e_{tj}$ measure the compatibility of decoder state $\mathbf{s}_{t-1}$ (or $\mathbf{s}_t$) with encoder state $\mathbf{h}_j$. The softmax turns scores into a probability distribution — a soft, differentiable "pointer" to the relevant input positions.

Scoring functions#

Additive (Bahdanau) attention:

$$ e_{tj} = \mathbf{v}^\top\tanh(\mathbf{W}_1\mathbf{s}_{t-1} + \mathbf{W}_2\mathbf{h}_j) $$

a small feedforward network scoring each pair.

Multiplicative (Luong) attention (Luong et al., 2015):

  • dot: $e_{tj} = \mathbf{s}_t^\top\mathbf{h}_j$
  • general: $e_{tj} = \mathbf{s}_t^\top\mathbf{W}\mathbf{h}_j$

Dot-product scoring is cheaper and can be computed for all positions at once with matrix multiplication.

Scaled dot-product (Transformer): $e = \mathbf{q}^\top\mathbf{k}/\sqrt{d_k}$. Scaling by $\sqrt{d_k}$ keeps the variance of scores around 1 when vectors have $d_k$ components of unit variance; without it, large dimensions produce huge scores, a saturated softmax and vanishing gradients.

Queries, keys and values#

The general abstraction, introduced explicitly in the Transformer:

  • a query $\mathbf{q}$ — what I am looking for;
  • keys $\mathbf{k}_j$ — what each position offers, used for matching;
  • values $\mathbf{v}_j$ — the content returned if selected.
$$ \text{Attention}(\mathbf{q}, \mathbf{K}, \mathbf{V}) = \sum_j\text{softmax}_j\left(\frac{\mathbf{q}^\top\mathbf{k}_j}{\sqrt{d_k}}\right)\mathbf{v}_j $$

Think of a soft dictionary lookup: instead of retrieving the single value whose key exactly matches, retrieve a blend of values weighted by key similarity. In seq2seq attention, the query is the decoder state and keys and values are both the encoder states.

Implementation#

python
import torch
import torch.nn as nn
import torch.nn.functional as F

def scaled_dot_product_attention(Q, K, V, mask=None):
    """Q: (B, Tq, d)  K: (B, Tk, d)  V: (B, Tk, dv)"""
    scores = Q @ K.transpose(-2, -1) / K.size(-1) ** 0.5       # (B, Tq, Tk)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float("-inf"))   # ignore padding / future
    weights = F.softmax(scores, dim=-1)
    return weights @ V, weights                                 # (B, Tq, dv)

class AdditiveAttention(nn.Module):
    def __init__(self, d_dec, d_enc, d_att=64):
        super().__init__()
        self.W1, self.W2 = nn.Linear(d_dec, d_att, bias=False), nn.Linear(d_enc, d_att, bias=False)
        self.v = nn.Linear(d_att, 1, bias=False)
    def forward(self, s, H, mask=None):          # s: (B, d_dec), H: (B, S, d_enc)
        e = self.v(torch.tanh(self.W1(s)[:, None] + self.W2(H))).squeeze(-1)   # (B, S)
        if mask is not None:
            e = e.masked_fill(mask == 0, float("-inf"))
        a = F.softmax(e, dim=-1)
        return (a[:, :, None] * H).sum(1), a    # context (B, d_enc), weights (B, S)

B, S, d = 2, 6, 16
H = torch.randn(B, S, d); s = torch.randn(B, d)
ctx, w = AdditiveAttention(d, d)(s, H)
print(ctx.shape, w.sum(-1))                        # weights sum to 1 for each example
out, w2 = scaled_dot_product_attention(torch.randn(B, 3, d), H, H)
print(out.shape, w2.shape)

In an attentional seq2seq decoder, the context vector $\mathbf{c}_t$ is concatenated with the decoder state (or input) to predict the next token.

Interpretability: alignment maps#

Plotting the attention weights $\alpha_{tj}$ as a matrix (target positions × source positions) reveals learned alignments: roughly diagonal for similar word orders, with swaps where languages reorder (e.g. adjective–noun order between English and French). These plots were a striking early glimpse into what neural translators learned.

Why attention was revolutionary#

  1. No bottleneck — the decoder accesses the full source at every step; long-sentence performance stopped collapsing.
  2. Short gradient paths — any output connects to any input through one attention step, easing learning of long-range dependencies.
  3. Content-based addressing — the model retrieves information by what it is, not where it is.
  4. Generality — attention works between any two sets of vectors: text and images (captioning), questions and documents (QA), a sequence and itself (self-attention).

In 2017, Vaswani et al. asked: if attention is so powerful, do we need recurrence at all? Their answer — "Attention Is All You Need" — introduced the Transformer, which we study in the NLP track.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

🔗 Deep Learning

Sequence-to-Sequence Models and the Encoder–Decoder Framework

Translation maps a sequence to another of different length. We build the encoder–decoder architecture, train it with teacher forcing, decode with greedy and beam search, and expose the bottleneck that motivated attention.

Intermediate⏱ 5 min#118
🔗 Deep Learning

Residual Connections: Why Very Deep Networks Became Trainable

Deeper plain networks can train worse than shallower ones. Residual connections fix this by learning corrections to the identity. We explain the degradation problem, the gradient highway, and variants from ResNet to transformers.

Intermediate⏱ 5 min#120
🔗 Deep Learning

Gated Recurrent Units (GRU) and Choosing a Recurrent Cell

The GRU simplifies the LSTM to two gates and one state while keeping long-term memory. We derive its equations, compare it with LSTMs empirically, and give practical guidance for recurrent models.

Intermediate⏱ 5 min#117