Python 和 PyTorch 手动实现多头注意力机制(Multi-Head Attention)
·
以下是使用 Python 和 PyTorch 手动实现多头注意力机制(Multi-Head Attention)的详细代码及解释:
导入必要的库
python
运行
import torch
import torch.nn as nn
定义缩放点积注意力(Scaled Dot-Product Attention)函数
缩放点积注意力是多头注意力机制的核心部分,它计算 query、key 和 value 之间的注意力权重,并返回加权后的 value。
python
运行
def scaled_dot_product_attention(query, key, value, mask=None):
d_k = query.size(-1)
scores = torch.matmul(query, key.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attention_weights = torch.softmax(scores, dim=-1)
return torch.matmul(attention_weights, value)
代码解释:
d_k表示 query 和 key 的维度。scores通过矩阵乘法计算 query 和 key 的转置的乘积,然后除以 \(\sqrt{d_k}\) 进行缩放。- 如果提供了掩码(mask),将掩码中值为 0 的位置对应的分数设为负无穷,这样在 softmax 后这些位置的注意力权重接近 0。
attention_weights使用 softmax 函数计算注意力权重。- 最后返回注意力权重与 value 的矩阵乘积。
定义多头注意力类(Multi-Head Attention)
python
运行
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
super(MultiHeadAttention, self).__init__()
self.num_heads = num_heads
self.d_model = d_model
self.d_k = d_model // num_heads
assert d_model % num_heads == 0
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.W_o = nn.Linear(d_model, d_model)
def split_heads(self, x, batch_size):
x = x.view(batch_size, -1, self.num_heads, self.d_k)
return x.transpose(1, 2)
def forward(self, query, key, value, mask=None):
batch_size = query.size(0)
Q = self.split_heads(self.W_q(query), batch_size)
K = self.split_heads(self.W_k(key), batch_size)
V = self.split_heads(self.W_v(value), batch_size)
attention_output = scaled_dot_product_attention(Q, K, V, mask)
attention_output = attention_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.W_o(attention_output)
代码解释:
- 初始化函数
__init__:d_model是输入和输出的维度。num_heads是头的数量,d_k是每个头的维度,确保d_model能被num_heads整除。- 定义了四个线性变换层
W_q、W_k、W_v和W_o,分别用于将输入的 query、key、value 进行线性变换,以及将多头注意力的输出进行线性变换得到最终结果。
split_heads方法:将线性变换后的张量按照头的数量进行拆分,调整维度顺序以便后续计算。forward前向传播函数:- 对输入的 query、key、value 进行线性变换并拆分成多个头。
- 调用
scaled_dot_product_attention函数计算每个头的注意力输出。 - 将各个头的输出进行维度调整和拼接,再通过
W_o进行线性变换得到最终的多头注意力输出。
使用示例
python
运行
# 示例输入
batch_size = 2
seq_length = 3
d_model = 6
num_heads = 2
query = torch.rand(batch_size, seq_length, d_model)
key = torch.rand(batch_size, seq_length, d_model)
value = torch.rand(batch_size, seq_length, d_model)
multihead_attn = MultiHeadAttention(d_model, num_heads)
output = multihead_attn(query, key, value)
print(output.shape)
![]()
上述示例代码中,先定义了输入张量的形状,然后创建了 MultiHeadAttention 实例,并传入 query、key 和 value 进行计算,最后打印输出张量的形状。通过这样的手动实现,可以更清晰地理解多头注意力机制的工作原理和计算流程。
更多推荐


所有评论(0)