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.

Seq2Seq: Fixed Vector vs. Attention Mechanism Fixed vector (information bottleneck) x₁ x₂ x₃ x₄ c (fixed vector) y₁ y₂ y₃ y₄ All decoding steps share the same c Attention mechanism (dynamic weighting) h₁ h₂ h₃ h₄ y₁ y₂ y₃ y₄ Each step generates a different context vector c_t Thicker line → larger attention weight → the more important the information at that position

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.

Scaled Dot-Product Attention computation process Q [batch, seq_q, d_k] K [batch, seq_k, d_k] MatMul QKᵀ Scale ÷ √d_k Softmax Normalization V [batch, seq_k, d_v] MatMul × V Output Mask (optional) -∞ padding Output = softmax(QKᵀ / √d_k) · V Output shape: [batch, seq_q, d_v] — each Query position gets a d_v-dimensional context vector

Example

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):
    """
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
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
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
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
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
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
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
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
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
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 torch
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
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 viann.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

The 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.

Other Extensions