Chuyển đến nội dung chính

Lesson 5: Attention Mechanism — Self-attention & Multi-head Attention

Attention mechanism from the root: Scaled Dot-Product Attention, Multi-head Attention, why Attention solves the problem of RNN. Code from scratch with PyTorch.

🧠 AI & ML — Lesson 4 Lesson 5: Attention Mechanism — Self-attention & Multi-head Attention

AI & LLM: From Basics to Advanced

Part 2: Transformer architecture

xdev.asia

Overview

Attention Mechanism is the heart of every modern LLM. This article will explain from visualization to mathematics, then code from scratch with PyTorch. After this article, you will understand why the phrase "Attention is All You Need" is so revolutionary.


1. Problem: Long-range Dependencies

Consider the sentence: "The animal didn't cross the street because it was too tired."

"It" here refers to "animal" or "street"? Humans immediately understand "animal" — but LSTM must "remember" the word "animal" through many sequential processing steps.

With Attention: the "it" token can directly look at every other token and calculate which tokens are most important to it — without an intermediary.


2. Core idea: Query, Key, Value

Attention is based on the database retrieval metaphor:

  • Query (Q): "What am I looking for?" — what the current token wants to know
  • Key (K): "What can I offer?" — each token "advertises" its content
  • Value (V): "This is my real content" — extracted information
Trực quan:
- Q của "it" hỏi: "tôi là gì?"
- K của "animal" trả lời tốt với query đó → attention score cao
- V của "animal" được đưa vào để tổng hợp nghĩa cho "it"

3. Scaled Dot-Product Attention

Recipe

$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$

Step by step explanation:

  1. QK^T — calculate dot product between query and every key → similarity scores
  2. / sqrt(d_k) — scale to avoid gradient vanishing when d_k is large
  3. softmax(...) — convert scores to probability distribution (sum = 1)
  4. * V — weighted sum of values according to attention weights

Why divide sqrt(d_k)?

When d_k is large (e.g. 64), dot products become very large → softmax is saturated → gradient is very small. Divide sqrt(d_k) keep variance stable:

# Nếu q, k ~ N(0,1), thì qk^T ~ N(0, d_k)
# Std = sqrt(d_k) → chia sqrt(d_k) để std = 1

Code: Scaled Dot-Product Attention

import torch
import torch.nn as nn
import torch.nn.functional as F
import math

def scaled_dot_product_attention(Q, K, V, mask=None):
    """
    Q: (batch, heads, seq_len, d_k)
    K: (batch, heads, seq_len, d_k)
    V: (batch, heads, seq_len, d_v)
    """
    d_k = Q.size(-1)

    # 1. Similarity scores: (batch, heads, seq_len_q, seq_len_k)
    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)

    # 2. Causal mask (dùng trong decoder — không nhìn tương lai)
    if mask is not None:
        scores = scores.masked_fill(mask == 0, float('-inf'))

    # 3. Attention weights
    attn_weights = F.softmax(scores, dim=-1)  # (B, H, T, T)

    # 4. Weighted sum of values
    output = torch.matmul(attn_weights, V)    # (B, H, T, d_v)

    return output, attn_weights

4. Multi-head Attention

Instead of running 1 attention with dimensional d_model, Multi-head Attention runs h attention heads in parallel, each head learning a different "aspect":

  • Head 1: syntactic relationship (subject-verb)
  • Head 2: coreference ("it" → "animal")
  • Head 3: semantic similarity
  • ...
MultiHead(Q,K,V) = Concat(head_1, ..., head_h) * W_O

head_i = Attention(Q*W_i^Q, K*W_i^K, V*W_i^V)

