Multi-Head Attention Mechanisms
# Multi-Head Attention Mechanisms
## Introduction & Motivation
Multi-head attention: parallel attention with different representation subspaces. Improves model expressiveness. Applications: transformers, NLP, vision.
Motivation: Attend to different features simultaneously.
Applications: Transformer backbone, universal architecture.
---
## Core Concepts & Theory
### Multiple Heads
Parallel attention computations.
### Representation Subspaces
Different attention perspectives.
### Head Aggregation
Concatenate and project outputs.
### Position Bias
Add position-specific biases.
---
## Mathematical Formulation
Single Head:
$$ ext{Attention}(Q, K, V) = ext{softmax}(\frac{QK^T}{\sqrt{d_k}}) V$$
Multi-Head:
$$ ext{MultiHead}(Q,K,V) = ext{Concat}(h_1,...,h_h)W^O$$
Head:
$$h_i = ext{Attention}(QW_i^Q, KW_i^K, VW_i^V)$$
---
## Advanced Theory & Extensions
### Sparse Attention
Reduce attention connections.
### Low-Rank Attention
Approximate full attention.
### Dynamic Heads
Adaptive number of heads.
---
## Computational Considerations
Single head: O(T²·D).
Multi-head: O(h·T²·D/h) = O(T²·D).
Overhead: Linear projection costs.
---
## Practical Implementation Strategies
### Head Size
Typically d_model / num_heads.
### Scaling Factor
1/sqrt(d_k) for numerical stability.
### Dropout
Regularize attention weights.
---
## Benchmark Datasets & Evaluation
Machine Translation: BLEU scores.
Text Classification: Accuracy metrics.
Question Answering: Exact match, F1.
---
## Key Challenges & Limitations
### Attention Sparsity
Most attention to few positions.
### Computational Cost
Quadratic complexity in sequence length.
### Head Redundancy
Multiple heads may learn similar patterns.
---
## Hyperparameter Tuning
Number of heads: 8-16.
Head dimension: 64-128.
Dropout: 0.1-0.3.
---
## Real-World Applications & Case Studies
Machine Translation: Seq2Seq with attention.
Question Answering: Attend to relevant passages.
Image Classification: Attend to important regions.
---
## Integration with Other Methods
Multi-head + positional encoding; + feed-forward networks.
---
## Summary & Key Takeaways
Multi-head attention improves expressiveness.
Principles:
1. Parallel heads: Different perspectives.
2. Representation subspaces: Reduce dimension per head.
3. Concatenation: Combine information.
4. Scaling: Numerical stability.
5. Versatility: Universal architecture.
---
## Appendix: Practical Labs
### Lab 1: Multi-Head Computation
import numpy as np
def multi_head_attention(query, key, value, num_heads=8):
"""Compute multi-head attention"""
batch_size, seq_len, d_model = query.shape
d_k = d_model // num_heads
# Reshape for heads
query = query.reshape(batch_size, seq_len, num_heads, d_k).transpose(0, 2, 1, 3)
key = key.reshape(batch_size, seq_len, num_heads, d_k).transpose(0, 2, 1, 3)
value = value.reshape(batch_size, seq_len, num_heads, d_k).transpose(0, 2, 1, 3)
# Compute attention per head
scores = query @ key.transpose(0, 1, 3, 2) / np.sqrt(d_k)
attention = np.exp(scores) / np.sum(np.exp(scores), axis=-1, keepdims=True)
# Apply to values
output = attention @ value
# Reshape and concat
output = output.transpose(0, 2, 1, 3).reshape(batch_size, seq_len, d_model)
return output
np.random.seed(42)
q = np.random.randn(2, 10, 512)
k = np.random.randn(2, 10, 512)
v = np.random.randn(2, 10, 512)
out = multi_head_attention(q, k, v, num_heads=8)
assert out.shape == (2, 10, 512)
print("✓ Multi-head attention working")### Lab 2: Head Dimension
def compute_head_dimension(d_model, num_heads):
"""Compute dimension per head"""
if d_model % num_heads != 0:
raise ValueError("d_model must be divisible by num_heads")
return d_model // num_heads
d_k = compute_head_dimension(512, 8)
assert d_k == 64
print(f"✓ Head dimension: {d_k}")### Lab 3: Attention Visualization
import numpy as np
def visualize_attention_heads(attention_weights, num_heads=8):
"""Analyze attention patterns across heads"""
batch_size, num_h, seq_len, _ = attention_weights.shape
# Compute average attention distance
distances = []
for head in range(num_h):
attn = attention_weights[0, head] # First sample
# Compute attention distance
pos_indices = np.arange(seq_len)
for i in range(seq_len):
avg_dist = np.sum(attn[i] * np.abs(pos_indices - i))
distances.append(avg_dist)
return distances
np.random.seed(42)
attn = np.random.rand(2, 8, 10, 10)
attn = attn / attn.sum(axis=-1, keepdims=True)
dists = visualize_attention_heads(attn)
assert len(dists) == 80 # 8 heads × 10 positions
print("✓ Attention visualization working")### Lab 4: Head Analysis
import numpy as np
def analyze_head_similarity(attention_heads):
"""Compute similarity between attention heads"""
num_heads, seq_len, _ = attention_heads.shape
similarities = np.zeros((num_heads, num_heads))
for i in range(num_heads):
for j in range(num_heads):
# Cosine similarity
flat_i = attention_heads[i].flatten()
flat_j = attention_heads[j].flatten()
cos_sim = np.dot(flat_i, flat_j) / (np.linalg.norm(flat_i) * np.linalg.norm(flat_j) + 1e-8)
similarities[i, j] = cos_sim
return similarities
np.random.seed(42)
heads = np.random.rand(8, 10, 10)
sims = analyze_head_similarity(heads)
assert sims.shape == (8, 8)
print("✓ Head similarity analysis working")---