PyTorch Attention Mechanism
The attention mechanism is one of the most important concepts in deep learning.
The attention mechanism enables models to learn to "focus on" the most relevant parts of the input, achieving great success in fields such as natural language processing and computer vision.
This section details the core principles of the attention mechanism, PyTorch implementations, and various attention variants.
Applicable version:The code in this article is written based on PyTorch 2.0+.
nn.MultiheadAttentionofbatch_firstThe parameter was introduced in PyTorch 1.9; earlier versions require manual dimension handling.
1. Attention Mechanism Basics
1.1 Why Attention Is Needed
Traditional sequence-to-sequence (Seq2Seq) models have a fundamental problem:The Encoder needs to compress all information into a fixed-length vector.For long sequences, this vector becomes an information bottleneck—the longer the sentence, the more severe the information loss.
The core idea of the attention mechanism is to allow the Decoder to "see" all hidden states of the Encoder when generating each output, and dynamically assign different attention weights based on the current context. This is like a human translator looking back at the corresponding part of the source text each time they translate a word.
1.2 The Essence of Attention Mechanism
The attention mechanism can be viewed as aweighted sumoperation. Given a query (Query), key (Key), and value (Value), weights are assigned by computing the similarity between the Query and each Key, and then a weighted sum of the Values is computed:
\[\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V\]$$
where:
- Q(Query): query vector, representing "what information I am looking for"
- K(Key): key vector, indicating "what information I have here" (used for matching)
- V(Value): value vector, representing "what I can actually provide"
- $$d_k$$: dimension of Key, $$\sqrt{d_k}$$ is used for scaling to prevent dot-product values from being too large, which would cause vanishing softmax gradients
Why divide by $$\sqrt{d_k}$$? When $$d_k$$ is large, the variance of the dot product of Q and K grows linearly with the dimension, causing the softmax input values to be too large and the gradients to approach zero. After scaling, the variance is restored to 1, ensuring normal gradient flow.
Example
import torch.nn as nn
import torch.nn.functional as F
import math
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Scaled dot-product attention
Parameters:
Q: Query tensor [batch, n_heads, seq_len_q, d_k]
K: Key tensor [batch, n_heads, seq_len_k, d_k]
V: Value tensor [batch, n_heads, seq_len_v, d_v]
mask: Mask tensor; masked positions are set to False/0
Returns:
output: Attention output [batch, n_heads, seq_len_q, d_v]
attention_weights: Attention weights [batch, n_heads, seq_len_q, seq_len_k]
"""
d_k = Q.size(-1)
# 1. Compute QK^T / sqrt(d_k)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# 2. Apply mask (optional)
# Set masked positions to a very small value; after softmax they approach 0
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 3. Softmax normalization → attention weights
attention_weights = F.softmax(scores, dim=-1)
# 4. Weighted sum
output = torch.matmul(attention_weights, V)
return output, attention_weights
# Test
batch_size, n_heads = 2, 4
seq_len_q, seq_len_k = 5, 6
d_k, d_v = 8, 16
Q = torch.randn(batch_size, n_heads, seq_len_q, d_k)
K = torch.randn(batch_size, n_heads, seq_len_k, d_k)
V = torch.randn(batch_size, n_heads, seq_len_k, d_v)
output, attn_weights = scaled_dot_product_attention(Q, K, V)
print(f"Output shape: {output.shape}") # [2, 4, 5, 16]
print(f"Attention weights shape: {attn_weights.shape}") # [2, 4, 5, 6]
print(f"Weight row sum (should be 1.0): {attn_weights[0, 0, 0].sum().item():.4f}")
The essence of attention weights is a probability distribution—after softmax, each row sums to 1. The larger the weight, the more the model "focuses" on that position.
2. PyTorch Attention Module
2.1 Multi-Head Attention
Multi-Head AttentionIt allows the model to simultaneously attend to information from different representation subspaces at different positions. It projects Q, K, and V through different linear projections into multiple subspaces, computes attention independently in each subspace, and finally concatenates the results and applies another linear transformation.
This is like having multiple people examine the same text from different perspectives simultaneously—some focus on grammatical structure, some on semantic relationships, and some on long-distance dependencies. The multi-head mechanism enables the model to capture richer patterns.
$$\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h)W^O$$
$$\text{where head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$$
Example
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0, "d_model must be divisible by n_heads"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads # Dimension of each head
# Linear projections for Q, K, V (compute all heads at once, more efficient)
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)
# Output projection
self.w_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, query, key, value, mask=None):
"""
Parameters:
query: [batch, seq_len_q, d_model]
key: [batch, seq_len_k, d_model]
value: [batch, seq_len_k, d_model]
mask: [batch, 1, 1, seq_len_k] or [batch, 1, seq_len_q, seq_len_k]
"""
batch_size = query.size(0)
# 1. Linear projection, then split into heads
# [batch, seq_len, d_model] → [batch, seq_len, n_heads, d_k]
# → [batch, n_heads, seq_len, d_k]
Q = self.w_q(query).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = self.w_k(key).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = self.w_v(value).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 2. Scaled dot-product attention
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 3. Weighted sum
context = torch.matmul(attn_weights, V)
# 4. Merge heads: [batch, n_heads, seq_len, d_k] → [batch, seq_len, d_model]
context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
# 5. Output projection
output = self.w_o(context)
return output, attn_weights
# Test
d_model, n_heads = 128, 8
seq_len, batch = 10, 4
layer = MultiHeadAttention(d_model, n_heads)
query = torch.randn(batch, seq_len, d_model)
key = torch.randn(batch, seq_len, d_model)
value = torch.randn(batch, seq_len, d_model)
output, attn_weights = layer(query, key, value)
print(f"Output shape: {output.shape}") # [4, 10, 128]
print(f"Attention weights shape: {attn_weights.shape}") # [4, 8, 10, 10]
print(f"Each head dimension d_k = {d_model // n_heads}")
2.2 PyTorch Built-in MultiheadAttention
PyTorch provides a highly optimizednn.MultiheadAttention, which uses fused kernels at the underlying level and is faster than manual implementations in most scenarios. It is recommended for use in production environments.
Example
import torch.nn as nn
class TransformerAttention(nn.Module):
"""Self-attention layer using PyTorch's built-in MultiheadAttention"""
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
self.attention = nn.MultiheadAttention(
embed_dim=d_model,
num_heads=n_heads,
dropout=dropout,
batch_first=True, # Input format is [batch, seq, features]
# PyTorch 2.0+ can enable the Flash Attention backend:
# attn_implementation="flash_attention_2" # requires installing flash-attn
)
self.layernorm = nn.LayerNorm(d_model)
def forward(self, x, key_padding_mask=None):
"""
Self-attention: Q, K, V are all x
Parameters:
x: [batch, seq_len, d_model]
key_padding_mask: [batch, seq_len], True indicates that position is padding
"""
attn_output, attn_weights = self.attention(
x, x, x, # self-attention
key_padding_mask=key_padding_mask
)
# Pre-Norm residual connection (more stable training than Post-Norm)
output = self.layernorm(x + attn_output)
return output, attn_weights
# Test
d_model, n_heads = 128, 8
seq_len, batch = 10, 4
model = TransformerAttention(d_model, n_heads)
x = torch.randn(batch, seq_len, d_model)
# Create padding mask: True indicates that position should be ignored
key_padding_mask = torch.zeros(batch, seq_len, dtype=torch.bool)
key_padding_mask[0, 7:] = True # The last 3 positions of the first sample are padding
key_padding_mask[1, 5:] = True # The last 5 positions of the second sample are padding
output, attn_weights = model(x, key_padding_mask=key_padding_mask)
print(f"Output shape: {output.shape}") # [4, 10, 128]
print(f"Attention weight shape: {attn_weights.shape}") # [4, 10, 10]
Note the mask types:
nn.MultiheadAttentionofkey_padding_maskThe conventionTrue = ignoreis used, opposite to the convention in some custom implementations where 0 = ignore. Be sure to confirm when using.
3. Variants of Attention Mechanism
3.1 Self-Attention
Self-attentionis a special form of the attention mechanism—Q, K, V all come from the same input sequence. It lets every position in the sequence directly attend to all other positions, thereby capturing global dependencies. This is the core component of the Transformer and the key to its advantage over RNNs: RNNs need to pass information step by step, while self-attention does it in one step.
Example
import torch.nn as nn
import torch.nn.functional as F
import math
class SelfAttention(nn.Module):
"""Multi-head self-attention layer, supports causal mask (for autoregressive generation)"""
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
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.out_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def _split_heads(self, x, batch_size):
"""[batch, seq, d_model] → [batch, n_heads, seq, d_k]"""
return x.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
def forward(self, x, causal=False):
"""
Self-attention: Q, K, V all come from the same input x
Parameters:
x: [batch, seq_len, d_model]
causal: whether to use a causal mask (prevent seeing future positions)
"""
batch_size, seq_len, _ = x.size()
# Project + split into heads
Q = self._split_heads(self.w_q(x), batch_size)
K = self._split_heads(self.w_k(x), batch_size)
V = self._split_heads(self.w_v(x), batch_size)
# Compute attention scores
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# Causal mask: set the upper triangular matrix to -inf to prevent seeing future tokens
if causal:
causal_mask = torch.triu(
torch.ones(seq_len, seq_len, device=x.device, dtype=torch.bool),
diagonal=1
)
scores = scores.masked_fill(causal_mask.unsqueeze(0).unsqueeze(0), float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# Weighted sum + merge heads
context = torch.matmul(attn_weights, V)
context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model)
return self.out_proj(context), attn_weights
# Test
d_model, n_heads = 256, 8
seq_len, batch = 20, 2
attn = SelfAttention(d_model, n_heads)
x = torch.randn(batch, seq_len, d_model)
# Without causal mask (Encoder scenario)
output_enc, weights_enc = attn(x, causal=False)
print(f"[Encoder] Output: {output_enc.shape}, weights: {weights_enc.shape}")
# With causal mask (Decoder / autoregressive generation scenario)
output_dec, weights_dec = attn(x, causal=True)
print(f"[Decoder] Output: {output_dec.shape}, weights: {weights_dec.shape}")
# Verify causality: the first token should not attend to later tokens
print(f"causal=True, weights[0,0,0,5:] (should be all 0): {weights_dec[0, 0, 0, 5:].tolist()[:5]}")
3.2 Cross-Attention
In cross-attention, Q comes from one sequence (e.g., Decoder), while K and V come from another sequence (e.g., Encoder). This is the bridge of the Encoder-Decoder architecture—it lets the Decoder dynamically "look up" all Encoder outputs when generating each token.
Example
import torch.nn as nn
import torch.nn.functional as F
import math
class CrossAttention(nn.Module):
"""Cross-attention: Q comes from the target sequence, K/V come from the source sequence"""
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
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.out_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, target, source, mask=None):
"""
Parameters:
target: target sequence [batch, target_len, d_model] → provides Q
source: source sequence [batch, source_len, d_model] → provides K, V
mask: optional mask
"""
batch_size = target.size(0)
target_len = target.size(1)
source_len = source.size(1)
# Project + split into heads
Q = self.w_q(target).view(batch_size, target_len, self.n_heads, self.d_k).transpose(1, 2)
K = self.w_k(source).view(batch_size, source_len, self.n_heads, self.d_k).transpose(1, 2)
V = self.w_v(source).view(batch_size, source_len, self.n_heads, self.d_k).transpose(1, 2)
# Attention computation
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# Weighted sum + merge heads
context = torch.matmul(attn_weights, V)
context = context.transpose(1, 2).contiguous().view(batch_size, target_len, self.d_model)
return self.out_proj(context), attn_weights
# Test
d_model, n_heads = 256, 8
batch = 2
target_len, source_len = 10, 15
cross_attn = CrossAttention(d_model, n_heads)
target = torch.randn(batch, target_len, d_model) # Decoder output
source = torch.randn(batch, source_len, d_model) # Encoder output
output, weights = cross_attn(target, source)
print(f"Target sequence shape: {target.shape}")
print(f"Source sequence shape: {source.shape}")
print(f"Output shape: {output.shape}") # [2, 10, 256]
print(f"Weights shape: {weights.shape}") # [2, 8, 10, 15]
# Meaning of weights: how much each target position attends to each source position
Self-attention vs. cross-attention:In self-attention, Q/K/V have the same shape ($$seq \times seq$$ weight matrix); in cross-attention, the length of Q can differ from that of K ($$target \times source$$ weight matrix).
3.3 Positional Encoding
The attention mechanism itself haspermutation invariance—shuffling the input order does not change the output. But language and images are ordered, so positional information must be explicitly injected through positional encoding. The original Transformer uses sine/cosine functions to generate fixed positional encodings:
$$PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right), \quad PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$
Example
import torch.nn as nn
import math
class PositionalEncoding(nn.Module):
"""Sinusoidal positional encoding (original Transformer scheme)"""
def __init__(self, d_model, max_len=5000, dropout=0.1):
super().__init__()
self.dropout = nn.Dropout(p=dropout)
# Precompute the positional encoding matrix [max_len, d_model]
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
# div_term: controls the decay speed of different frequency dimensions
div_term = torch.exp(
torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)
)
pe[:, 0::2] = torch.sin(position * div_term) # Even dimensions
pe[:, 1::2] = torch.cos(position * div_term) # Odd dimensions
pe = pe.unsqueeze(0) # [1, max_len, d_model]
self.register_buffer('pe', pe) # Not trainable
def forward(self, x):
"""
x: [batch, seq_len, d_model]
Return: x + positional_encoding
"""
x = x + self.pe[:, :x.size(1), :]
return self.dropout(x)
# Test
d_model = 128
seq_len = 50
batch = 4
pe = PositionalEncoding(d_model)
x = torch.randn(batch, seq_len, d_model)
output = pe(x)
print(f"Input shape: {x.shape}")
print(f"Output shape: {output.shape}")
print(f"Positional encoding shape: {pe.pe.shape}") # [1, 5000, 128]
4. Visual Applications of Attention Mechanism
Attention mechanisms are not only applicable to sequence data, but also shine in computer vision. The following introduces three classic visual attention modules:
4.1 Channel Attention
Channel attention(e.g., SENet) lets the network learn "which channels are more important". It compresses spatial information by global pooling on each channel, then learns inter-channel dependencies through two fully connected layers.
Example
import torch.nn as nn
class SEBlock(nn.Module):
"""Squeeze-and-Excitation Block (channel attention)"""
def __init__(self, channels, reduction=16):
super().__init__()
self.squeeze = nn.AdaptiveAvgPool2d(1) # Global average pooling
self.excitation = nn.Sequential(
nn.Linear(channels, channels // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channels // reduction, channels, bias=False),
nn.Sigmoid() # Output channel weights from 0 to 1
)
def forward(self, x):
# x: [batch, channels, H, W]
b, c, _, _ = x.size()
# Squeeze: [B, C, H, W] → [B, C, 1, 1] → [B, C]
y = self.squeeze(x).view(b, c)
# Excitation: learn channel weights [B, C]
y = self.excitation(y).view(b, c, 1, 1)
# Channel weighting: multiply each channel by its corresponding weight
return x * y.expand_as(x)
# Test
se = SEBlock(channels=256, reduction=16)
x = torch.randn(4, 256, 32, 32)
output = se(x)
print(f"Input: {x.shape} → Output: {output.shape}")
4.2 Spatial Attention
Spatial attentionlets the network learn "which positions are more important". It compresses along the channel dimension (taking the mean and max), then uses convolution to learn a spatial weight map.
Example
import torch.nn as nn
class SpatialAttention(nn.Module):
"""Spatial attention module"""
def __init__(self, kernel_size=7):
super().__init__()
padding = kernel_size // 2
# Input has 2 channels (mean + max), output has 1 channel (attention map)
self.conv = nn.Conv2d(2, 1, kernel_size, padding=padding, bias=False)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
# x: [batch, channels, H, W]
# Channel compression: take mean and max → one [B, 1, H, W] each
avg_out = torch.mean(x, dim=1, keepdim=True)
max_out, _ = torch.max(x, dim=1, keepdim=True)
# Concatenate → [B, 2, H, W]
scale = torch.cat([avg_out, max_out], dim=1)
# Convolution + Sigmoid → spatial weight map [B, 1, H, W]
scale = self.sigmoid(self.conv(scale))
# Spatial weighting
return x * scale
# Test
spatial_attn = SpatialAttention(kernel_size=7)
x = torch.randn(4, 256, 32, 32)
output = spatial_attn(x)
print(f"Input: {x.shape} → Output: {output.shape}")
4.3 CBAM(Convolutional Block Attention Module)
CBAM applies channel attention and spatial attention inseries: first select important channels through channel attention, then select important positions through spatial attention. The two are complementary, and the effect is better than using either alone.
Example
import torch.nn as nn
class CBAM(nn.Module):
"""CBAM: channel attention + spatial attention"""
def __init__(self, channels, reduction=16, kernel_size=7):
super().__init__()
# Channel attention
self.channel_attn = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(channels, channels // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channels // reduction, channels, bias=False),
nn.Sigmoid()
)
# Spatial attention
self.spatial_attn = nn.Sequential(
nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False),
nn.Sigmoid()
)
def forward(self, x):
# 1. Channel attention
b, c, _, _ = x.size()
ca = self.channel_attn(x).view(b, c, 1, 1)
x = x * ca
# 2. Spatial attention
avg_out = torch.mean(x, dim=1, keepdim=True)
max_out, _ = torch.max(x, dim=1, keepdim=True)
sa = self.spatial_attn(torch.cat([avg_out, max_out], dim=1))
x = x * sa
return x
# Test
cbam = CBAM(channels=256)
x = torch.randn(4, 256, 32, 32)
output = cbam(x)
print(f"Input: {x.shape} → Output: {output.shape}")
5. Applications of Attention Mechanism in NLP
5.1 Transformer Encoder
The Transformer Encoder is composed of N identical layers stacked together, each containing two sublayers:multi-head self-attentionandfeed-forward network. Each sublayer uses residual connections and layer normalization.
Example
import torch.nn as nn
import math
class TransformerEncoderLayer(nn.Module):
"""Single Transformer Encoder layer"""
def __init__(self, d_model, n_heads, d_ff, dropout=0.1):
super().__init__()
# Multi-head self-attention
self.self_attn = nn.MultiheadAttention(
d_model, n_heads, dropout=dropout, batch_first=True
)
# Feed-forward network (two linear layers + activation)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.GELU(), # GELU is more commonly used than ReLU
nn.Dropout(dropout),
nn.Linear(d_ff, d_model)
)
# Layer normalization (Pre-Norm style, more stable training)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, key_padding_mask=None):
# Sublayer 1: Self-attention
residual = x
x = self.norm1(x) # Pre-Norm
attn_out, _ = self.self_attn(x, x, x, key_padding_mask=key_padding_mask)
x = residual + self.dropout(attn_out)
# Sublayer 2: Feed-forward network
residual = x
x = self.norm2(x) # Pre-Norm
x = residual + self.dropout(self.ffn(x))
return x
class TransformerEncoder(nn.Module):
"""Multi-layer Transformer Encoder"""
def __init__(self, d_model, n_heads, d_ff, num_layers, dropout=0.1):
super().__init__()
self.layers = nn.ModuleList([
TransformerEncoderLayer(d_model, n_heads, d_ff, dropout)
for _ in range(num_layers)
])
self.final_norm = nn.LayerNorm(d_model)
def forward(self, x, key_padding_mask=None):
for layer in self.layers:
x = layer(x, key_padding_mask)
return self.final_norm(x)
# Test
d_model, n_heads, d_ff = 128, 4, 512
num_layers, seq_len, batch = 3, 20, 4
encoder = TransformerEncoder(d_model, n_heads, d_ff, num_layers)
x = torch.randn(batch, seq_len, d_model)
# Padding mask: True = ignore this position
key_padding_mask = torch.zeros(batch, seq_len, dtype=torch.bool)
key_padding_mask[0, 15:] = True
output = encoder(x, key_padding_mask=key_padding_mask)
print(f"Input: {x.shape} → Output: {output.shape}")
5.2 Attention Visualization
Visualizing attention weights is an important means of understanding and debugging Transformer models. Through heatmaps, you can intuitively see what information the model attends to at each position.
Example
import matplotlib.pyplot as plt
import numpy as np
def visualize_attention(attention_weights, tokens=None, save_path=None):
"""
Visualizing multi-head attention weights
Parameters:
attention_weights: [n_heads, seq_len, seq_len]
tokens: token list (optional)
save_path: save path (optional)
"""
n_heads = attention_weights.shape[0]
seq_len = attention_weights.shape[1]
fig, axes = plt.subplots(1, n_heads, figsize=(n_heads * 3, 3.5))
if n_heads == 1:
axes = [axes]
for head in range(n_heads):
ax = axes[head]
attn = attention_weights[head].detach().cpu().numpy()
im = ax.imshow(attn, cmap='Blues', aspect='auto', vmin=0, vmax=attn.max())
ax.set_title(f'Head {head + 1}', fontsize=10, fontweight='bold')
ax.set_xlabel('Key')
ax.set_ylabel('Query')
if tokens is not None:
ax.set_xticks(range(len(tokens)))
ax.set_yticks(range(len(tokens)))
ax.set_xticklabels(tokens, rotation=90, fontsize=8)
ax.set_yticklabels(tokens, fontsize=8)
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches='tight')
plt.show()
# Simulated attention weights after training
n_heads, seq_len = 4, 8
attention_weights = torch.softmax(torch.randn(n_heads, seq_len, seq_len), dim=-1)
tokens = ['The', 'cat', 'sat', 'on', 'the', 'mat', '.', '[PAD]']
visualize_attention(attention_weights, tokens)
6. Complete Attention Classifier
6.1 Attention-Based Text Classification
The following is a complete text classification model: it uses a Transformer Encoder to extract features, then compresses variable-length sequences into fixed-length vectors via attention pooling, and finally feeds them into a classifier.
Example
import torch.nn as nn
import math
class AttentionClassifier(nn.Module):
"""Text classification model based on Transformer + attention pooling"""
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes,
n_heads=4, num_layers=2, dropout=0.1):
super().__init__()
self.embed_dim = embed_dim
# Word embedding
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.pos_encoding = PositionalEncoding(embed_dim, dropout=dropout)
# Transformer encoder
encoder_layer = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=n_heads,
dim_feedforward=hidden_dim,
dropout=dropout,
batch_first=True,
activation='gelu'
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers)
# Attention pooling: learn importance weights for each position
self.attn_pool = nn.Sequential(
nn.Linear(embed_dim, 1) # Output a scalar weight for each position
)
# Classification head
self.classifier = nn.Sequential(
nn.Linear(embed_dim, hidden_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, num_classes)
)
def forward(self, x):
"""
x: [batch, seq_len] — token ids
Returns: logits [batch, num_classes], attn_weights [batch, seq_len]
"""
# Word embedding + positional encoding
mask = (x == 0) # padding mask
x = self.embedding(x) * math.sqrt(self.embed_dim)
x = self.pos_encoding(x)
# Transformer encoding
x = self.transformer(x, src_key_padding_mask=mask)
# Attention pooling: compute weights for each position, weighted sum
scores = self.attn_pool(x).squeeze(-1) # [batch, seq_len]
scores = scores.masked_fill(mask, float('-inf')) # Ignore padding
attn_weights = torch.softmax(scores, dim=1) # [batch, seq_len]
pooled = torch.bmm(attn_weights.unsqueeze(1), x).squeeze(1) # [batch, embed_dim]
# Classification
logits = self.classifier(pooled)
return logits, attn_weights
# Test
model = AttentionClassifier(
vocab_size=10000, embed_dim=128, hidden_dim=256, num_classes=5
)
# Simulated input (with padding)
x = torch.randint(1, 10000, (16, 50))
x[:, -10:] = 0 # The last 10 positions are padding
logits, attn_weights = model(x)
print(f"Input: {x.shape}")
print(f"Logits: {logits.shape}") # [16, 5]
print(f"Attention weights: {attn_weights.shape}") # [16, 50]
6.2 Comparison of Common Attention Variants
| Type | Q Source | K/V Source | Core Use | Typical Application |
|---|---|---|---|---|
| Self-Attention | Itself | Itself | Global interaction within a sequence | Transformer、BERT、ViT |
| Cross-Attention | Target sequence | Source sequence | Cross-sequence information fusion | Machine translation, image captioning, Stable Diffusion |
| Channel Attention | Global pooled features | Same as above | Learn channel importance | SENet、CBAM |
| Spatial Attention | Channel-compressed features | Same as above | Learn spatial position importance | CBAM、Spatial Transformer |
| Sparse Attention | Partial positions | Partial positions | Reduce computational complexity for long sequences | Longformer、BigBird |
7. Best Practices and Common Issues
7.1 Usage Tips
- Mask consistency: Ensure that the padding mask and causal mask have consistent direction (True = ignore vs 1 = keep). Different API conventions may differ.
- Number of heads selection: Usually set as a divisor of $$d_{model}$$, commonly 8 or 16. Too many heads reduce the per-head dimension $$d_k$$, potentially harming performance.
- Residual connection: Be sure to use residual connections + layer normalization before and after attention layers; this is a necessary condition for training deep Transformers.
- Learning rate warmup: Gradient fluctuations are large in the early stage of Transformer training; using a warmup + cosine decay scheduler can significantly improve stability.
- Flash Attention: PyTorch 2.0+ supports via
nn.functional.scaled_dot_product_attentionautomatically invoking Flash Attention, significantly optimizing both memory usage and speed.
7.2 Common Issues
| Problem | Cause | Solution |
|---|---|---|
| Attention weights are all uniformly distributed | Unnormalized input or learning rate too large | Check input scale, use LayerNorm, reduce learning rate |
| Training loss does not decrease | Missing positional encoding or incorrect mask | Confirm positional encoding is added, check mask direction |
| Out of memory (long sequences) | Attention matrix has $$O(n^2)$$ complexity | Use Flash Attention, gradient checkpointing, or sparse attention |
| Attention degradation during inference | Inconsistent mask between training/inference | Ensure that the padding mask uses the same logic during training and inference |
Other ExtensionsThe attention mechanism is a core technology of modern deep learning. From NLP to CV, from speech to protein structure prediction, almost all cutting-edge models are based on attention. Understanding its principles and implementation is the key first step to mastering models such as Transformer, ViT, Stable Diffusion, and GPT.