Code: Multi-head Attention

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super().__init__()
        assert d_model % num_heads == 0

        self.d_model = d_model
        self.num_heads = num_heads
        self.d_k = d_model // num_heads  # dimension per head

        # Linear projections
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def split_heads(self, x, batch_size):
        """(B, T, d_model) → (B, H, T, d_k)"""
        x = x.view(batch_size, -1, self.num_heads, self.d_k)
        return x.transpose(1, 2)

    def forward(self, query, key, value, mask=None):
        batch_size = query.size(0)

        # 1. Linear projections + split into heads
        Q = self.split_heads(self.W_q(query), batch_size)
        K = self.split_heads(self.W_k(key),   batch_size)
        V = self.split_heads(self.W_v(value), batch_size)

        # 2. Scaled dot-product attention
        x, attn_weights = scaled_dot_product_attention(Q, K, V, mask)

        # 3. Concatenate heads: (B, H, T, d_k) → (B, T, d_model)
        x = x.transpose(1, 2).contiguous()
        x = x.view(batch_size, -1, self.d_model)

        # 4. Final linear projection
        return self.W_o(x), attn_weights


# Test
d_model, num_heads, seq_len, batch = 512, 8, 20, 4
mha = MultiHeadAttention(d_model, num_heads)
x = torch.randn(batch, seq_len, d_model)
out, weights = mha(x, x, x)  # Self-attention: Q=K=V=x
print(f"Output: {out.shape}")        # (4, 20, 512)
print(f"Attn weights: {weights.shape}")  # (4, 8, 20, 20)

5. Self-attention vs Cross-attention

Self-attentionCross-attention
Q fromSame sequenceDecoder sequence
K, V fromSame sequenceEncoder output
Used inEncoder, DecoderDecoder
PurposeTokens are related to each otherDecoder attend encoder
# Self-attention: Q = K = V = encoder_output
self_attn_out = mha(encoder_out, encoder_out, encoder_out)

# Cross-attention: Q từ decoder, K/V từ encoder
cross_attn_out = mha(decoder_out, encoder_out, encoder_out)

6. Causal (Masked) Self-attention

In the decoder, tokens are not "forward-looking" when generated. We use causal mask:

def create_causal_mask(seq_len):
    """Upper triangular matrix = 0 (future tokens bị mask)"""
    mask = torch.tril(torch.ones(seq_len, seq_len))
    return mask  # 1=attend, 0=mask

# Ví dụ seq_len=4:
# [[1, 0, 0, 0],
#  [1, 1, 0, 0],
#  [1, 1, 1, 0],
#  [1, 1, 1, 1]]

7. Visualize Attention

import matplotlib.pyplot as plt
import seaborn as sns

def plot_attention(attn_weights, tokens, head=0):
    """Plot attention heatmap cho head cụ thể"""
    fig, ax = plt.subplots(figsize=(8, 6))
    # attn_weights: (batch, heads, seq, seq)
    weights = attn_weights[0, head].detach().numpy()
    sns.heatmap(weights, xticklabels=tokens, yticklabels=tokens,
                cmap='Blues', ax=ax)
    ax.set_title(f'Attention Head {head}')
    plt.tight_layout()
    plt.show()

# Ví dụ
tokens = ["The", "cat", "sat", "on", "the", "mat"]
# plot_attention(attn_weights, tokens, head=0)

8. Complexity: Attention vs RNN

Time ComplexitySpace ComplexitySequential Ops
Self-AttentionO(n² · d)O(n²)O(1)
RNNO(n · d²)O(d)O(n)
CNNO(k · n · d²)O(k · d)O(log n)

Attention advantages: O(1) sequential operations → fully parallelizable on GPU.

Disadvantage: O(n²) — with long sequences (n=4096+), quadratic cost is very expensive. This is the reason why Flash Attention and Sparse Attention research was born.


Summary

Attention cho phép:
✅ Mọi token attend trực tiếp mọi token khác
✅ Parallel processing (không tuần tự như RNN)
✅ Không có information bottleneck
✅ Học nhiều loại quan hệ khác nhau (multi-head)

Next post: Putting it all together — complete Transformer architecture with Encoder, Decoder, Positional Encoding and Feed-Forward layers.