ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

手写 Self-Attention:QKV、Softmax 与稀疏注意力

手写 Self-Attention:QKV、Softmax 与稀疏注意力 1. 先把 Self-Attention 拆成一个日常场景狗都能看懂这个说法我第一次听见的时候是有点不服气的直到我自己被一个实习生问住——他问我为什么注意力分数要做 softmax我张口就说因为要归一化然后他追问那为什么不直接除以总和我卡了三秒。那一刻我才意识到很多讲 Self-Attention 的科普都喜欢从公式出发Attention(Q,K,V) softmax(QK^T/√d_k)V这一行字背下来只要十秒但真正理解它为什么长这样可能要花一个月。这篇东西的目标很明确让任何一个会写 Python、懂一点矩阵乘法的人读完能自己从零手写出一份能跑通的 Self-Attention并且知道每一行为什么这么写出了问题该往哪儿查。它解决的问题说起来也很朴素——让模型在处理一个序列里的某个位置时能够自动判断该看哪儿。这件事在机器翻译里是翻译这个词的时候该参照源句的哪些词在图像超分里是补这个像素的时候该参考哪些区域的结构在推荐系统里是推这个商品的时候该看用户历史里的哪几笔行为。Self-Attention 之所以成了过去几年最通用的建模模块就是因为它把这件该看哪儿的判断权交给了数据本身而不是像卷积核那样由人手工设定只看左右各一个邻居。适合读这篇的人有三类第一类是刚接触 Transformer、想把公式背后的机制搞明白的学生第二类是工程上要用到注意力模块、但只会调nn.MultiheadAttention却不敢改的开发者第三类是像我这样被某个具体任务比如图像超分、长文摘要、语音对齐里的注意力效率问题卡住想搞清楚稀疏化到底在稀疏什么的老手。三类人的出发点不一样但底层需要理解的东西是同一套所以我打算从最直观的比喻一路讲到能落地的代码再顺手聊聊自适应稀疏注意力这个最近挺热的方向。1.1 一场会议室里的信息加权我比较喜欢的入门比喻是会议室讨论。假设会议室里坐了 8 个人每个人手里都有一份自己的信息。现在轮到第 3 号发言他要说出一句综合了全场信息的话但他不可能把 8 个人的话原样复述一遍于是他先做一件事给每个人打一个这个人对我现在要说的内容有多重要的分然后按照这个分数把所有人的信息加权混在一起。分数高的他说的话里就多带一点那个人的观点分数低的基本忽略。这就是 Self-Attention 的全部核心动作。第 3 号是序列里第 3 个位置的 token每个人手里的信息是序列里其他 token 的表示我该说点什么决定了他怎么打分。关键在于这个打分的依据不是固定的而是由内容决定的——如果第 3 个词是它指代的是前面某个名词那模型就希望它在前面那个名词上分配高权重如果第 3 个词是个纯粹的连接词那它可能只需要看看紧邻的上下文就够了。权重由内容动态计算这是自注意力区别于所有固定模式算子的根本特征。再往下拆一层每个人用来打分的依据其实分成了两套东西一套是他自己想找什么另一套是别人能提供什么。前者叫 Query查询后者叫 Key键而别人真正要贡献的内容叫 Value值。打分的过程就是拿我的 Query 去跟每个人的 Key 做匹配匹配度越高说明这个人的信息越符合我当前的需要那我就从他那里多取一点 Value。这个 QKV 三分的设计是整个机制里最容易被跳过、但其实最关键的一步后面我会专门用一节展开。1.2 卷积、循环网络和注意力各自的短板要理解 Self-Attention 为什么会被发明出来得先看看在它之前大家是怎么处理序列的。循环神经网络的路子是我一个一个往后读把前面读到的内容压缩进一个隐藏状态里。这个做法的问题很直观信息要沿着时间步一层层传递传到第 100 个位置的时候第 1 个位置的信息已经被反复挤压过很多次了早就变形了。而且它是串行的第 10 步必须等第 9 步算完GPU 上并行度很差长序列训练慢得让人想砸键盘。当然后来的 LSTM、GRU 加门控缓解了遗忘问题但串行的本质没变。卷积网络的路子是我用一个固定大小的窗口在序列上滑每次只看窗口内的几个邻居。它的优点是并行度高、局部模式抓得准但它有个硬伤感受野要靠堆层数来扩大。一个 3×3 的卷积核堆两层才看到 5 个位置堆十层才看到 21 个位置。如果两个词相隔 500 个位置你的网络得堆几百层才能让它们有一次交互机会这在工程上完全不现实。深度可分离卷积、空洞卷积都能缓解但谁跟谁该交互仍然是由人预先设定好的模式不是数据说了算。注意力机制走的是第三条路一步到位任意两个位置直接计算关联度。序列长度是 1000那第 1 个位置和第 1000 个位置之间的距离就是 1 次矩阵乘法不经过任何中间层。代价也很清楚——所有位置两两配对计算量和内存都是序列长度的平方级别。这个平方在后面讨论长序列和图像超分的时候会变成主角因为它既是注意力强大的来源也是它最要命的地方。我自己的经验是理解注意力的过程最好不要一上来就想着它比 RNN 好而应该想着它用平方级的代价换来了什么。想清楚这一点后面看到各种稀疏化、线性化的方案时你就知道它们本质上都是在往回收这个代价而不是在改变注意力的语义。1.3 一整套流程其实只有五步撇开所有的包装一个标准的单头 Self-Attention 前向过程真的只有五步我把它写下来给第一次接触的人一个整体的抓手输入序列 X形状是[batch, seq_len, d_model]分别乘以三个可学习的权重矩阵得到 Q、K、V 三份表示。用 Q 和 K 做矩阵乘法得到形状为[batch, seq_len, seq_len]的分数矩阵第 i 行第 j 列表示第 i 个位置对第 j 个位置的关注程度未归一化。对分数除以 √d_k 做缩放再按需要加上掩码。在最后一个维度上做 softmax把每一行变成一组加起来等于 1 的权重。用这组权重去加权求和 V得到输出再经过一个输出投影返回。五步里真正需要动脑筋的是第 2 步和第 3 步为什么是点积、为什么是缩放、为什么是 softmax 而不是别的归一化。这三问一旦通了剩下的都是工程细节。接下来我就按这个顺序把每一块拆开讲最后再给一份完整可运行的 PyTorch 代码。2. Q、K、V 到底是什么2.1 三个投影矩阵的分工以及为什么不能共用先回答一个新手最容易犯的迷糊既然 Q、K、V 都是从同一个输入 X 变来的那为什么不干脆直接用 X 自己省掉三个矩阵答案藏在不对称性里。考虑一句话这个苹果很好吃因为它很甜。当模型处理它这个词的时候它需要去找前面能解释它的名词——这时候它作为 Query 想要的是一种能被我指代的实体的特征而苹果作为 Key 提供的是一种我是个可以被指代的实体的特征至于苹果的 Value提供的是苹果这个词本身的语义内容。注意这三者的角色完全不同Query 表达的是需求Key 表达的是索引标签Value 表达的是实际内容。如果强行让三者共用一个表示就会产生一个别扭的结果——一个词用来被查找的特征和它用来提供内容的特征必须是同一套模型没法在这两个角色之间做区分。举个更技术的例子在代码补全任务里一个左括号(作为 Key 时的关键信息可能是我的类型是开括号、我在第 12 层嵌套但它作为 Value 时要提供的内容可能是这里开启了一个函数调用参数列表。这两种信息用同一套向量表达会互相挤占维度效果打折。所以三个线性投影各自独立是必要的。它们的形状通常是[d_model, d_k]其中d_k可以等于d_model也可以更小多头情况下通常是d_model / num_heads。有一点值得强调Q、K、V 的投影维度必须一致因为要做点积但投影的输出维度本身没有必须等于输入维度的要求这是很多人第一次读代码时会疑惑的地方。另外实践中这三个投影一般不带偏置或者带偏置差别不大但输出投影把多头拼接后的结果映射回d_model通常带偏置因为它承担的是融合的角色。顺手提一个工程上的经验如果你在做小规模任务参数量吃紧把所有注意力层的 QKV 投影合并成一个大矩阵是常见做法PyTorch 的in_proj_weight就是这么干的一个[3*d_model, d_model]的矩阵切三份用。这样做除了省一点显存更重要的是让一次矩阵乘法算完三个投影GPU 上的 kernel 启动次数少实际速度提升有时候比理论计算量节省更明显。2.2 为什么非得除以根号 d_k这一节我想认真写因为这是最容易被背下来但不懂的点。假设 Q 和 K 的每一个维度都是独立的随机变量均值 0、方差 1这是初始化做得比较规范时的近似情况。那两个长度为 d_k 的向量做点积结果是 d_k 个乘积项相加。每一项的均值是 0方差是 1两个独立标准正态相乘方差为 1。d_k 个这样的项相加点积的方差就变成了 d_k标准差是 √d_k。取 d_k 64标准差就是 8。也就是说分数矩阵里的数值会在大约 -24 到 24 这个区间里晃±3 个标准差。把这个量级的数喂进 softmax 会发生什么softmax 是exp(x_i) / Σexp(x_j)指数函数对大的输入极其敏感。当分数在 20 多和 -20 多之间分布时exp(20) ≈ 4.85e8而exp(-20) ≈ 2e-9两者差了 17 个数量级。结果就是 softmax 的输出几乎变成一个 one-hot 向量——最大那个位置接近 1其余接近 0。这带来两个后果。第一个后果是梯度消失softmax 输出接近 one-hot 的时候它的雅可比矩阵几乎全是 0反向传播回去的梯度被压得极小那些本该被稍微关注一下的位置永远学不到东西。第二个后果是训练不稳定早期的随机初始化下某个位置的分数偶然高一点就会形成赢家通吃注意力过早固化到一个错误的位置上模型很难再纠正回来。除以 √d_k 之后点积的方差被拉回 1 左右标准差恢复到 1分数落在 -3 到 3 这个舒服的区间softmax 输出是一个分布平滑、梯度健康的向量。这个操作的本质是方差归一化不是让数值变小一点这么简单。我把常用的几个 d_k 值对应的标准差列出来你感受一下量级d_k点积标准差 √d_k未缩放时分数大致范围缩放后范围164±12±3648±24±312811.3±34±325616±48±351222.6±68±3一个很自然的追问是为什么是除以 √d_k而不是除以 d_k如果除以 d_k点积的方差会变成 d_k / d_k² 1/d_k分数被压得太小softmax 输出接近均匀分布注意力就糊了等于没在关注任何东西。所以 √d_k 是唯一能让方差恰好回到 1 的缩放系数这不是拍脑袋定的经验值而是推导出来的。提示如果你在做温度采样的相关工作会看到 softmax 里加一个温度系数 τ形式上跟这个 1/√d_k 很像。区别在于这里的缩放系数是由维度决定的固定值用来保证数值健康而 τ 是人为调节分布尖锐程度的手段两者目的不同不要混为一谈。2.3 两种掩码不能搞混掩码是注意力实现里出错率最高的地方没有之一。我见过太多人在调试时盯着损失曲线怀疑人生最后发现是 mask 写反了。需要区分的是两种截然不同的掩码。第一种是填充掩码padding mask。一批数据里的句子长度不一样短的要用 0 补齐到同一长度。这些补齐位是没有任何语义的不应该参与注意力计算否则模型会去关注一堆空。做法是在分数矩阵上把所有指向填充位置的列置为负无穷softmax 之后这些位置的权重自然就是 0。它的形状通常是[batch, 1, 1, seq_len]其中1是广播出来的头维度和查询维度。第二种是因果掩码causal mask。自回归生成任务里预测第 t 个词的时候不能让模型看到 t 之后的内容否则就是作弊。做法是把分数矩阵的上三角不含对角线置为负无穷形状是[seq_len, seq_len]通常用torch.tril生成。它的几何意义很直白每个位置只能看到自己和左边的内容。两种掩码在实现上有一个共同的坑必须用masked_fill填负无穷不能填 0。因为 softmax 之前填 0 只是让分数变成 0而 0 经过 exp 是 1仍然会分到权重只有填负无穷exp 之后才是 0。新手写代码时经常写成scores scores * mask这在掩码位是 0、其他位是 1 的情况下勉强能用但一旦你的 mask 语义反过来有效位是 0就全错了而且这种错误在训练时表现得非常隐蔽——loss 会下降只是永远达不到应有的水平。还有一个进阶的坑负无穷在混合精度下可能变成 NaN。如果整个注意力行全部被掩掉比如某些实现里 padding 全在一行softmax 的分母是 0就会产生 NaN。稳妥的做法是用一个很大的负数比如-1e9代替float(-inf或者确保每一行至少有一个有效位置。注意PyTorch 的nn.MultiheadAttention用的是attn_mask和key_padding_mask两个参数语义跟上面说的两种掩码一一对应但前者用布尔值或浮点值都有不同含义True 表示屏蔽还是保留取决于你传的是 bool 还是 float这是官方文档里最容易读错的地方。自己写一遍就知道为什么要小心了。2.4 多头注意力到底多在哪多头注意力这个名字容易让人以为是把同一个注意力算了好几遍然后平均其实完全不同。它的做法是把d_model维的表示切成h份每份d_k d_model / h维每一份独立做一次完整的注意力计算最后把h个输出拼回d_model维。这里最关键的认知是每个头拥有独立的 Q、K、V 投影矩阵所以它们看到的是同一个输入的不同子空间。为什么要这么做单个头在做 softmax 的时候注意力权重是一组和为 1 的分布也就是说它在每个位置上只能表达一种关注模式。但语言里的依赖关系往往同时存在多种一个头在跟踪主谓一致另一个头在跟踪指代关系第三个头在盯着相邻词的搭配。如果只有一个头这些模式会被迫挤在同一组权重里互相妥协最后谁都表达不清楚。多头相当于给模型提供了多个并行的观察视角。每个头的维度变小了比如 512 维切成 8 个头每个头 64 维计算量与单头持平——因为虽然做了 8 次但每次的矩阵是 1/8 宽。这是一个很优雅的设计表达能力的提升几乎是免费的。实践中有一个常见的误解需要澄清不是头越多越好也不是头越少越好。头数太少多种模式挤在一起头数太多每个头的维度太小点积的表达能力下降而且不同头之间容易学出冗余的模式你在可视化注意力图时经常能看到好几个头长得几乎一样。经验上d_model / num_heads落在 32 到 128 之间比较稳低于 32 就要警惕了。另外还有个现象值得一提多头注意力的不同头确实会自发分工。在可视化里经常能看到靠近底层的头更关注局部相邻位置中层的头出现明显的句法依赖模式比如指向动词的主语高层的头则倾向于关注语义相关的远距离内容。这种分工不是你设计的是训练自己长出来的。我自己在做一个小规模翻译任务时观察过8 个头里有 2 个头的注意力图几乎是斜对角线的说明它们退化成了只看前一个词的作用其实已经跟卷积差不多了。3. 从零手写一遍3.1 张量形状约定与准备工作在写代码之前先把形状约定讲死因为绝大多数维度报错都源于这里的不一致。我习惯用这一套Bbatch sizeL序列长度seq_lend_model模型隐藏维度h注意力头数d_k每个头的维度等于d_model // h输入张量形状是[B, L, d_model]。经过 QKV 投影后为了做多头并行我们会把d_model拆成h × d_k然后把头维度提到前面[B, d_model] - [B, L, h, d_k] - [B, h, L, d_k]。这一步用view加transpose实现注意view之后必须transpose而不能直接reshape因为你要保证d_model的切分方式是前 d_k 维是第 1 个头接着 d_k 维是第 2 个头view是按最后一维连续切分的符合要求。分数矩阵的形状是[B, h, L, L]这是整个模块里内存占用最大的张量。输出的形状要还原回[B, L, d_model]还原过程是transpose(1, 2)之后contiguous().view(B, L, d_model)。这里的.contiguous()千万别省transpose之后内存不连续直接view会报错。环境上PyTorch 1.9 以上都行我用的是 2.x。如果想要查看注意力权重记得把模型设成eval()模式并关掉 dropout否则每次前向的注意力图都不一样可视化会看起来一团糟。3.2 单头版本逐行拆解先写最朴素的单头版本代码短每一行都能对上前面讲的原理。import math import torch import torch.nn as nn import torch.nn.functional as F class SingleHeadAttention(nn.Module): def __init__(self, d_model, d_k): super().__init__() self.d_k d_k self.w_q nn.Linear(d_model, d_k, biasFalse) self.w_k nn.Linear(d_model, d_k, biasFalse) self.w_v nn.Linear(d_model, d_k, biasFalse) self.w_o nn.Linear(d_k, d_model, biasTrue) def forward(self, x, maskNone, return_attnFalse): # x: [B, L, d_model] q self.w_q(x) # [B, L, d_k] k self.w_k(x) # [B, L, d_k] v self.w_v(x) # [B, L, d_k] scores torch.matmul(q, k.transpose(-2, -1)) # [B, L, L] scores scores / math.sqrt(self.d_k) # 方差归一化 if mask is not None: # mask 中 1 表示保留0 表示屏蔽 scores scores.masked_fill(mask 0, float(-inf)) attn F.softmax(scores, dim-1) # 每行和为 1 out torch.matmul(attn, v) # [B, L, d_k] out self.w_o(out) # [B, L, d_model] if return_attn: return out, attn return out几个细节值得停下来看。k.transpose(-2, -1)这个写法比写死transpose(1, 2)更通用因为前者不关心前面有几个维度将来加头维度时不用改。缩放那一步我特意拆成两行写就是为了让读代码的人看见这个操作有些人喜欢写成一行q k.transpose(-2,-1) / math.sqrt(d_k)紧凑但容易在 review 时被忽略。masked_fill那行里的mask 0是我个人偏好的语义写法——用 1 表示这个位置有效。但要注意很多数据集或框架返回的 mask 恰好相反用 True 表示需要屏蔽。接手的代码里一定要先打印几个样本确认语义别凭直觉假设。输出投影w_o带偏置前三个不带。这不是硬性规定但我见过的多数实现是这么做的理由是输出投影承担的是把多头信息融合回原空间的任务偏置能提供一点额外的表达自由度而 QKV 投影的作用是构造中间表示后面马上要做点积偏置的影响会被相对削弱。当然这也是可以调的超参不放心的话两种都试试。踩坑备忘masked_fill之后立刻 softmax如果某一行全部被屏蔽输出就是 NaN。写完之后拿一个全屏蔽的样例测一下看看你的代码是崩溃还是静默产生 NaN早发现早处理。3.3 多头版本与官方实现对齐单头理解清楚之后多头就是在维度上做文章。class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model 必须能整除头数 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads 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.w_o nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) self._reset_parameters() def _reset_parameters(self): for proj in (self.w_q, self.w_k, self.w_v, self.w_o): nn.init.xavier_uniform_(proj.weight) if proj.bias is not None: nn.init.zeros_(proj.bias) def _split_heads(self, t): # [B, L, d_model] - [B, h, L, d_k] B, L, _ t.shape return t.view(B, L, self.num_heads, self.d_k).transpose(1, 2) def forward(self, x, attn_maskNone, key_padding_maskNone, return_attnFalse): B, L, _ x.shape q self._split_heads(self.w_q(x)) # [B, h, L, d_k] k self._split_heads(self.w_k(x)) v self._split_heads(self.w_v(x)) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if attn_mask is not None: # 形状 [L, L] 或 [B*h, L, L]True 表示屏蔽 scores scores.masked_fill(attn_mask, float(-inf)) if key_padding_mask is not None: # 形状 [B, L]True 表示该位置是 padding scores scores.masked_fill(key_padding_mask[:, None, None, :], float(-inf)) attn F.softmax(scores, dim-1) attn self.dropout(attn) out torch.matmul(attn, v) # [B, h, L, d_k] out out.transpose(1, 2).contiguous().view(B, L, self.d_model) out self.w_o(out) if return_attn: return out, attn return out这份实现和 PyTorch 官方nn.MultiheadAttention在数值上是可以对齐的方法是对着拷贝权重。官方把 QKV 三个投影拼成了一个[3*d_model, d_model]的大矩阵我们切一下就行torch.manual_seed(0) d_model, num_heads, B, L 64, 8, 2, 10 mha nn.MultiheadAttention(d_model, num_heads, batch_firstTrue).eval() mine MultiHeadAttention(d_model, num_heads).eval() with torch.no_grad(): W mha.in_proj_weight # [3*d_model, d_model] Wq, Wk, Wv W.chunk(3, dim0) mine.w_q.weight.copy_(Wq) mine.w_k.weight.copy_(Wk) mine.w_v.weight.copy_(Wv) b mha.in_proj_bias bq, bk, bv b.chunk(3, dim0) mine.w_q.bias.copy_(bq) mine.w_k.bias.copy_(bk) mine.w_v.bias.copy_(bv) mine.w_o.weight.copy_(mha.out_proj.weight) mine.w_o.bias.copy_(mha.out_proj.bias) x torch.randn(B, L, d_model) with torch.no_grad(): ref, _ mha(x, x, x) got mine(x) print(最大绝对误差:, (ref - got).abs().max().item())这里有两件事必须确认否则对不上会白折腾。第一是batch_firstTrue不加这个参数的话官方接口默认序列维度在前形状约定完全对不上。第二是eval()——nn.MultiheadAttention的 dropout 默认是 0但如果你的实现里 dropout 开了而模型没切换模式误差会很大。跑出来误差在 1e-6 量级就算对齐了。那为什么还要自己写因为当你需要改注意力模式的时候——比如加相对位置编码、换成稀疏注意力、给不同的头设置不同的窗口大小——官方的黑盒接口就不好用了。自己手里有一份可控的实现改起来是分钟级的事。3.4 用手算验证一遍建立直觉代码能跑通不代表理解了。我的建议是拿一个极小的例子手动算一遍看看数是不是符合预期。假设d_k 2序列长度是 3Q 和 K 分别是q torch.tensor([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]) k torch.tensor([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]])不缩放的点积矩阵是[[1, 0, 1], [0, 1, 1], [1, 1, 2]]。注意第三行第三列的 2它跟其他值不在一个量级。缩放系数是 √2 ≈ 1.414缩放后是[[0.71, 0, 0.71], [0, 0.71, 0.71], [0.71, 0.71, 1.41]]。可以看到经过缩放最大最小值的差距被压缩了softmax 之后的分布不会那么极端。再把q[2]的第一维从 1 改成 5然后比较缩放和不缩放两种情况下 softmax 输出的差别。你会看到不缩放时输出几乎是[0, 0, 1]缩放后是一个明显更平缓的分布。这个动手实验做过一遍比读十遍公式管用。还有一个小实验值得做把d_k从 2 改成 64用随机初始化的 Q、K 跑一遍打印scores.std()在不缩放和缩放两种情况下的值。不缩放时大约在 8 附近正好是 √64缩放后在 1 附近。这个数字直接把前面的方差推导验证了比什么都直观。4. 常见问题与排查技巧实录4.1 症状对照表先看现象再查代码注意力模块出问题时症状往往不直接指向原因。我把自己和同事踩过的坑整理成一张对照表出问题的时候按症状往下查比盲目打印张量高效得多。症状最可能的原因排查动作loss 完全不降或者降得很慢缩放项漏了softmax 饱和打印 softmax 输出的最大最小值看是否接近 0/1loss 前期正常某个点突然爆成 NaNmask 里出现整行全屏蔽检查每行 mask 的有效元素个数是否为 0验证集效果好但推理时输出不变模型没切 evaldropout 在作怪确认.eval()和torch.no_grad()生成的文本重复同一个词因果掩码没加模型看到了未来打印注意力矩阵看上三角是否非零训练正常但复现不出同样结果dropout 未固定种子或多头顺序问题固定随机种子检查权重初始化显存随序列长度暴涨O(L²) 注意力矩阵计算B*h*L*L*4字节数考虑梯度检查点或稀疏化同一批数据两次前向结果差异大BatchNorm 或者 LayerNorm 位置不当确认归一化层在残差之前还是之后这张表里的前四项我基本上每次接触一个新的注意力实现都会遇到至少一个。特别是loss 不降和因果掩码漏加它们的共同特点是训练过程看起来没问题只是效果差很容易被归结为数据不够或者模型太小白白浪费几天时间。4.2 掩码写反的三种典型表现掩码问题之所以难查是因为它不会报错只会静静地让模型学不好。我把三种典型表现展开说一下。第一种是该屏蔽的没屏蔽。比如 padding 位置的 mask 没生效模型在算注意力的时候把补齐的 0 也当成正常内容去加权末尾位置永远在关注一堆空白。这时候你会发现短句子的效果明显不如长句子因为短句子补的 padding 占比更大。排查方法很简单构造一个长度差异很大的 batch观察同一句话单独跑和放进 batch 里跑的输出差异如果差异显著就是 padding 处理有问题。第二种是不该屏蔽的屏蔽了。最常见的是因果掩码用错方向——把下三角屏蔽了结果每个位置只能看到未来的内容。这种错误在训练时 loss 可能下降得还挺快因为模型确实看到了答案但一到推理阶段就完全崩掉因为它什么都没有。判断方法是在推理时打印第一个位置输出的 logits如果它猜出了后面才会出现的内容基本可以确定掩码反了。第三种是形状广播出的错。比如key_padding_mask的形状是[B, L]但你没有手动补齐成[B, 1, 1, L]就直接跟[B, h, L, L]的分数矩阵做运算PyTorch 的广播规则可能恰好让你的代码跑通但语义错误。这种情况尤其危险因为不报错。我的习惯是每次新增 mask 都显式地写unsqueeze哪怕理论上广播能自动处理写清楚比省几个字符重要得多。提示调试掩码最快的方法不是看代码是构造一个只有 3 个位置、其中 1 个是 padding的极小输入手动跑一遍把注意力矩阵打印出来看。矩阵的形状和数值对不对一眼就能看出来。4.3 显存和速度那个甩不掉的平方前面说过注意力的成本是平方级的具体是多少我们算一笔账。注意力矩阵的形状是[B, h, L, L]。用 fp32 存储一个元素 4 字节单个矩阵的内存占用是B × h × L² × 4字节。取一组常见配置B8h12对应 768 维、12 头L512算下来是8 × 12 × 512² × 4 100,663,296字节大约 96 MB。看起来还好。但把这个张量换成 L2048、h16、用 fp16 存就是8 × 16 × 2048² × 2 1,073,741,824约 1 GB。而且训练时不止存一份——前向要存 softmax 的输入和输出反向还要存中间量实际占用是这个数字的好几倍。序列再翻一倍到 4096占用变成 4 倍直接就是 4 GB 级别单卡基本放不下。这就是为什么长文本建模一直是注意力机制最头疼的问题。工程上的应对手段按投入产出比排序大概是这样的梯度检查点gradient checkpointing用时间换空间把中间激活值不存反向时重算。序列长但 batch 小的时候很有效代价是训练慢 20% 到 30%。混合精度训练注意力矩阵用 fp16/bf16 存直接省一半显存。要注意 softmax 前的数值范围负无穷记得换成大负数防止溢出。FlashAttention 这类融合算子核心思路是不把完整的L × L矩阵写回显存而是在 SRAM 里分块计算边算边累积。这是目前最有效的手段之一很多框架已经内置改一行调用就能用上。稀疏化直接减少参与计算的元素个数把 O(L²) 降到 O(L·k)。这一块是接下来要重点聊的方向。顺带说个实测经验注意力层的实际耗时瓶颈往往不在矩阵乘法本身而在内存搬运。L × L的矩阵算出来要写回 HBM再读进来做 softmax再写回再读进来跟 V 相乘来来回回好几趟。这也是为什么 FlashAttention 这类算子融合的方案能带来好几倍的加速——它减少的是访存次数不只是计算量。5. 进阶方向自适应稀疏 Self-Attention 在解决什么5.1 把稠密注意力的浪费点说清楚稠密注意力有一个隐含假设每个位置都可能跟其他所有位置有关联所以全都算一遍。但真实数据里这个假设明显过强。拿自然语言来说一个句子里大部分词跟远处的词没什么关系。我在可视化里看过很多次一张[L, L]的注意力图往往是稀疏的——权重集中在几条对角线带局部上下文、几个垂直/水平的条带少数几个被反复关注的 token比如标点和特殊的句首标记其余大片区域接近 0。既然算出来是 0那算它的过程就是纯粹的浪费。图像超分场景里的浪费更夸张。一张 512×512 的图如果按 patch 切分成 16×16 的 token就是 1024 个 token稠密注意力矩阵是 1024×1024一百万个元素。但对超分任务来说绝大多数平坦区域墙、天空、皮肤的像素只需要看局部邻域就能很好地补出来真正需要长距离信息的是那些有重复纹理和结构的地方——比如一排窗户、一段栅栏、一块砖墙。让所有位置都付出平方级的代价只为了少数位置能获得长距离感受野这个账明显不划算。这就引出了稀疏注意力的基本思路让每个位置只跟少数几个最相关的位置交互。如果每个位置只跟 k 个位置算注意力复杂度就从 O(L²) 降到 O(L·k)。L 越大这个节省越可观——L 从 512 涨到 4096稠密方案的成本涨 64 倍而 k 固定为 64 的稀疏方案只涨 8 倍。5.2 稀疏化的几条主流路线稀疏化怎么选决定了方案的性质。我把它分成三类这个分类是我自己按选择依据整理的不严格对应某一篇论文。第一类是模式固定的稀疏。人为指定每个 token 看哪些位置比如只看向左右各 w 个邻居局部窗口或者每隔 d 个位置看一个空洞或者按块划分只在块内做注意力。这类方案实现最简单、最可控速度也确实快。缺点是它的稀疏模式是预设的、与内容无关的——如果某个任务恰好需要跨越很远的关联固定模式就抓不到。Swin Transformer 用的窗口注意力就是这一类的代表在视觉任务里效果很稳。第二类是内容驱动的稀疏也叫自适应稀疏。让模型自己判断该关注哪儿。常见做法有两种一是给每个 query 算一遍粗略的相关性分数然后只保留 top-k 个 key 参与后续的精确计算相当于两阶段的检索二是学习一个门控或者阈值把低于阈值的注意力权重直接砍掉。这类方案的关键词就是自适应——稀疏的模式随输入内容变化不同的图、不同的句子走不同的路径。代价是引入了额外的计算算 top-k 本身要开销而且稀疏模式动态变化会影响 GPU 的并行效率理论 FLOPs 省了很多但实际加速可能有限这是工程上公认的难点。第三类是低秩或核近似。不直接稀疏化矩阵而是假设注意力矩阵可以用低秩分解近似或者用核函数把 softmax 换掉把计算顺序调整成先算K^T V再乘 Q复杂度变成线性。这类方案的数学上很漂亮但在实际任务上的效果经常打折扣因为低秩假设对某些任务并不成立。我这几年下来的感受是学术上更偏爱第二类工程上更偏爱第一类。前者有故事可讲、指标好看后者好实现、好调优、好部署。真正落地的时候往往是两者的混合——先用固定窗口把绝大部分计算限制住再在窗口内部做内容驱动的选择。5.3 图像超分里的自适应稀疏方案在做什么回到热词里提到的方向自适应稀疏注意力用于高效图像超分。这个组合挺有意思因为超分是个输出像素远多于输入的任务——输入一张 256×256 的图输出 1024×1024计算压力本来就大再叠上平方复杂度的注意力基本没法用。所以这个方向的核心矛盾就是如何在保住重建质量的前提下把注意力的开销压下来。按公开文献里比较常见的思路推演这类方案大体有这么几个着力点我按自己的理解梳理一下细节各篇可能不同这里讲的是共性的设计逻辑。第一是分辨率的非均匀对待。图像的平坦区域和纹理区域的信息需求天差地别。平坦区域丢了细节也没人看得出来纹理区域差一点点就糊了。所以很多方案会让注意力在空间上不均匀分配——平坦处走大窗口的粗粒度聚合甚至直接跳过纹理和边缘处保留细粒度的长程连接。这就是自适应最直接的体现稀疏的程度由内容决定不由位置决定。第二是动态的 top-k 或阈值剪枝。给每个 query token 算一遍与其他 token 的相似度只保留最相关的若干个参与计算。在超分任务里这个最相关往往对应着重复纹理——比如一扇窗户的某个 patch跟同一面墙上其他窗户的 patch 高度相似把它们连起来做注意力能显著提升结构一致性。而对一片纯色天空的 patch几乎不需要跨区域的信息。第三是稀疏模式的硬件友好化。这一条是纯工程。稀疏本身省了计算但如果稀疏的模式是随机的、不规则的GPU 的内存访问会变得极其碎片化实际速度可能比稠密还慢。所以真正能落地的方案通常会把稀疏约束在块级别或者固定粒度的分组上让每一步操作仍然是规整的矩阵运算。判断一个稀疏方案能不能用不能只看论文里报的 FLOPs 下降比例一定要看实际推理延迟。这句话我踩过坑一个理论上节省 70% 计算的方案实测只快了 15%因为访存模式太乱。第四是渐进式的注意力范围。在网络浅层用局部窗口因为浅层主要提取局部纹理深层逐步扩大范围甚至做全局聚合因为深层需要整合结构信息。这个思路借鉴了卷积网络的感受野设计经验在超分这种对空间结构敏感的任务上比较自然。如果你要动手复现这类方案我的建议是从最简单的一步开始先把稠密注意力换成固定窗口注意力跑通、调好把 baseline 建立起来。然后再往窗口里面加内容驱动的选择逻辑。一上来就实现完整的自适应稀疏很容易在调试稀疏索引的时候迷失方向最后连 baseline 是什么样都不知道。注意做稀疏注意力的时候务必确认你的实现里被剪掉的位置是真的不参与计算而不只是权重被置零。后者省不了任何计算只省了一点点有效信息是自欺欺人。检查方法是把剪枝后的实际 FLOPs 或者推理时间打出来看。我在实际项目里的体会是稀疏注意力这类工作理论分析和代码实现之间隔着一条很宽的沟。你需要在纸面上推导清楚复杂度降到了多少然后亲手量一遍真实的耗时两个数字对不上是常态而对不上的原因通常藏在内存访问模式里不在数学公式里。每次我遇到明明剪了一半为什么没变快的情况最后查出来的都是索引构造和数据重排的开销吃掉了计算上的节省。所以我的习惯是任何稀疏方案上线前先在一个固定的、有代表性的输入尺寸上把端到端时间量三遍取中位数再决定要不要往下一步走。
返回列表