☰
手算Self-Attention:QKV、多头与自适应稀疏注意力
2026/10/2 14:57:40 网站建设 项目流程

1. 为什么Self-Attention值得你花半小时搞明白

如果你最近两年翻过任何一篇深度学习论文,或者在技术群里潜水过一段时间,肯定见过Self-Attention这个词。它几乎成了Transformer架构的代名词,从自然语言处理一路杀到计算机视觉、语音识别、推荐系统,甚至在图像超分辨率这种看起来跟"序列"没多大关系的任务里,也开始大量出现它的身影,比如最近冒出来的自适应稀疏自注意力(adaptive sparse self-attention)就被用在了高效图像超分辨率上。问题在于,很多人第一次看到Q、K、V三个字母的时候,脑子是懵的——这仨到底是干嘛的?为什么点积一下就能"注意"了?softmax又是怎么冒出来的?

这篇内容就是冲着这个来的。我不打算堆公式吓人,也不打算从信息论讲起,而是用最直白的生活类比加上手算一遍的方式,把Self-Attention从"能看懂"带到"能自己写出来"。不管你是刚入门深度学习的小白,还是已经调过几次Transformer但一直没深究原理的工程师,看完之后应该都能在草稿纸上把整个过程推一遍。全文涉及的核心关键词就是Self-Attention,围绕它的机制、计算、实现和衍生思路展开,最后还会顺带聊聊它在图像超分这种非序列任务里怎么变形使用。

先说清楚这篇内容适合谁:如果你完全没接触过神经网络,可能需要先补一下矩阵乘法和softmax的基本概念;如果你已经会调用nn.MultiheadAttention但不知道里面发生了什么,这篇正好补上那块缺口;如果你是被"adaptive sparse self-attention"这类新词搞得一头雾水的算法工程师,第五部分会让你明白它的来龙去脉。我不讲废话,直接进入正题。

2. 用一个生活场景把Self-Attention的骨架搭起来

2.1 从"翻译一句话"开始理解注意力的动机

想象你在翻译一句话,英文是"The animal didn't cross the street because it was too tired."。中文翻译时要决定"it"指的是什么。你会自然而然地往回看,发现"it"大概率指"animal",而不是"street"。这个"往回看、判断该关注哪个词"的动作,就是注意力机制的直觉来源。

传统的循环神经网络(RNN)是把句子一个词一个词地读,前面的信息要靠隐状态一路传下来,传到最后容易丢。而Self-Attention的做法更直接:句子里的每个词都去跟包括自己在内的所有词"打招呼",算一下彼此的相关性,然后根据相关性大小把其他词的信息按比例拿过来融合到自己身上。这样"it"在编码自己的时候,会大量吸收"animal"的信息,少量吸收"street"的信息,最后它的表示就带上了正确的语义倾向。

关键点在于"自己跟自己打招呼"这件事。如果Query来自解码器、Key和Value来自编码器,那叫交叉注意力(Cross-Attention)。而当Q、K、V全部来自同一个序列的时候,就是我们说的Self-Attention——自己注意自己。

2.2 Query、Key、Value到底是什么关系

很多人卡在Q、K、V这三个名字上。我换个说法:

  • Query(查询):我想找什么。比如"it"这个词想知道"我到底指谁"。
  • Key(键):我身上贴的标签。每个词都挂一个标签,说明自己"是什么角色",比如"animal"的标签偏向"可被指代的生物"。
  • Value(值):我实际能提供的信息内容。一旦某个词被匹配上,它真正贡献出来的那部分信息就是Value。

用一个更生活的例子:你去图书馆找书。你心里想的检索词是Query,每本书书脊上的分类标签是Key,书里的实际内容是Value。你拿检索词去跟每本书的标签比对,匹配度高的书,你就把它的内容搬回家。Self-Attention就是把这件事对句子里每个词都做一遍,而且是并行的。

这三者其实都是从同一个输入向量线性变换出来的:给输入X分别乘三个可学习的权重矩阵W_Q、W_K、W_V,就得到Q、K、V。为什么不让它们直接用X本身?因为让模型自己去学"我该用什么方式提问、什么方式标记、什么方式提供内容",表达能力更强。这是设计上的一个关键取舍,用可学习参数换灵活性。

3. 手算一遍:Self-Attention到底是怎么算出来的

