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:
QK^T— calculate dot product between query and every key → similarity scores/ sqrt(d_k)— scale to avoid gradient vanishing when d_k is largesoftmax(...)— convert scores to probability distribution (sum = 1)* 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-attention | Cross-attention | |
|---|---|---|
| Q from | Same sequence | Decoder sequence |
| K, V from | Same sequence | Encoder output |
| Used in | Encoder, Decoder | Decoder |
| Purpose | Tokens are related to each other | Decoder 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 Complexity | Space Complexity | Sequential Ops | |
|---|---|---|---|
| Self-Attention | O(n² · d) | O(n²) | O(1) |
| RNN | O(n · d²) | O(d) | O(n) |
| CNN | O(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.