一、为什么需要注意力机制?

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(全连接层)叠加使用
📥 输入                    📤 输出
[] ──┐              ┌── [] ← 融入了所有词的信息
[] ──┤  Self-       ├── [] ← 融入了所有词的信息
[] ──┤  Attention   ├── [] ← 融入了所有词的信息
[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]= [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¹

= α'₁,₁ × 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拼接所有头的输出

= [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拼接
          [|| ... | 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])
Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