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")

---

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account