手撕注意力机制与多头注意力:核心原理与代码实现
·
一、引言
在深度学习领域,尤其是自然语言处理中,注意力机制已经成为革命性的技术。今天,我们将深入理解注意力机制和多头注意力的核心原理,并用Python从零开始实现它们。
二、注意力机制
核心思想:
注意力机制模拟了人类大脑的注意力分配过程,让模型能够动态地关注输入中不同部分的重要性。其核心公式为:
Attention(Q, K, V) = softmax(QK^T / √d_k) V
核心代码:
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class ScaledDotProductAttention(nn.Module):
"""缩放点积注意力机制"""
def __init__(self, dropout=0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
def forward(self, q, k, v, mask=None):
"""
前向传播
Args:
q: 查询向量 [batch_size, seq_len, d_k]
k: 键向量 [batch_size, seq_len, d_k]
v: 值向量 [batch_size, seq_len, d_v]
mask: 注意力掩码
"""
# 计算查询和键的点积相似度,并缩放
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.size(-1))
# 应用掩码(将掩码位置设为负无穷,softmax后为0)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 计算注意力权重(softmax归一化)
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 用注意力权重对值向量进行加权求和
output = torch.matmul(attn_weights, v)
return output, attn_weights
注意力机制的工作原理:
- 相似度计算:通过点积衡量查询和键的相似程度
- 缩放处理:防止点积值过大导致梯度问题
- 权重归一化:使用softmax将相似度转换为概率分布
- 加权求和:用注意力权重对值向量进行加权
三、多头注意力机制
核心思想:
多头注意力将输入投影到多个子空间,在每个子空间中独立计算注意力,最后将结果合并。这种设计允许模型同时关注来自不同表示子空间的信息。
核心代码;
class MultiHeadAttention(nn.Module):
def __init__(self, hidden_dim, num_heads):
super(MultiHeadAttention, self).__init__()
assert hidden_dim % num_heads == 0
self.hidden_dim = hidden_dim
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads
# 线性变换层,用于将输入分别映射到query、key和value
self.q_linear = nn.Linear(hidden_dim, hidden_dim)
self.k_linear = nn.Linear(hidden_dim, hidden_dim)
self.v_linear = nn.Linear(hidden_dim, hidden_dim)
# 最终的线性变换层,用于将多头注意力的结果进行融合
self.out_linear = nn.Linear(hidden_dim, hidden_dim)
def forward(self, query, keys, values):
batch_size = query.size(0)
# 将输入通过线性变换得到query、key和value
q = self.q_linear(query)
k = self.k_linear(keys)
v = self.v_linear(values)
# 将query、key和value分割成多头
q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
# 计算分数,这里通过矩阵乘法和缩放操作
scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5)
# 对分数进行softmax归一化,得到注意力权重
attention_weights = F.softmax(scores, dim=-1)
# 根据注意力权重加权求和得到多头的输出
output = torch.matmul(attention_weights, v)
# 将多头的输出合并起来
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.hidden_dim)
# 通过最终的线性变换层得到最终输出
output = self.out_linear(output)
return output
更多推荐
所有评论(0)