3.1 打分、归一化、加权求和三步走

整个Self-Attention的数学流程,拆开就是三步:

  1. 打分:用Q和K做点积,得到每个词对其他词的关注分数。scores = Q @ K.T
  2. 归一化:把分数除以√d_k(d_k是Key的维度),再过一个softmax,变成加起来等于1的权重。
  3. 加权求和:用这些权重去乘V,得到每个词的新表示。output = softmax(scores) @ V

写成公式就是大家最常见的那一行:

Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V

看起来简单,但每一处都有讲究。我们拿一组具体数字走一遍。

假设一句话有3个词,每个词用4维向量表示(方便手算)。为了不失一般性,我直接设定Q、K、V如下(实际中它们是学出来的):

Q = [[1, 0, 1, 0], [0, 1, 0, 1], [1, 1, 0, 0]] K = [[1, 0, 0, 1], [0, 1, 1, 0], [1, 0, 1, 0]] V = [[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]]

第一步,算QK^T。以第一个词(Q第一行[1,0,1,0])对第一个词(K第一行[1,0,0,1])为例,点积 = 1×1 + 0×0 + 1×0 + 0×1 = 1。对第二个词(K第二行[0,1,1,0]):1×0+0×1+1×1+0×0 = 1。对第三个词(K第三行[1,0,1,0]):1×1+0×0+1×1+0×0 = 2。所以第一个词的分数是[1, 1, 2]。

第二步,除以√d_k。这里d_k=4,√4=2。分数变成[0.5, 0.5, 1]。为什么要除?后面单独讲,先记住这个动作。

第三步,softmax。对[0.5, 0.5, 1]做softmax:

  • exp(0.5)=1.6487,exp(0.5)=1.6487,exp(1)=2.7183
  • 和 = 6.0157
  • 权重 = [0.274, 0.274, 0.452]

意思是第一个词在编码自己时,27.4%吸收自己的信息,27.4%吸收第二个词,45.2%吸收第三个词。

第四步,加权V。用[0.274, 0.274, 0.452]去乘V的三行:

  • 第一维:0.274×1 + 0.274×5 + 0.452×9 = 0.274 + 1.37 + 4.068 = 5.712
  • 后面三维同理可得:0.274×2+0.274×6+0.452×10 = 6.712;0.274×3+0.274×7+0.452×11 = 7.712;0.274×4+0.274×8+0.452×12 = 8.712

所以第一个词的新表示是[5.712, 6.712, 7.712, 8.712]。其余两个词照同样流程各算一遍。这就是一次完整的Self-Attention,全部是矩阵乘法,天然可以并行。

提示:手算时一定要自己列一遍,光看公式不会有肌肉记忆。我第一次真正理解就是在纸上把3×4的矩阵乘完那一下。

3.2 为什么非得除以根号d_k

这是面试里最常被问、也最容易答错的点。除以√d_k是为了防止点积结果过大,导致softmax进入饱和区。

假设Q和K的每个分量都是均值0、方差1的独立随机变量,那么它们的点积(d_k项求和)的方差就是d_k。d_k越大,点积的数值范围就越宽。当数值很大时,softmax的输出会趋近于one-hot,也就是某个位置接近1,其余接近0。这会导致两个后果:一是梯度几乎消失,因为softmax在饱和区导数极小;二是模型过早地把注意力死锁在个别位置上,丧失探索能力。

除以√d_k相当于把点积的方差重新拉回到1附近,让softmax工作在一个梯度健康的区间。这不是玄学,是可以推导的:Var(q·k) = d_k,除以√d_k后Var就变成1。所以这个除法不是随便找个数,而是精确地对应了"把方差归一化"这个目的。

我见过有人改成除以d_k,结果训练明显变慢甚至不收敛,就是因为过度压缩了分数差异。记住:是√d_k,不是d_k。

3.3 多头注意力不是玄学,就是多个人同时看

一个注意力头只能学到一种关注模式。但语言里的关系是多样的——有语法依赖、有指代关系、有语义相似,一个头忙不过来。多头注意力(Multi-Head Attention)的做法是:把Q、K、V在特征维度上切成h份,每份单独做一次Self-Attention,最后把结果拼接再线性变换。

好处很直观:不同的头可以关注不同类型的关系。有的头盯着主谓一致,有的头盯着远距离指代,有的头盯着邻近词。这跟卷积网络里多个卷积核各看各的纹理是一个思路。

