深入浅出:自注意力机制与多头注意力机制完全指南
一、为什么需要注意力机制?
1.1 从RNN/LSTM的局限说起
在自注意力出现之前,处理序列数据(如文字、语音)主要依赖RNN和LSTM。
RNN的工作方式----像“传话游戏”:
"我 昨天 在 北京 吃 的 烤鸭 很 好吃"
↓ ↓ ↓ ↓ ↓ ↓ ↓ ↓ ↓
第1步→第2步→第3步→第4步→第5步→第6步→第7步→第8步→第9步
(必须按顺序,一步一步处理)
问题一:长距离信息丢失
"这部电影的导演,虽然之前拍过很多烂片,
但这次的作品,剧情紧凑,演员出色,
特效震撼,总体来说非常值得一看,它真的很棒"
问题:"它"指的是什么?
答案:要联系到最开头的"这部电影"
困境:信息经过层层传递,到最后已经严重衰减!
距离近(相邻词):信息保留 ████████ 80% ✓
距离中(隔几个词):信息保留 █████ 50% △
距离远(隔很多词):信息保留 █ 10% ✗
问题二:无法并行计算
串行处理(RNN): 并行处理(理想):
处理"我" ⏳ "我" ━━━━━━━→ ✓
↓等待... "爱" ━━━━━━━→ ✓ 同时完成!
处理"爱" ⏳ "吃" ━━━━━━━→ ✓
↓等待... "烤鸭" ━━━━━━━→ ✓
处理"吃" ⏳
↓等待...
处理"烤鸭" ✓
(GPU大量闲置浪费!)
1.2 自注意力的解决思路
自注意力的核心思想:让序列中每个元素都能直接“联系”其他所有元素,不需要中间人传话
RNN(接力传话): Self-Attention(直接通话):
我→爱→吃→烤鸭 我 ←──────────────→ 烤鸭
(信息逐步衰减) 我 ←──────→ 吃
(任意两词,直接建立联系)
- Self-Attention就像把传话游戏改成了“群聊”,所有人可以同时互相交流,不需要层层转达
二、自注意力机制详解
2.1 基本概念
Self-Attention的核心特性:
- 输入N个向量,输出N个向量(数量不变)
- 每个输出向量都融合了整个序列的上下文信息
- 可以与FC(全连接层)叠加使用
📥 输入 📤 输出
[a¹] ──┐ ┌── [b¹] ← 融入了所有词的信息
[a²] ──┤ Self- ├── [b²] ← 融入了所有词的信息
[a³] ──┤ Attention ├── [b³] ← 融入了所有词的信息
[a⁴] ──┘ └── [b⁴] ← 融入了所有词的信息
数量不变!但内容已经"升级"了!
生活类比:就像5个同学开完班会,每个人的认知都升级了,但同学数量还是5个。开会前的“我”和开会后的“我”,是同一个人,但想法更丰富了
2.2 与CNN的关系
Self-Attention可以理解为感受野可以自学习的CNN:
对比 CNN Self-Attention
感受野 固定,只看周围几个 自由,能看所有位置
关注范围 人工设定 数据自己学习
比喻 固定焦距相机 自动变焦相机
全局建模 需要多层堆叠 天生全局视野
2.3 三个核心角色:Q、K、V
Self-Attention引入了三个关键矩阵:
🔍 图书馆搜索类比:
Q(Query,查询)= 你输入的搜索词 "我想找什么?"
K(Key,键) = 每本书的书名标签 "我有什么标签?"
V(Value,值) = 每本书的正文内容 "我的实际内容是什么?"
流程:
搜索词(Q) 匹配 书名(K) → 得到相关度分数α
然后按分数α,提取对应书的内容(V)
→ 得到真正需要的信息!
生成Q、K、V的方式:
原始输入向量 a
├── × Wq矩阵 → Q(查询向量)
├── × Wk矩阵 → K(键向量)
└── × Wv矩阵 → V(值向量)
⚠️ 重要:Wq、Wk、Wv 是模型唯一需要学习的参数!
2.4 计算流程:如何产生输出b¹?
第一步:计算关联度α
方法一:点积(Transformer使用)
a¹ → × Wq → q¹
a² → × Wk → k²
↘
q¹ · k² = α₁,₂ (a2对a1的影响程度)
点积计算示例:
q¹ = [1, 2, 3]
k² = [4, 5, 6]
点积 = 1×4 + 2×5 + 3×6 = 4 + 10 + 18 = 32
方法二:Additive(加法注意力)
q¹ ──┐
├→ 拼接[q¹,k²] → 激活函数 → α₁,₂
k² ──┘
点积像两人直接握手感受契合度
Additive像两人通过第三方中间人评判相关性
第二步:除以√dk,防止极端化
🌡️ 不除以√dk(温度低):
某个α特别大 → Softmax后≈1.0
其他α很小 → Softmax后≈0.0
→ 模型变得"极端",梯度消失,难以训练 ❌
🌡️ 除以√dk(温度适中):
所有α缩小 → Softmax后分布均匀
→ 梯度正常,模型正常学习 ✓
第三步:Softmax归一化
原始α值(大小不一): Softmax后(变成概率):
α₁,₁ = 32 α'₁,₁ = 0.25 ┐
α₁,₂ = 18 →Softmax→ α'₁,₂ = 0.15 ├ 加起来 = 1.0 ✓
α₁,₃ = 45 α'₁,₃ = 0.50 │
α₁,₄ = 7 α'₁,₄ = 0.10 ┘
Softmax会放大差距:重要的更重要,不重要的更不重要!
第四步:加权融合V,得到输出b¹
b¹ = α'₁,₁ × v¹ + α'₁,₂ × v² + α'₁,₃ × v³ + α'₁,₄ × v⁴
↑25%关注自己 ↑15%关注a2 ↑50%关注a3 ↑10%关注a4
含义:b¹是按注意力比例,融合了所有词信息的"升级版a¹"
2.5 矩阵化计算(完整公式)
把所有计算打包成矩阵,一次性算完,高效!
完整计算公式:
┌ QKᵀ ┐
Attention│ ──── │ V = Output
└ √dk ┘
步骤分解:
1. Q = 输入 × Wq 生成查询矩阵
2. K = 输入 × Wk 生成键矩阵
3. V = 输入 × Wv 生成值矩阵
4. A = QKᵀ / √dk 计算注意力分数
5. A' = Softmax(A) 归一化(注意力矩阵)
6. O = A' × V 加权融合,得到输出
注意力矩阵A的可视化:
k¹ k² k³ k⁴
q¹ [ 0.25 0.15 0.50 0.10 ] ← a1对所有词的关注比例
q² [ 0.10 0.40 0.30 0.20 ] ← a2对所有词的关注比例
q³ [ 0.60 0.05 0.25 0.10 ] ← a3对所有词的关注比例
q⁴ [ 0.20 0.30 0.15 0.35 ] ← a4对所有词的关注比例
每行加起来 = 1.0 ✓ (这就是注意力矩阵A')
2.6 位置编码:给词语贴上“座位号”
问题:Self-Attention天生无序!
句子A:"猫追老鼠" → 向量集合{猫,追,老鼠}
句子B:"老鼠追猫" → 向量集合{猫,追,老鼠}
对Self-Attention来说:这两句话没区别!😱
(因为它只知道有哪些词,不知道顺序)
解决方案:位置编码(Positional Encoding)
核心思想:给每个位置贴上独特的"座位号"标签
原来(无位置信息): 加入位置编码后:
a¹ = "我"的语义 a¹ + e¹ = "我"的语义 + "第1位"的标记
a² = "爱"的语义 → a² + e² = "爱"的语义 + "第2位"的标记
a³ = "烤鸭"的语义 a³ + e³ = "烤鸭"的语义 + "第3位"的标记
类比:给包裹加上地址标签,既知道内容,又知道位置 📦
四种位置编码方法对比:
方法 原理 优点 缺点 代表模型
Sinusoidal 数学正弦/余弦函数 无需训练,支持任意长度 不够灵活 原版Transformer
Position Embedding 模型自学习 灵活,效果好 长度受训练限制 BERT、GPT
Floater 连续函数 支持小数位置,更精细 较复杂 特殊任务
RNN 循环网络 天然有序 失去并行优势 较少使用
2.7 Self-Attention与FC层的配合
Self-Attention = 集体开大会(处理整个序列的关系)
FC层 = 老师单独辅导(深度加工单个词的特征)
流水线:
输入序列
↓
【Self-Attention层】
所有词互相交流,融合全局信息
↓
【FC层】(对每个词独立处理)
对"我"深度加工 → 更精准的"我"
对"烤鸭"深度加工 → 更精准的"烤鸭"
↓
输出(更高质量的表示)
二者叠加:先集体讨论,再单独精炼 = 效果最佳!
三、多头注意力机制详解
3.1 为什么需要多头?
单头注意力就像只派一位专家分析文章----视角单一,容易遗漏:
句子:"我喜欢在晴天去北京的故宫游览"
单头注意力(一种视角):
→ 可能只捕捉到"我"和"游览"(主谓关系)✓
→ 忽略了"晴天"和"故宫"(时间地点关系)✗
→ 忽略了"北京"和"故宫"(位置归属关系)✗
多头注意力(多种视角):
头1 → 主谓关系:我 ←→ 游览 ✓
头2 → 地点关系:北京 ←→ 故宫 ✓
头3 → 时间关系:晴天 ←→ 游览 ✓
头4 → 修饰关系:北京的 ←→ 故宫 ✓
→ 多角度全面理解!✓
盲人摸象:一个盲人只摸到象腿,说大象像柱子。多个盲人从不同角度同时摸,合并结果,才能还原大象的真实形象!
3.2 类比CNN的多通道
CNN多通道处理图片: 多头注意力处理句子:
通道1:检测横向边缘 ━━━ 头1:捕捉语法关系
通道2:检测竖向边缘 ┃┃┃ 头2:捕捉语义关系
通道3:检测颜色信息 🎨 头3:捕捉指代关系
通道4:检测纹理信息 ≋≋≋ 头4:捕捉时序关系
多个通道/头,各司其职,捕捉不同特征!
3.3 完整计算流程
第一步:生成h组(Q,K,V)
原始输入
↓ ↓ ↓ ↓ ↓ ↓ ↓ ↓ (h=8个头)
Wq¹ Wq² Wq³ Wq⁴ Wq⁵ Wq⁶ Wq⁷ Wq⁸
Wk¹ Wk² Wk³ Wk⁴ Wk⁵ Wk⁶ Wk⁷ Wk⁸
Wv¹ Wv² Wv³ Wv⁴ Wv⁵ Wv⁶ Wv⁷ Wv⁸
↓ ↓ ↓ ↓ ↓ ↓ ↓ ↓
(Q,K,V)¹ (Q,K,V)² ... (Q,K,V)⁸
每组(Q,K,V)使用不同的变换矩阵,学习不同的特征
第二步:h组并行计算Attention
(Q¹,K¹,V¹) → 注意力计算 → O¹ (语法视角的输出)
(Q²,K²,V²) → 注意力计算 → O² (语义视角的输出)
(Q³,K³,V³) → 注意力计算 → O³ (指代视角的输出)
...
(Q⁸,K⁸,V⁸) → 注意力计算 → O⁸ (第8种视角的输出)
所有头并行计算,互不干扰!
第三步:Concat拼接所有头的输出
O¹ = [0.5, 0.3, 0.2] 头1的理解(3维)
O² = [0.8, 0.1, 0.6] 头2的理解(3维)
O³ = [0.4, 0.7, 0.5] 头3的理解(3维)
拼接后:
[0.5, 0.3, 0.2, | 0.8, 0.1, 0.6, | 0.4, 0.7, 0.5]
← 头1的视角 → ← 头2的视角 → ← 头3的视角 →
(维度变为3倍:9维)
类比:把所有专家的报告装订成一本合集
第四步:线性变换,恢复原始维度
拼接后(维度 = num_head × d_v)
↓ × Wo(线性变换矩阵)
输出(维度恢复为原始词向量维度)
类比:总编辑把厚厚的合集报告,精炼成标准格式的最终报告
3.4 多头注意力完整图解
输入词向量(d_model维)
↓
┌───────────────┼───────────────┐
↓ ↓ ↓
头1(h=1) 头2(h=2) ... 头h(h=H)
×Wq¹,Wk¹,Wv¹ ×Wq²,Wk²,Wv² ×Wqh,Wkh,Wvh
↓ ↓ ↓
Attention₁ Attention₂ Attentionh
↓ ↓ ↓
O¹ O² Oh
└───────────────┼───────────────┘
↓
Concat拼接
[O¹ | O² | ... | Oh]
↓
× Wo线性变换
↓
最终输出(d_model维)✓
3.5 MHA vs MQA : 多头注意力的进化
为了提升推理速度、降低内存消耗,研究者提出了多查询注意力(MQA):
MHA(多头注意力): MQA(多查询注意力):
头1: Q¹ K¹ V¹ 头1: Q¹ ─┐
头2: Q² K² V² 头2: Q² ├─ 共享 K V
头3: Q³ K³ V³ 头3: Q³ │
头4: Q⁴ K⁴ V⁴ 头4: Q⁴ ─┘
每个头都有独立K和V 只有Q是多头,K和V全头共享
(参数多,内存大) (参数少,内存小,速度快)
四、代码实现
import torch.nn as nn # 导入PyTorch的神经网络模块,包含各种网络层
import numpy as np # 导入NumPy,用于数值计算
import torch # 导入PyTorch主库
import math # 导入数学库,用于计算平方根等数学运算
# ============================================================
# 多头注意力机制(Multi-Head Attention, MHA)
# 每个"头"都有自己独立的Q、K、V变换矩阵
# 类比:多个专家各自用不同的分析框架独立分析同一段文字
# ============================================================
class MHA(nn.Module):
def __init__(self, num_head, dimension_k, dimension_v, d_k, d_v, d_o):
"""
初始化多头注意力机制
参数说明:
- num_head: 头的数量,即"几个专家同时分析"
- dimension_k: 输入Q和K的词向量维度(每个词用多少维数字表示)
- dimension_v: 输入V的词向量维度
- d_k: 每个头处理Q和K时的维度(单个专家的分析维度)
- d_v: 每个头处理V时的维度
- d_o: 最终输出的维度
"""
super().__init__() # 调用父类nn.Module的初始化方法(固定写法)
# 保存超参数,方便后续forward函数使用
self.num_head = num_head # 保存头的数量
self.d_k = d_k # 保存每个头的K/Q维度
self.d_v = d_v # 保存每个头的V维度
self.d_o = d_o # 保存输出维度
# 定义线性变换层(全连接层),用于生成Q、K、V矩阵
# fc_q: 将输入q从dimension_k维 → 变换到 (num_head × d_k)维
# 相当于:同时为所有头生成Q矩阵
# 例如:8个头,每个头16维 → 输出128维
self.fc_q = nn.Linear(dimension_k, num_head * d_k)
# fc_k: 将输入k从dimension_k维 → 变换到 (num_head × d_k)维
# 相当于:同时为所有头生成K矩阵
self.fc_k = nn.Linear(dimension_k, num_head * d_k)
# fc_v: 将输入v从dimension_v维 → 变换到 (num_head × d_v)维
# 相当于:同时为所有头生成V矩阵
self.fc_v = nn.Linear(dimension_v, num_head * d_v)
# fc_o: 最终输出层,将拼接后的多头结果 → 变换到目标输出维度d_o
# 相当于:总编辑把所有专家报告整合成标准格式
self.fc_o = nn.Linear(num_head * d_v, d_o)
# Softmax层:将注意力分数转换为概率分布(所有值加起来=1)
# dim=2 表示在第3个维度(序列长度维度)上做归一化
self.softmax = nn.Softmax(dim=2)
def forward(self, q, k, v, mask):
"""
前向传播:定义数据如何流过网络
参数说明:
- q: Query矩阵,形状为 [batch, n_q, dimension_q]
表示"我想查询什么"
- k: Key矩阵,形状为 [batch, n_k, dimension_k]
表示"我能提供什么标签"
- v: Value矩阵,形状为 [batch, n_v, dimension_v]
表示"我的实际内容是什么"
- mask: 遮蔽矩阵,用于遮盖不该被注意到的位置
(例如:不让模型看到未来的词)
"""
# 获取输入张量的维度信息
# batch: 批次大小(一次处理多少个样本)
# n_q: Q序列的长度(有多少个查询词)
# dimension_q: 每个查询词的向量维度
batch, n_q, dimension_q = q.size()
# n_k: K序列的长度(有多少个键词)
batch, n_k, dimension_k = k.size()
# n_v: V序列的长度(有多少个值词)
batch, n_v, dimension_v = v.size()
# ---- 线性变换:生成多头的Q、K、V ----
# q经过fc_q变换:[batch, n_q, dimension_q] → [batch, n_q, num_head*d_k]
# 相当于:一次性为所有头生成Q矩阵
q = self.fc_q(q)
# k经过fc_k变换:[batch, n_k, dimension_k] → [batch, n_k, num_head*d_k]
k = self.fc_k(k)
# v经过fc_v变换:[batch, n_v, dimension_v] → [batch, n_v, num_head*d_v]
v = self.fc_v(v)
# ---- 重新排列维度,拆分成多个头 ----
# q的变换过程详解:
# 1. view(batch, n_q, num_head, d_k)
# 把最后一维拆开:[batch, n_q, num_head*d_k] → [batch, n_q, num_head, d_k]
# 类比:把"8个专家的分析结果"拆开成"每个专家各自的结果"
# 2. permute(2, 0, 1, 3)
# 调换维度顺序:[batch, n_q, num_head, d_k] → [num_head, batch, n_q, d_k]
# 把"头数"维度移到最前面,方便后续并行计算
# 3. contiguous()
# 确保内存连续(permute后内存可能不连续,需要此操作)
# 4. view(-1, n_q, d_k)
# 合并前两维:[num_head, batch, n_q, d_k] → [num_head*batch, n_q, d_k]
# 把所有头和所有批次合并,方便矩阵运算
q = q.view(batch, n_q, self.num_head, self.d_k).permute(2, 0, 1, 3).contiguous().view(-1, n_q, self.d_k)
# k和v做同样的维度变换
k = k.view(batch, n_k, self.num_head, self.d_k).permute(2, 0, 1, 3).contiguous().view(-1, n_k, self.d_k)
v = v.view(batch, n_v, self.num_head, self.d_v).permute(2, 0, 1, 3).contiguous().view(-1, n_v, self.d_v)
# ---- 计算注意力分数 ----
# torch.matmul(q, k.transpose(-1, -2)):
# q形状:[num_head*batch, n_q, d_k]
# k转置后:[num_head*batch, d_k, n_k]
# 矩阵相乘结果:[num_head*batch, n_q, n_k]
# 每个元素表示:某个查询词对某个键词的关注程度(原始分数)
# 除以sqrt(d_k):防止点积值过大导致softmax饱和,梯度消失
# 类比:把过于极端的分数缩放到合理范围
attention = torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(self.d_k)
# 扩展mask维度以匹配多头的数量
# mask原形状:[batch, n_q, n_k]
# repeat(num_head, 1, 1)后:[num_head*batch, n_q, n_k]
# 让每个头都使用同样的遮蔽矩阵
mask = mask.repeat(self.num_head, 1, 1)
# 将mask加到注意力分数上
# mask中被遮蔽的位置值为-inf(负无穷)
# 加上-inf后,这些位置经过softmax会变成0
# 效果:被遮蔽的位置完全不会被注意到
# 类比:把不该看的答案用黑色墨水涂掉
attention = attention + mask
# 对注意力分数做Softmax归一化
# 将原始分数转换为概率分布(每行加起来=1)
# 类比:把"关注程度原始分"转换为"关注比例百分比"
attention = self.softmax(attention)
# ---- 用注意力权重加权求和Value ----
# attention形状:[num_head*batch, n_q, n_k]
# v形状:[num_head*batch, n_v, d_v] (n_k == n_v)
# 相乘结果:[num_head*batch, n_q, d_v]
# 含义:每个查询词按注意力比例,融合所有值词的信息
# 类比:按各专家的参考程度,加权融合所有专家的实际内容
output = torch.matmul(attention, v)
# ---- 将多个头的输出拼接在一起 ----
# 1. view(num_head, batch, n_q, d_v)
# 恢复头数维度:[num_head*batch, n_q, d_v] → [num_head, batch, n_q, d_v]
# 2. permute(1, 2, 0, 3)
# 调换维度:[num_head, batch, n_q, d_v] → [batch, n_q, num_head, d_v]
# 把头数维度移到倒数第二位,准备拼接
# 3. contiguous()
# 确保内存连续
# 4. view(batch, n_q, -1)
# 拼接所有头:[batch, n_q, num_head, d_v] → [batch, n_q, num_head*d_v]
# 类比:把所有专家的报告横向拼接成一份合集
output = output.view(self.num_head, batch, n_q, self.d_v).permute(1, 2, 0, 3).contiguous().view(batch, n_q, -1)
# 最终线性变换:将拼接后的输出映射到目标维度d_o
# [batch, n_q, num_head*d_v] → [batch, n_q, d_o]
# 类比:总编辑把合集报告整合成标准格式的最终报告
output = self.fc_o(output)
# 返回注意力权重矩阵和最终输出
return attention, output
# ============================================================
# 多查询注意力机制(Multi-Query Attention, MQA)
# 与MHA的区别:所有头共享同一套K和V,只有Q是多头的
# 优点:大幅减少内存占用,加快推理速度
# 类比:多个专家用各自的问题(Q),但查阅同一份参考资料(K,V)
# ============================================================
class MQA(nn.Module):
def __init__(self, num_head, dimension_k, dimension_v, d_k, d_v, d_o):
"""
初始化多查询注意力机制
参数与MHA相同,但K和V只有一组(不再是num_head组)
"""
super().__init__() # 调用父类初始化方法
# 保存超参数
self.num_head = num_head # 头的数量
self.d_k = d_k # 每个头的K/Q维度
self.d_v = d_v # 每个头的V维度
self.d_o = d_o # 输出维度
# Q的线性变换层:和MHA一样,为所有头生成Q
# 输出维度:num_head * d_k(多头)
self.fc_q = nn.Linear(dimension_k, num_head * d_k)
# K的线性变换层:与MHA不同!只生成1组K(不乘以num_head)
# 所有头共享这一组K
# 输出维度:d_k(单头)← 关键区别!
self.fc_k = nn.Linear(dimension_k, d_k)
# V的线性变换层:同样只生成1组V(所有头共享)
# 输出维度:d_v(单头)← 关键区别!
self.fc_v = nn.Linear(dimension_v, d_v)
# 输出线性变换层:和MHA一样
self.fc_o = nn.Linear(num_head * d_v, d_o)
# Softmax归一化层
self.softmax = nn.Softmax(dim=2)
def forward(self, q, k, v, mask):
"""
前向传播(大部分与MHA相同,关键区别在K和V的处理上)
"""
# 获取输入维度信息(与MHA相同)
batch, n_q, dimension_q = q.size() # Q的批次、序列长度、向量维度
batch, n_k, dimension_k = k.size() # K的批次、序列长度、向量维度
batch, n_v, dimension_v = v.size() # V的批次、序列长度、向量维度
# ---- 线性变换 ----
# Q的变换:和MHA一样,生成多头Q
# [batch, n_q, dimension_q] → [batch, n_q, num_head*d_k]
q = self.fc_q(q)
# K的变换:与MHA不同!只生成1组K
# [batch, n_k, dimension_k] → [batch, n_k, d_k] ← 没有乘以num_head
k = self.fc_k(k)
# V的变换:同样只生成1组V
# [batch, n_v, dimension_v] → [batch, n_v, d_v] ← 没有乘以num_head
v = self.fc_v(v)
# ---- Q的维度重排(和MHA完全相同)----
# [batch, n_q, num_head*d_k] → [num_head*batch, n_q, d_k]
# 把多头Q拆分开,每个头独立处理
q = q.view(batch, n_q, self.num_head, self.d_k).permute(2, 0, 1, 3).contiguous().view(-1, n_q, self.d_k)
# ---- K和V的处理:与MHA完全不同!----
# MHA中:每个头有独立的K和V
# MQA中:所有头共享同一组K和V,通过repeat复制给每个头使用
# k.repeat(self.num_head, 1, 1):
# 把K复制num_head份,让每个头都能使用
# [batch, n_k, d_k] → [num_head*batch, n_k, d_k]
# 类比:把同一份参考资料复印num_head份,分发给每个专家
k = k.repeat(self.num_head, 1, 1)
# 同样地,把V复制num_head份
# [batch, n_v, d_v] → [num_head*batch, n_v, d_v]
v = v.repeat(self.num_head, 1, 1)
# ---- 以下步骤与MHA完全相同 ----
# 计算注意力分数:Q与K的点积,除以sqrt(d_k)缩放
# [num_head*batch, n_q, d_k] × [num_head*batch, d_k, n_k]
# → [num_head*batch, n_q, n_k]
attention = torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(self.d_k)
# 扩展mask到多头维度:[batch, n_q, n_k] → [num_head*batch, n_q, n_k]
mask = mask.repeat(self.num_head, 1, 1)
# 加上mask(被遮蔽位置变成-inf,softmax后变成0)
attention = attention + mask
# Softmax归一化,得到注意力概率分布
attention = self.softmax(attention)
# 用注意力权重对V做加权求和
# [num_head*batch, n_q, n_k] × [num_head*batch, n_v, d_v]
# → [num_head*batch, n_q, d_v]
output = torch.matmul(attention, v)
# 将多头输出拼接:[num_head*batch, n_q, d_v] → [batch, n_q, num_head*d_v]
output = output.view(self.num_head, batch, n_q, self.d_v).permute(1, 2, 0, 3).contiguous().view(batch, n_q, -1)
# 最终线性变换:[batch, n_q, num_head*d_v] → [batch, n_q, d_o]
output = self.fc_o(output)
# 返回注意力权重和最终输出
return attention, output
# ============================================================
# 测试代码:创建随机数据,验证MHA和MQA的运行结果
# ============================================================
# 定义超参数
batch = 10 # 批次大小:一次处理10个样本
num_head = 8 # 头的数量:8个"专家"同时分析
n_q, n_k, n_v = 2, 4, 4 # 序列长度:Q有2个词,K和V各有4个词
dimension_q, dimension_k, dimension_v = 128, 128, 64 # 词向量维度:Q和K用128维,V用64维
d_k, d_v, d_o = 16, 16, 8 # 每个头的维度:K和V用16维,输出用8维
# 创建随机输入数据(模拟真实的词向量输入)
# torch.randn:生成符合正态分布的随机张量
q = torch.randn(batch, n_q, dimension_q) # Q矩阵:[10, 2, 128]
k = torch.randn(batch, n_k, dimension_k) # K矩阵:[10, 4, 128]
v = torch.randn(batch, n_v, dimension_v) # V矩阵:[10, 4, 64]
# 创建因果遮蔽矩阵(Causal Mask)
# 用于自回归任务(如文本生成),防止模型看到"未来"的词
# torch.full:创建全部填充为-inf的矩阵
# 形状:[batch, n_q, n_k] = [10, 2, 4]
mask = torch.full((batch, n_q, n_k), -np.inf)
# torch.triu:保留矩阵的上三角部分,其余置0
# diagonal=1:从对角线上方1位开始保留(不包括对角线本身)
# 效果示例(n_q=2, n_k=4时):
# [[ 0, -inf, -inf, -inf], ← 第1个查询词:可以看第1个键,看不到后面的
# [ 0, 0, -inf, -inf]] ← 第2个查询词:可以看第1、2个键,看不到后面的
mask = torch.triu(mask, diagonal=1)
# ---- 测试多头注意力(MHA)----
# 实例化MHA模型
mha = MHA(num_head, dimension_k, dimension_v, d_k, d_v, d_o)
# 前向传播,得到注意力权重和输出
attention, output = mha(q, k, v, mask)
# 打印输出形状,验证结果
# attention形状应为:[num_head*batch, n_q, n_k] = [80, 2, 4]
# output形状应为:[batch, n_q, d_o] = [10, 2, 8]
print(attention.size(), output.size()) # 预期:torch.Size([80, 2, 4]) torch.Size([10, 2, 8])
# ---- 测试多查询注意力(MQA)----
# 实例化MQA模型
mqa = MQA(num_head, dimension_k, dimension_v, d_k, d_v, d_o)
# 前向传播,得到注意力权重和输出
attention, output = mqa(q, k, v, mask)
# 打印输出形状(与MHA相同,但内部K和V的处理方式不同)
# attention形状:[num_head*batch, n_q, n_k] = [80, 2, 4]
# output形状:[batch, n_q, d_o] = [10, 2, 8]
print(attention.size(), output.size()) # 预期:torch.Size([80, 2, 4]) torch.Size([10, 2, 8])
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)