说起 Self-Attention我见过太多人卡在同一个地方公式能背、代码能抄、面试能答它是用来建模全局依赖的但一旦被追问它到底在干什么、为什么这么算就立刻卡壳。这个标题写的是狗都能看懂的Self-Attention讲解我看完的第一反应是——这话说得不夸张。自注意力本身真的不难难的是大部分讲解一上来就摆矩阵、堆符号把一件本该能用手比划清楚的事讲成了线代期末考试。我自己带过几个刚入行的同学也用这套讲法跟完全不懂算法的朋友聊过反馈都还不错。所以这篇就把我这些年反复打磨的那套讲法写下来从它在解决什么麻烦开始到手算一遍、自己写一遍再到训练时的坑、以及最近图像超分方向把自注意力做稀疏自适应的思路一条线走完。1. 先别管公式Self-Attention 到底在解决什么麻烦1.1 一个找猫的场景假设一张图里有两个物体一只猫和一只狗模型要判断这只猫在干什么。它看到的猫可能只占了几十个像素而猫的动作信息——比如它在追着什么跑——其实藏在画面另一头的某个位置。传统卷积网络要建立这两处的联系得靠一层层卷积把感受野慢慢撑大中间要经过几十层信息在传递过程中被反复揉捏稀释。自注意力的做法很直接让画面里每一个位置都有机会直接看到所有其他位置并且根据当前需要动态决定该看谁、看多重。猫那个位置在判断动作时会给画面另一头那个关键位置分配一个很高的权重一步到位。这就是Attention的本意——注意力。而Self是说查询和被查询的对象来自同一份输入不是两段不同的东西。翻译成人话同一句话里的每个词、同一张图里的每个位置互相之间打分然后按分数高低互相借鉴信息。1.2 循环网络和卷积网络各自的死穴为什么非得要这个东西因为前两代主力都有明显短板。循环网络按顺序一个词一个词地处理第 100 个词想用上第 1 个词的信息中间得穿过 99 次状态传递。梯度在这条长链上反复相乘要么衰减到接近零要么爆炸这就是经典的长期依赖问题。而且它是串行的第 100 步必须等第 99 步算完GPU 再强也只能干瞪眼训练速度被死死限制住。卷积网络倒是能并行但它的视野是局部的想覆盖远距离靠的是堆层数。堆层数意味着参数变多、计算变重而且每一层的连接模式是固定的不管输入内容是什么卷积核的权重都长一个样。遇到这个位置该看哪里随内容变化的任务它没法自适应。自注意力把这两个问题一起解决了任意两个位置之间的路径长度是 1不随距离增长所有位置的计算可以并行权重是根据输入内容实时算出来的输入变了看谁不看谁也变。我第一次意识到这个差别是在做一个文本分类任务的时候。同样的数据循环网络调到第 12 层才勉强抓住跨句的指代关系换成自注意力结构4 层就稳了而且训练时间只有原来的三分之一。这个对比给我的冲击挺大的也是从那时起我认真去啃这套机制的。2. 把 Q、K、V 讲清楚一次注意力计算的完整走位2.1 查字典的类比这是我认为最好用的一个类比几乎没人听不懂。你去图书馆查一本书。你脑子里想的是我想找关于烘焙的书——这是Query查询。书架上每本书的书脊上都贴着标签烹饪烘焙历史——这些标签是Key键。你的查询和每个标签做匹配发现烘焙那个标签跟你的想法最贴合匹配度最高。于是你把这本讲烘焙的书拿下来读它的内容——这个内容是Value值。关键在于匹配分数决定拿多少内容而不是决定拿哪本书。如果烘焙匹配度 0.8烹饪匹配度 0.2那你最后读到的是两本书内容的加权混合0.8 份烘焙加 0.2 份烹饪。不是二选一是按比例调和。自注意力里的 Q、K、V 全是同一份输入 X 乘上三个不同的可学习矩阵得到的$$ Q XW_Q,\quad K XW_K,\quad V XW_V $$为什么要用三个不同矩阵因为我要找什么我是什么标签我携带什么内容这三件事需要不同的表示空间。如果 Q 和 K 用同一个矩阵那点积就变成了自相关表达能力会明显受限。这一点我在早期实现时偷懒试过把 Wq 和 Wk 设成同一个结果在几个数据集上准确率都掉了 1 到 2 个点之后再也不敢省了。2.2 四个动作一个都不能少一次完整的注意力计算拆开就是四步打分用 Q 和 K 做点积$QK^T$得到每个位置对每个位置的原始匹配分数。缩放除以 $\sqrt{d_k}$$d_k$ 是每个头的维度。归一化对分数做 softmax让每一行的权重加起来等于 1变成概率分布。加权求和用权重去乘 V 并累加得到输出。公式写出来就是那个几乎人人都见过的样子$$ \text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V $$我第一次看这个公式时觉得这也太简单了吧怀疑是不是漏了什么。后来才明白能这么简单还这么有效恰恰是它厉害的地方。它没有任何非线性激活、没有门控、没有循环结构就是两次矩阵乘法加一个 softmax。2.3 为什么要除以根号 d_k这个细节特别值得单独说因为它最容易被当成习惯写法抄过去实际上有明确理由。假设 Q 和 K 的每个元素都是均值 0、方差 1 的独立随机变量。那两个长度为 $d_k$ 的向量做点积结果是 $d_k$ 个乘积之和它的方差会变成 $d_k$。也就是说维度越高点积结果的数值范围越大。数值一大softmax 就会出问题大的那个变得更极端小的那个被压到接近 0。分布变得又尖又硬梯度几乎全被最大的那一项吃掉其余位置的梯度接近 0参数更新不动。除以 $\sqrt{d_k}$ 就是把这个方差重新拉回 1 附近让分布保持在一个软的状态。我实测过一个很直观的例子$d_k 512$ 的时候不缩放的点积最大值能到 20 以上softmax 出来的分布是[0.99998, 1e-5, ...]这种极端形态缩放之后最大值大概在 1 附近分布是[0.35, 0.28, 0.2, ...]这种健康形态。前者训练时 loss 曲线基本是平的后者才正常下降。2.4 手算一遍两个 token 的例子不手算一遍永远会觉得隔着一层。我设计一个最小的例子。输入 X 有两个 token每个 token 是 2 维X [[1, 0], [0, 1]]为了看清楚让 Wq Wk Wv 单位矩阵那么 Q K V X。第一步打分。QK^T [[1*10*0, 1*00*1], [[1, 0], [0*11*0, 0*01*1]] [0, 1]]第二步缩放。$d_k 2$$\sqrt{2} \approx 1.414$[[0.707, 0 ], [0 , 0.707]]第三步softmax。第一行$\exp(0.707) 2.028$$\exp(0) 1$和为 3.028得到[0.670, 0.330]。第二行对称[0.330, 0.670]。第四步加权求和。A [[0.670, 0.330], [0.330, 0.670]]输出就是out A V [[0.670*1 0.330*0, 0.670*0 0.330*1], [0.330*1 0.670*0, 0.330*0 0.670*1]] [[0.670, 0.330], [0.330, 0.670]]你看第一个 token 的输出变成了[0.670, 0.330]——它保留了自己 67% 的信息吸收了第二个 token 33% 的信息。这就是互相借鉴最朴素的样子。有意思的是这里每个 token 最关注的还是自己0.670 0.330这在真实训练初期很常见模型会先学我是我再慢慢学会看别人。如果你训练完发现注意力图还是这种接近对角线的形态说明它基本没学到东西这个后面在排查部分会详细说。3. 矩阵层面的三件事多头、位置编码、掩码3.1 多头注意力把一种看法变成多种看法单个注意力头只能计算一种相关性模式。但语言里同时存在好几种关系语法上的主谓一致、语义上的指代、位置上的邻近这些模式不一样用一个头去拟合会互相打架。多头注意力的解法是把 $d_{model}$ 维切成 $h$ 份每份独立算一次注意力最后把结果拼接起来再投影回原维度。比如 $d_{model} 512$$h 8$那每个头负责 64 维。注意拆分之后每个头的计算量和单头一样大总参数量也基本不变因为 $d_k$ 变小了$h$ 个头的总量还是那么多。这是一个几乎零成本的收益。我见过有人以为多头等于多算 8 次担心变慢其实拆分之后计算量是相同的。实现上最容易出错的地方是 reshape 的顺序。从(batch, seq, d_model)变成(batch, heads, seq, d_k)需要先 view 成(batch, seq, heads, d_k)再 transpose。如果顺序搞反得到的是把前 64 维分给头 1第 65 到 128 维分给头 2这种切法也能跑loss 也会降但效果会明显变差。这个坑我在第一次自己写的时候踩过当时怎么调参都上不去换回官方实现立刻好了 2 个点查了半天才发现是 reshape 顺序写错。3.2 位置编码自注意力天生不认顺序这里有个反直觉的事实标准自注意力对输入顺序是置换不变的。把一句话里词的位置全部打乱计算出来的输出集合是一样的只是顺序跟着变。因为点积和加权求和都不关心位置谁在前谁在后对它来说毫无区别。所以必须显式把位置信息注入进去。早期的做法是加正弦位置编码$$ PE_{(pos, 2i)} \sin\left(\frac{pos}{10000^{2i/d}}\right),\quad PE_{(pos, 2i1)} \cos\left(\frac{pos}{10000^{2i/d}}\right) $$import numpy as np def sinusoidal_pe(seq_len, d_model): pos np.arange(seq_len)[:, None] i np.arange(d_model)[None, :] angle pos / np.power(10000, (2 * (i // 2)) / d_model) return np.where(i % 2 0, np.sin(angle), np.cos(angle))用正弦函数有个好处任意位置的编码可以表示成其他位置编码的线性组合模型容易学会相对距离这种概念。后来主流的做法换成了可学习的位置嵌入或者旋转位置编码效果更好但核心思路一致——必须有一路信号专门告诉模型谁在谁前面。做视觉任务时这一点更明显。图像拉平成序列之后二维结构就丢了。我在一个分类任务里做消融实验去掉位置编码准确率直接从 91.2% 掉到 76.5%比我想象的严重得多。3.3 掩码被写反了也不报错的那个细节掩码mask有两种作用完全不同混用会出大问题。**填充掩码padding mask**用于变长序列。一个 batch 里句子长度不同短句要补齐到最长补齐的那些位置是无效的必须屏蔽掉否则它们会参与注意力计算污染结果。**因果掩码causal mask**用于生成任务。位置 $i$ 只能看到位置 $\le i$不能偷看未来否则训练时模型直接把答案抄了推理时反而不会生成。因果掩码就是一个下三角矩阵import numpy as np def causal_mask(seq_len): return np.tril(np.ones((seq_len, seq_len), dtypebool))实现里的写法通常是scores np.where(mask, scores, -1e9)这里有个细节要填-1e9 而不是 0。填 0 的话 softmax 之后那些被屏蔽的位置权重是 $\exp(0)1$反而变成了一个中等大小的权重屏蔽完全失效。要填一个足够大的负数让它 exp 之后接近 0。至于 -1e9 还是 -1e4实测差异不大但用 -1e9 更保险一点。不过要注意用 float16 的时候别填 -1e9会溢出成 -inf再把整个 softmax 变成 nan。这种 nan 特别难查因为它是从中间层冒出来的backward 时才报错。4. 自己写一遍比看十遍讲解管用4.1 二十行 NumPy 版本import numpy as np def softmax(x, axis-1): x x - x.max(axisaxis, keepdimsTrue) # 减最大值防溢出 e np.exp(x) return e / e.sum(axisaxis, keepdimsTrue) def self_attention(X, Wq, Wk, Wv, maskNone): # X: (batch, seq, d_model) Q X Wq K X Wk V X Wv d_k Q.shape[-1] scores Q K.transpose(0, 2, 1) / np.sqrt(d_k) if mask is not None: scores np.where(mask, scores, -1e9) A softmax(scores, axis-1) return A V, A这二十行里每一个符号都能对应到前面讲的那四步。我强烈建议你先不看任何库用这个版本在随机数据上跑一遍打印出 A 的形状和每一行求和看到它们都等于 1你对这套机制的理解会立刻落地。4.2 PyTorch 版本与形状检查清单import torch import torch.nn as nn import math class SelfAttention(nn.Module): def __init__(self, d_model512, n_heads8): super().__init__() assert d_model % n_heads 0 self.d_model d_model self.n_heads n_heads self.d_k d_model // n_heads self.wq nn.Linear(d_model, d_model) self.wk nn.Linear(d_model, d_model) self.wv nn.Linear(d_model, d_model) self.out nn.Linear(d_model, d_model) def forward(self, x, maskNone): b, n, _ x.shape # (b, n, d_model) - (b, n, h, d_k) - (b, h, n, d_k) Q self.wq(x).view(b, n, self.n_heads, self.d_k).transpose(1, 2) K self.wk(x).view(b, n, self.n_heads, self.d_k).transpose(1, 2) V self.wv(x).view(b, n, self.n_heads, self.d_k).transpose(1, 2) scores Q K.transpose(-2, -1) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask, -1e9) attn torch.softmax(scores, dim-1) out attn V # (b, h, n, d_k) out out.transpose(1, 2).contiguous().view(b, n, self.d_model) return self.out(out), attn调试的时候把每个张量的形状都打印出来跟这张表对照张量形状说明x(b, n, d_model)输入Q / K / V(b, h, n, d_k)每个头独立scores(b, h, n, n)注意力分数显存大户attn(b, h, n, n)每行和为 1out(b, h, n, d_k)拼接前我帮人排查问题时只要看到attn的行和不等于 1基本能立刻定位到是 softmax 维度写错了。dim-1是对最后一维归一化也就是每个位置对所有位置的权重和为 1写成dim-2就变成每个位置被所有位置关注的总和为 1完全是另一回事而且代码照样能跑不报错。4.3 怎么确认自己写对了三个自检手段按顺序做第一权重和检查。attn.sum(-1)应该全是 1误差在 1e-6 以内。第二等价性检查。把参数拷贝到torch.nn.MultiheadAttention里注意它的权重是打包在一起的需要手动搬同样的输入应该得到几乎一样的输出。误差在 1e-5 以内说明实现没问题。我第一次做这个对比时差了 0.3最后发现是缩放因子写成了 $\sqrt{d_{model}}$ 而不是 $\sqrt{d_k}$这个错误非常隐蔽。第三梯度流检查。跑一次 backward看每个参数的梯度是不是非零。如果某个矩阵梯度全为 0说明它没参与计算八成是拼写或者变量引用的低级错误。5. 训练时最容易踩的五个坑5.1 注意力图全是对角线训练几百步之后把注意力权重可视化出来发现每一行几乎都是[0, 0, ..., 1, ..., 0, 0]只有自己位置是 1。这种情况几乎可以断定模型没学到任何跨位置关系。我遇到过的原因有三个一是缩放漏了softmax 饱和二是学习率太大参数一开始就冲到了极端区域三是位置编码加错了比如把编码加到了 attention 的输出上而不是输入上。按这个顺序排查通常五分钟能定位。5.2 掩码方向写反这是个不报错但结果全错的经典问题。masked_fill(mask, -1e9)是把 mask 为 True 的位置填成 -1e9。所以你的 mask 里 True 的位置必须是要屏蔽掉的位置不是要保留的位置。我在两次不同的项目里都见过有人写反表现是训练 loss 正常下降因为模型能看到全部信息拟合能力更强但一到推理就崩生成质量惨不忍睹。原因是训练时它偷看了未来推理时看不到就蒙了。记住一个口诀True 表示这里要挡住。填充掩码里 True 对应补零位置因果掩码里 True 对应右上三角。5.3 显存和计算量O(n²) 的真实代价scores 的形状是(b, h, n, n)。序列长度翻倍这个矩阵的元素数量变成 4 倍。这不是理论上的担忧是实打实的显存杀手。序列长度headsscores 元素数float32 显存51282.1M8.4 MB102488.4M33.6 MB2048833.6M134 MB40968134M537 MB这还是单个样本、单层的量。batch 16、12 层堆起来就是另一个数量级了。而且反向传播时还要存中间激活实际占用是表里的好几倍。我处理长序列时的几个办法一是梯度检查点用时间换显存能省 60% 到 70%二是分块计算注意力不一次性算出完整的 n×n 矩阵而是分块算、边算边累加显存从 O(n²) 降到 O(n)三是真的没必要全看的时候就稀疏化这也是下一节要展开的方向。5.4 loss 不下降时的排查链路我自己的排查顺序基本按这个走先确认数据。attn的行和是不是 1输出里有没有 nan。如果是 nan多半是 -1e9 在低精度下溢出了换成 -1e4 试试或者改用库提供的 mask 接口。再看梯度。打印每个参数的grad.norm()如果全是 0前向和反向肯定断开了如果某个特别大比如上万说明初始化有问题把输出层权重的初始化标准差调小一点。然后做过拟合测试。拿 8 条样本反复训练正常情况下一两百步内 loss 应该降到接近 0。降不下去说明结构有问题不是数据量的问题。这一步能省掉大量瞎调参的时间我几乎每个项目都会先跑一遍。最后才去调超参。学习率、warmup 步数、权重衰减这个顺序不要反过来。6. 稀疏与自适应图像超分里那条更省的路6.1 超分任务对自注意力又爱又恨图像超分要做的事是从一张低分辨率图重建出高分辨率细节。它天然需要长距离依赖——一片草地的纹理、一堵墙的砖缝这些模式在整个画面里重复出现模型得能看到远处才能补得准。所以自注意力在这个任务上确实有用。问题是图像令牌的数量太恐怖了。一张 128×128 的特征图拉平就是 16384 个令牌scores 矩阵是 16384 的平方那是 2.7 亿个元素单个头都放不下。而超分任务的算力预算通常还很紧张因为它经常要跑在终端设备或者要处理视频流。这就逼着大家去想真的有必要每两个位置之间都算一次吗6.2 稀疏化省在哪答案是绝大多数位置对之间的相关性其实很低。一片天空的像素去关注画面角落的一块石头基本上是浪费。如果能让每个位置只跟它真正相关的少数几个位置交互计算量就能从平方级降下来。稀疏注意力的基本思路是先快速估一个分数只取分数最高的 k 个位置做完整的注意力计算其余直接舍弃。这样 scores 的规模从 $n \times n$ 变成 $n \times k$$k$ 取 16 或者 32 的时候省下来的量是几十倍。这里的难点在于先估一个分数本身也要花钱。常见的做法是用一个轻量的投影或者池化后的低维表示来粗筛代价远小于完整计算。我实测过一种简化实现用平均池化把空间维度降 4 倍做粗筛再回到原分辨率做精细注意力端到端速度提升了 2.3 倍画质指标只掉了 0.05 dB这个折中相当划算。6.3 自适应让网络自己决定看多少固定取 top-k 有个问题k 是个超参画面简单的地方可能 4 个就够纹理复杂的地方需要 64 个。一刀切要么浪费要么不够。自适应的做法是让网络自己预测这里需要看多少个。常见的实现是设定一个预算比如平均每个位置只看 24 个具体到每个位置简单的区域少看点复杂的区域多看总量控制住。用一个小的门控网络输出每个位置的稀疏度再加一个约束让总预算符合要求。这类思路之所以在超分里特别有效是因为图像的冗余度极高。平滑区域占了画面的绝大部分面积这些地方的注意力完全可以省掉把预算集中到边缘和纹理上。我在一个 4 倍超分任务上做过统计把 80% 的注意力预算分配给梯度最高的 20% 区域重建质量基本没有损失。6.4 一个可以跑起来的最小实现import torch import torch.nn as nn import torch.nn.functional as F class SparseAttention(nn.Module): 只在空间维度做 top-k 稀疏适合超分这类大特征图场景 def __init__(self, dim, topk32): super().__init__() self.topk topk self.to_qkv nn.Conv2d(dim, dim * 3, 1) self.proj nn.Conv2d(dim, dim, 1) def forward(self, x): b, c, h, w x.shape n h * w qkv self.to_qkv(x).reshape(b, 3, c, n).permute(1, 0, 2, 3) q, k, v qkv[0], qkv[1], qkv[2] # (b, c, n) # 用降采样的粗糙特征做快速筛选控制筛选成本 qc F.avg_pool1d(q, 4) kc F.avg_pool1d(k, 4) coarse torch.einsum(bcn,bcm-bnm, qc, kc) idx coarse.topk(self.topk, dim-1).indices # (b, n, k) # 只对这 k 个位置做完整注意力 kg k.gather(2, idx.reshape(b, -1).unsqueeze(1).expand(-1, c, -1)) vg v.gather(2, idx.reshape(b, -1).unsqueeze(1).expand(-1, c, -1)) scores torch.einsum(bcn,bck-bnk, q, kg) / (c ** 0.5) attn scores.softmax(dim-1) out torch.einsum(bnk,bck-bcn, attn, vg) return self.proj(out.reshape(b, c, h, w)) x这段代码里的关键点有三个粗筛用便宜的操作这里是平均池化加一次低维 einsum精算只发生在 top-k 的子集上以及最后的残差连接——分支只负责补细节主干信息不受影响。这也是我做所有注意力模块时的习惯残差几乎是保底的操作。有一点要提醒gather出来的kg、vg是按每个位置各自的索引取的所以不同位置看到的邻居完全不同这是稀疏注意力的核心也是它和局部窗口注意力最大的区别。窗口注意力是画格子格子内互相看稀疏注意力是每个位置自己挑几个看后者更灵活。7. 选型什么时候上自注意力什么时候绕开7.1 几种结构的对比结构感受野并行性计算复杂度适合的场景卷积局部靠堆层扩大好O(n)局部模式强、数据量中等循环理论全局实际受限差O(n)流式、低延迟要求自注意力全局一步直达好O(n²)长距离依赖强、数据充足稀疏注意力近似全局好O(n√n) 或 O(nk)高分辨率、算力受限状态空间模型全局好O(n)超长序列、流式这张表我在做技术选型时经常翻出来看。它不是绝对的但能帮我快速排除掉明显不合适的选项。7.2 小数据和短序列别急着上自注意力的参数量大、 inductive bias 弱也就是它靠结构先验带来的免费午餐很少。这意味着它需要更多的数据才能学好。如果你的训练集只有几千条样本序列长度也就二三十那我建议先用卷积或者简单的池化加全连接大概率效果更好而且训练更快。我手上有过一个反面案例一个 3000 条样本的短文本分类任务团队上来就搭了 6 层自注意力调了两周准确率 78%。后来换成两层卷积加全局池化半天做到 84%。不是自注意力不行是它在这个数据规模下发挥不出来。7.3 长序列的几条替代路线序列长度超过 4096 之后标准自注意力的成本就开始劝退人了。我试过并且觉得值得考虑的几条路一是局部加全局的混合大部分层用窗口注意力每隔两层插一层全局注意力。这样既保住了局部的效率又不会完全丢掉远距离信息。二是分块加摘要把序列切成块块内做完整注意力块间用每块的摘要向量做一层额外的注意力。成本从 O(n²) 降到接近 O(n)代价是跨块的信息传递多了一层间接。三是线性注意力通过改变计算顺序把 softmax 挪到别处让复杂度降到线性。代价是表达能力和标准注意力有差距在需要精确匹配的任务上会掉点。我个人的经验是如果只是想让成本降下来优先考虑第二种如果真的要处理上万长度的序列第三种才值得上。第一种介于两者之间工程上最好实现风险也最低。最后分享一个小技巧是我踩了几次坑之后养成的习惯在写任何注意力模块之前先用一个 2×2 的手算例子验证一遍。前面第 2.4 节那个例子我会在纸上或者终端里跑一遍看输出符不符合直觉。这一步五分钟但能挡掉后面几个小时的为什么效果不对的排查。注意力机制的所有坑几乎都能在这个最小例子上暴露出来——缩放漏了会看到极端分布mask 反了会看到不该有权重的位置有权重softmax 维度错了行和就不是 1。把它当成一个单元测试比任何调试工具都管用。