代价是要控制好每个头的维度。假设模型维度d_model=512,头数h=8,那每个头的维度d_k = 512/8 = 64。总计算量跟单头差不多,但表达能力提升明显。这也是为什么Transformer原论文里d_model必须是头数的整数倍——不然切不均匀。

4. 亲手写一个Self-Attention:30行PyTorch代码

4.1 环境准备与依赖确认

代码部分我用PyTorch,装好基础的1.x或2.x任意版本都行。检查环境:

python -c "import torch; print(torch.__version__)"

如果你还没装,用pip装CPU版本足够跑通这个例子:

pip install torch numpy

不需要GPU,我们这个规模的数据在CPU上毫秒级出结果。为什么选PyTorch而不是别的?因为它的张量操作语义清晰,@直接对应矩阵乘法,跟公式一一对得上,读代码就像读数学。

4.2 从零实现核心计算

下面这段代码,我把每一步都对应到公式,方便你对照:

import torch import torch.nn as nn import torch.nn.functional as F import math class SelfAttention(nn.Module): def __init__(self, d_model): super().__init__() self.d_model = d_model # 三个线性层,把输入投影成Q、K、V 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.scale = math.sqrt(d_model) def forward(self, x): # x shape: [batch, seq_len, d_model] Q = self.W_q(x) # [B, L, D] K = self.W_k(x) # [B, L, D] V = self.W_v(x) # [B, L, D] # 打分:Q @ K^T,注意转置最后两维 scores = torch.matmul(Q, K.transpose(-2, -1)) # [B, L, L] # 缩放 scores = scores / self.scale # 归一化 attn = F.softmax(scores, dim=-1) # 加权求和 out = torch.matmul(attn, V) # [B, L, D] return out, attn # 测试一下 torch.manual_seed(0) x = torch.randn(2, 3, 4) # batch=2, 序列长度=3, 维度=4 sa = SelfAttention(d_model=4) out, attn = sa(x) print("输出形状:", out.shape) # 应为 [2, 3, 4] print("注意力矩阵:\n", attn[0]) # 第一个样本的3x3权重

关键细节说明:K.transpose(-2, -1)是把最后两个维度交换,这样[B, L, D] @ [B, D, L] = [B, L, L],得到每个位置对每个位置的分数。dim=-1的softmax保证每一行加起来是1。attn就是可视化注意力时画的那张热力图。

4.3 加掩码、加多头,把玩具变成能用的模块

上面是单头且不带掩码的版本。真正用在语言模型里还得处理两件事。

掩码(Mask):解码器不能看到未来的词,要做一个上三角为负无穷的掩码。

def forward_with_mask(self, x): Q, K, V = self.W_q(x), self.W_k(x), self.W_v(x) scores = torch.matmul(Q, K.transpose(-2, -1)) / self.scale # 上三角置为 -inf,softmax后变0 L = x.size(1) mask = torch.triu(torch.ones(L, L), diagonal=1).bool() scores = scores.masked_fill(mask, float('-inf')) attn = F.softmax(scores, dim=-1) return torch.matmul(attn, V), attn

多头:不需要自己写循环,PyTorch的nn.MultiheadAttention已经封装好了。

mha = nn.MultiheadAttention(embed_dim=64, num_heads=8, batch_first=True) x = torch.randn(2, 10, 64) out, weights = mha(x, x, x) # 自注意力,Q=K=V=x print(out.shape) # [2, 10, 64]

注意:batch_first=True是很多坑的来源。老版本默认序列维在前,写成[L, B, D],调不对形状时优先怀疑这里。

5. 从文本到图像:Self-Attention在超分任务里的变形

5.1 为什么图像超分也盯上了注意力

图像超分辨率(Super-Resolution, SR)要做的是把低分辨率图还原成高分辨率图,补出丢失的高频细节。传统卷积网络靠堆层扩大感受野,但卷积的"局部性"决定了它一次只看一小块,远处的纹理关系得靠很深的网络才能间接捕捉到。而Self-Attention天生就是全局的,每个像素位置都能跟图上任意位置直接发生关联,这对"纹理重复"这类现象特别有效——比如草地、砖墙、织物,规律性纹理在远处也能找到参照。

