注意力机制实现·代码逐行分析
1. 导入模块
import torch.nn.functional as F
import math
-
import torch.nn.functional as F:导入PyTorch的神经网络函数库,用于调用softmax等函数。 -
import math:导入数学库,用于计算平方根。
注意:代码中使用了torch.matmul等函数,因此实际运行时还需要import torch,但此处未显式写出,我们假设已导入。
2. 缩放点积注意力函数
def attention(Q, K, V, mask=None):
"""缩放点积注意力"""
定义函数attention,接受查询张量Q、键张量K、值张量V和一个可选的掩码mask。
文档字符串,该函数实现的是缩放点积注意力机制。
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(Q.size(-1))
-
torch.matmul, K.transpose(-2, -1)):计算查询和键的点积。Q的形状通常为(batch_size, ..., seq_len_q, d_k),K的形状为(batch_size, ..., seq_len_k, d_k)。对K的最后两个维度转置,得到形状(batch_size, ..., d_k, seq_len_k),然后矩阵乘法得到形状(batch_size, ..., seq_len_q, seq_len_k)的注意力分数矩阵。 -
/ math.sqrt(Q.size(-1)):除以d_k的平方根进行缩放,防止点积过大导致softmax梯度消失。Q.size(-1)即为d_k。
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
如果提供了掩码,则将scores中对应mask == 0的位置填充为一个极小的负数(-1e9)。这样在后续softmax中,这些位置的权重趋近于0,起到遮蔽作用。
attention_weights = F.softmax(scores, dim=-1)
在最后一个维度(seq_len_k)上应用softmax,将分数转换为概率分布,得到注意力权重,形状与scores相同。
output = torch.matmul(attention_weights, V)
将注意力权重与值张量V相乘,得到加权求和后的输出。V的形状为(batch_size, ..., seq_len_k, d_v),乘法结果形状为(batch_size, ..., seq_len_q, d_v)。
return output, attention_weights
返回注意力输出和注意力权重。
3. 多头注意力类
class MultiHeadAttention(torch.nn.Module):
定义一个继承自torch.nn.Module的类,实现多头注意力机制。
def __init__(self, d_model, n_heads):
super().__init__()
构造函数,接受模型维度d_model和头数n_heads。首先调用父类构造函数初始化。
self.d_model = d_model
self.n_heads = n_heads
保存型维度和头数。
self.d_k = d_model // n_heads
计算每个头的维度,假设d_model能被n_heads整除。
self.W_q = torch.nn.Linear(d_model, d_model)
self.W_k = torch.nn.Linear(d_model, d_model)
self.W_v = torch.nn.Linear(d_model, d_model)
创建三个线性变换层,分别用于将输入的查询、键、值从d_model维映射到d_model维。每个头的线性变换是共享的,但后续会通过重塑来分头处理。
self.W_o = torch.nn.Linear(d_model, d_model)
创建输出线性变换层,用于将多头合并后的结果再映射回d_model维。
def forward(self, query, key, value, mask=None):
定义前向传播函数,接受查询、键、值以及可选的掩码。通常输入形状为(batch_size, seq_len, d_model)。
batch_size = query.size(0)
获取批次大小,假设输入的第一维是batch维度。
# 线性变换并重塑
Q = self.W_q(query).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
-
self.W_q(query):对查询进行线性变换,输出形状(batch_size, seq_len_q, d_model)。 -
.view(batch_size, -1, self.n_heads, self.d_k):重塑为(batch_size, seq_len_q, n_heads, d_k)。 -
.transpose(1, 2):交换第1维(seq_len_q)和第2维(n_heads),得到形状(batch_size, n_heads, seq_len_q, d_k)。这样每个头独立处理。
K = self.W_k(key).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
类似地对键进行处理,得到形状(batch_size, n_heads, seq_len_k, d_k)。
V = self.W_v(value).view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
类似地对值进行处理,得到形状(batch_size, n_heads, seq_len_k, d_k)。
# 注意力计算
output, attention_weights = attention(Q, K, V, mask)
调用前面定义的attention函数,传入多头形式的Q、K、V和掩码。返回的输出形状为(batch_size, n_heads, seq_len_q, d_k),注意力权重形状为(batch_size, n_heads, seq_len_q, seq_len_k)。
# 合并多头
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
-
.transpose(1, 2):将维度从(batch_size, n_heads, seq_len_q, d_k)转换为(batch_size, seq_len_q, n_heads, d_k)。 -
.contiguous():确保张量在内存中是连续的,因为transpose可能产生非连续张量,后续view操作需要连续性。 -
.view(batch_size, -1, self.d_model):将后两维合并,n_heads * d_k = d_model,得到形状(batch_size, seq_len_q, d_model)。
return self.W_o(output), attention_weights
-
self.W_o(output):对合并后的输出应用输出线性变换,得到最终的多头注意力输出,形状为(batch_size, seq_len_q, d_model)。 -
同时返回注意力权重,用于可视化或分析。
4.注意力机制完整实现
import torch
import torch.nn.functional as F
import math
def attention(Q, K, V, mask=None):
"""缩放点积注意力"""
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(Q.size(-1))
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attention_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attention_weights, V)
return output, attention_weights
class MultiHeadAttention(torch.nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
self.W_q = torch.nn.Linear(d_model, d_model)
self.W_k = torch.nn.Linear(d_model, d_model)
self.W_v = torch.nn.Linear(d_model, d_model)
self.W_o = torch.nn.Linear(d_model, d_model)
def forward(self, query, key, value, mask=None):
batch_size = query.size(0)
# 线性变换并重塑
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)
# 注意力计算
output, attention_weights = attention(Q, K, V, mask)
# 合并多头
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.W_o(output), attention_weights
更多推荐

所有评论(0)