PyTorch torch.nn.MultiheadAttention Function
PyTorch torch.nn Reference Manual
torch.nn.MultiheadAttentionIt is a multi-head attention mechanism module in PyTorch.
It is a core component of the Transformer architecture, allowing the model to simultaneously attend to information from different representation subspaces at different positions.
Function Definition
torch.nn.MultiheadAttention(embed_dim, num_heads, dropout=0.0, bias=True, add_bias_kv=False, kdim=None, vdim=None, batch_first=True)
Parameter Description:
embed_dim(int): Input embedding dimension.num_heads(int): Number of attention heads.dropout(float): Dropout probability. Default is 0.kdim(int): Dimension of the key vectors. Default is None (same as embed_dim).vdim(int): Dimension of the value vectors. Default is None (same as embed_dim).batch_first(bool): If True, the first dimension of input and output is batch. Default is True.
Usage Examples
Example 1: Basic Usage
Example
import torch
import torch.nn as nn
# Create multi-head attention: 512-dimensional, 8 heads
mha = nn.MultiheadAttention(embed_dim=512, num_heads=8)
# Input: batch=4, sequence length=100, dimension=512
query = torch.randn(4, 100, 512)
key = torch.randn(4, 100, 512)
value = torch.randn(4, 100, 512)
# Forward pass
output, attn_weight = mha(query, key, value)
print("Query shape:", query.shape)
print("Output shape:", output.shape)
print("Attention weight shape:", attn_weight.shape)
import torch.nn as nn
# Create multi-head attention: 512-dimensional, 8 heads
mha = nn.MultiheadAttention(embed_dim=512, num_heads=8)
# Input: batch=4, sequence length=100, dimension=512
query = torch.randn(4, 100, 512)
key = torch.randn(4, 100, 512)
value = torch.randn(4, 100, 512)
# Forward pass
output, attn_weight = mha(query, key, value)
print("Query shape:", query.shape)
print("Output shape:", output.shape)
print("Attention weight shape:", attn_weight.shape)
Example 2: Self-Attention
Example
import torch
import torch.nn as nn
mha = nn.MultiheadAttention(embed_dim=256, num_heads=4)
# Use the same input as Q, K, V (self-attention)
x = torch.randn(2, 50, 256)
# self-attention: q=k=v=x
output, weights = mha(x, x, x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Attention weight shape:", weights.shape)
import torch.nn as nn
mha = nn.MultiheadAttention(embed_dim=256, num_heads=4)
# Use the same input as Q, K, V (self-attention)
x = torch.randn(2, 50, 256)
# self-attention: q=k=v=x
output, weights = mha(x, x, x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Attention weight shape:", weights.shape)
Example 3: Attention with Mask
Example
import torch
import torch.nn as nn
mha = nn.MultiheadAttention(embed_dim=128, num_heads=4)
# Input
x = torch.randn(1, 20, 128)
# Create an upper triangular mask (for decoder)
mask = torch.triu(torch.ones(20, 20), diagonal=1).bool()
output, _ = mha(x, x, x, attn_mask=mask)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Mask shape:", mask.shape)
import torch.nn as nn
mha = nn.MultiheadAttention(embed_dim=128, num_heads=4)
# Input
x = torch.randn(1, 20, 128)
# Create an upper triangular mask (for decoder)
mask = torch.triu(torch.ones(20, 20), diagonal=1).bool()
output, _ = mha(x, x, x, attn_mask=mask)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
print("Mask shape:", mask.shape)
Example 4: Complete Transformer Encoder Layer
Example
import torch
import torch.nn as nn
class TransformerLayer(nn.Module):
def __init__(self, d_model, nhead):
super(TransformerLayer, self).__init__()
self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_model * 4),
nn.ReLU(),
nn.Linear(d_model * 4, d_model)
)
def forward(self, x):
# Self-attention + residual
attn_out, _ = self.self_attn(x, x, x)
x = self.norm1(x + attn_out)
# FFN + residual
ffn_out = self.ffn(x)
x = self.norm2(x + ffn_out)
return x
# Test
layer = TransformerLayer(d_model=512, nhead=8)
x = torch.randn(4, 100, 512)
output = layer(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
import torch.nn as nn
class TransformerLayer(nn.Module):
def __init__(self, d_model, nhead):
super(TransformerLayer, self).__init__()
self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_model * 4),
nn.ReLU(),
nn.Linear(d_model * 4, d_model)
)
def forward(self, x):
# Self-attention + residual
attn_out, _ = self.self_attn(x, x, x)
x = self.norm1(x + attn_out)
# FFN + residual
ffn_out = self.ffn(x)
x = self.norm2(x + ffn_out)
return x
# Test
layer = TransformerLayer(d_model=512, nhead=8)
x = torch.randn(4, 100, 512)
output = layer(x)
print("Input shape:", x.shape)
print("Output shape:", output.shape)
Example 5: Viewing Attention Weights
Example
import torch
import torch.nn as nn
import numpy as np
mha = nn.MultiheadAttention(embed_dim=64, num_heads=2, batch_first=True)
# Short video sequence
x = torch.randn(1, 5, 64)
_, attn = mha(x, x, x)
attn = attn.squeeze(0) # Remove batch dimension
print("Attention weights of the first head (first 3 positions):")
print(attn[0, :3, :3].tolist())
print("nVisualization - attention of position 0 to all positions:")
print(np.array2string(attn[0, 0].numpy(), precision=2))
import torch.nn as nn
import numpy as np
mha = nn.MultiheadAttention(embed_dim=64, num_heads=2, batch_first=True)
# Short video sequence
x = torch.randn(1, 5, 64)
_, attn = mha(x, x, x)
attn = attn.squeeze(0) # Remove batch dimension
print("Attention weights of the first head (first 3 positions):")
print(attn[0, :3, :3].tolist())
print("nVisualization - attention of position 0 to all positions:")
print(np.array2string(attn[0, 0].numpy(), precision=2))
Frequently Asked Questions
Q1: How to choose num_heads?
embed_dim must be divisible by num_heads. Common values: 8, 12, 16.
Q2: Why are the three matrices Q, K, and V needed?
It allows the model to learn different projections and enhance expressive power.
Q3: What is key_padding_mask?
It is used to mask padding positions, preventing attention from being computed on padding.
Use Cases
- Transformer: Encoder and Decoder
- Self-attention model: BERT、GPT
- Sequence modeling: Alternative to RNN
Tip: When batch_first=True, the input shape is (batch, seq, embed_dim).
Other Extensions