所以把Transformer搬进超分任务,直觉上是合理的。但问题也来了:图像的像素数远大于句子的词数。一张256×256的图展平就是65536个token,注意力矩阵是65536×65536,显存直接爆炸。这就是为什么原生的Self-Attention在图像任务里不能硬套。

5.2 自适应稀疏自注意力在解决什么

最近这类工作(adaptive sparse self-attention for efficient image super-resolution)的核心思路就是:别让每个位置都跟所有位置算注意力,只挑真正相关的少数位置算。

具体来说,大致有这么几个方向:

  • 窗口化:只在局部窗口内做注意力,比如Swin Transformer那样把图切成小块,在块内自注意力,再通过移位实现跨块交互。这是用局部性换效率。
  • 稀疏化:不预先定死窗口,而是让网络自己学"哪些位置值得关注"。可能是通过学习一个稀疏掩码,或者用top-k只保留分数最高的若干连接。
  • 自适应:稀疏模式不是固定的,而是根据输入内容动态变化。平坦区域可能不需要太多注意力,纹理密集区域才需要加大力度。这就叫"自适应"。

这三者叠加,就是"自适应稀疏自注意力"。它的价值在于把一个O(N²)的问题压到接近O(N),同时在关键区域保留全局建模能力。相比固定窗口,自适应稀疏能更好地处理图像里结构分布不均的情况。

代价是什么?实现复杂度上去了。稀疏操作往往涉及不规则索引,GPU上的并行效率不如规整的稠密矩阵乘法,写得不好反而更慢。所以工程上的功夫在于:怎么让稀疏既省内存又真的省时间。这是这类方法落地时最容易翻车的地方。

6. 踩坑记录与常见问题速查

6.1 新手最容易绕晕的六个概念

困惑点一句话解答
Q、K、V为什么不用输入本身用可学习投影能学出更适合提问/标记/提供内容的表示
为什么除以√d_k归一化点积方差,防止softmax饱和导致梯度消失
softmax为什么作用在最后一个维度因为要对"每个位置对所有位置的分数"做归一化,行和为1
多头是把QKV切分还是复制切分,每个头维度是d_model/h
Self-Attention和Cross-Attention区别Q、K、V是否来自同一个序列
位置信息去哪了Self-Attention本身无序,需要额外加位置编码

这张表建议存下来,遇到概念模糊时扫一眼。说实话,我见过太多人把多头理解成"复制多份做平均",那完全是错的——多头是切分维度各自独立计算,最后拼接再投影。

6.2 实操里的四个真实坑

坑一:维度对不上。最常见的报错是矩阵乘法维度不匹配。记住Self-Attention里Q和K的最后一维必须相同(那是d_k),V的维度可以和它们不同(d_v)。写代码时养成打印shape的习惯,比盯着报错猜快十倍。

坑二:softmax方向搞反。如果你发现注意力矩阵每列和为1而不是每行和为1,说明dim写错了。行代表"某个位置关注其他所有位置",所以归一化应该在最后那个维度。

坑三:忘了scale。有人写着写着把除法漏了,小数据上可能看不出问题,一上大规模训练就发散。建议把scale直接写进__init__里,别每次手写。

坑四:图像任务直接套文本实现。把图像展平成超长序列去算全局注意力,显存分分钟爆。要么用窗口,要么用稀疏,要么用线性注意力变体。这一点在做超分、检测、分割时尤其重要,别拿NLP的模板硬套。

提示:调试Self-Attention时,先用序列长度3、维度4这种迷你尺寸跑通,再放大。小尺寸能手动验证结果,出错了也容易定位。

最后分享一个我自己验证实现是否正确的小技巧:把输入设成完全相同的一批向量,如果实现正确,注意力权重应该趋近于均匀分布(因为没有哪个位置比其他位置更"特别")。如果权重明显偏斜,那多半是scale或者投影初始化有问题。这个自检办法救过我好几次,尤其在改动掩码逻辑之后。

再往深了走,你可以去研究线性注意力、Performer那类用核函数近似softmax的工作,以及前面提到的自适应稀疏路线。它们的共同目标都是一个:在尽量不损失全局建模能力的前提下,把计算和内存降下来。理解了最基础的这个Attention(Q,K,V) = softmax(QK^T/√d_k)V,这些变体看起来就不会那么面目可憎了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询