ARTICLE DETAIL

资讯详情

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

注意力机制全解析:自注意力、多头、通道与空间注意力实战

注意力机制全解析:自注意力、多头、通道与空间注意力实战 1. 从一次模型调优说起注意力到底在算什么去年帮一个做时序预测的团队排查模型效果问题他们用 Transformer 做电力负荷预测训练集上 loss 降得很漂亮验证集却始终比一个简单的 LSTM 基线差一截。我把他们的模型代码拉下来看发现问题出在注意力层的实现上——他们直接照搬了一段网上抄来的注意力代码query、key、value 三个矩阵的维度没对齐mask 的位置也放错了导致模型实际上在偷看未来时刻的数据。训练时因为信息泄漏指标虚高推理时没有未来数据可看效果自然崩掉。这件事让我意识到一个很普遍的现象现在 Transformer 相关的教程铺天盖地但真正把注意力机制讲透、讲清楚每一步张量形状怎么变的并不多。大部分人停留在Q 乘 K 再 softmax 再乘 V这个公式层面一旦要自己改结构、加模块、排查问题就无从下手。所以这篇内容我打算把注意力机制这条线从最基础的版本一路讲到多头、通道、空间注意力把每一层的输入输出形状、设计动机、容易踩的坑都摊开来说。不管你是刚接触 Transformer 的新手还是已经在用 PyTorch、TensorFlow 做落地但总觉得理解不够扎实的从业者这篇应该都能帮你把这块知识补齐。核心关键词注意力机制、自注意力机制、多头注意力机制、通道注意力机制、空间注意力机制我会逐个拆解。2. 注意力机制的整体设计思路与选型考量2.1 为什么需要注意力这个抽象先说清楚一件事注意力机制不是从 Transformer 才有的。早在机器翻译的 encoder-decoder 架构里Bahdanau 那批人就已经用注意力来解决长序列信息压缩的问题了。当时的痛点是把一整个句子编码成一个固定长度的向量句子一长前面的信息就被稀释掉了。注意力机制的思路很朴素——解码每一步的时候不只看最终那个压缩向量而是回头去看编码器所有时刻的输出并且给每个时刻分配一个权重权重高的说明当前这一步更依赖那个位置的信息。这个分配权重的过程本质就是一次加权求和。你可以把它想成查字典手里拿着一个查询query然后去一堆键值对里找最匹配的键key匹配度越高对应的值value就分到越大的权重。这个类比是理解注意力最有效的入口后面所有变体基本都是在怎么算匹配度和怎么用这些权重上做文章。所以注意力要解决的核心问题就一个在信息很多的时候让模型学会有选择地关注而不是平均用力。这个思路在视觉、语音、时序、图数据上全都通用这也是它能成为通用组件的根本原因。2.2 三种主流注意力的分野标题里提到的这几种注意力其实对应三类不同的使用场景理解它们的分野比背公式重要得多。自注意力Self-Attentionquery、key、value 都来自同一个序列用来建模序列内部元素之间的依赖关系。它的特点是任意两个位置之间的距离都是 1不管隔多远一步就能建立联系这正好解决了 RNN 的长距离依赖问题。多头注意力Multi-Head Attention把自注意力并行做很多次每次用不同的投影矩阵让模型能在不同的子空间里分别关注不同类型的关系。通道注意力 / 空间注意力这两个主要出现在计算机视觉里。通道注意力关注哪些特征通道更重要空间注意力关注图像上哪些位置更重要它们和自注意力关注的维度不一样一个在通道维度上做加权一个在空间维度上做加权。很多人一开始会把这几类搞混觉得它们是一个东西的不同叫法。不是的。自注意力和多头注意力是处理序列的通道和空间注意力是处理特征图的它们的张量组织方式、加权轴、典型网络结构都不一样。下面我会分开讲。2.3 选型时的几个关键判断在实际落地里你需要根据任务类型选合适的注意力形式这里给几个我总结的判断依据任务类型推荐注意力形式原因文本/语音/时序建模多头自注意力需要建模长距离依赖多头能捕获多种关系图像分类/检测通道注意力SE、ECA计算量小插到 backbone 里就能涨点细粒度识别/分割空间注意力或 CBAM 组合需要定位关键区域图像生成/超分自注意力Swin、ViT需要全局感受野建模像素间关系多模态对齐交叉注意力query 和 key 来自不同模态我的经验是如果你不确定用哪种先从通道注意力SE 这类试起因为它便宜、稳、几乎不会让模型变差是最低风险的涨点手段。空间注意力和自注意力收益可能更高但调参成本和显存开销都更大。3. 自注意力机制核心细节与实现要点3.1 一步步拆解自注意力的计算自注意力的标准公式大家都会背$$\text{Attention}(Q,K,V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$但光背没用我把每一步的形状都写清楚。假设输入序列长度是 $L$每个 token 的嵌入维度是 $d_{model}$head 维度是 $d_k$。第一步线性投影得到 Q、K、Vimport torch import torch.nn as nn import math class SelfAttention(nn.Module): def __init__(self, d_model): super().__init__() self.d_model d_model 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) def forward(self, x, maskNone): # x: (batch, seq_len, d_model) Q self.W_q(x) # (batch, seq_len, d_model) K self.W_k(x) # (batch, seq_len, d_model) V self.W_v(x) # (batch, seq_len, d_model) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_model) # scores: (batch, seq_len, seq_len) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) # (batch, seq_len, seq_len) out torch.matmul(attn, V) # (batch, seq_len, d_model) return out, attn关键点在三处第一Q 乘 K 的转置得到的是一个 $L \times L$ 的矩阵。这个矩阵的每个元素 $(i, j)$ 表示位置 $i$ 的 token 对位置 $j$ 的关注程度。行是 query 位置列是 key 位置。搞不清这个方向mask 就很容易放反。第二除以 $\sqrt{d_k}$ 的作用。当 $d_k$ 比较大时Q 和 K 的点积结果方差会随维度线性增长数值过大会把 softmax 推到饱和区梯度几乎为零。除以 $\sqrt{d_k}$ 是把方差拉回到 1 附近保证 softmax 有正常的梯度。这个操作看起来小但没有它深层 Transformer 基本训不起来。第三softmax 是按最后一个维度做的也就是对每一行做归一化让每个 query 对所有 key 的权重加起来等于 1。3.2 位置编码为什么绕不开自注意力有个先天缺陷它是置换不变的。把输入序列的顺序打乱输出只是跟着打乱注意力算出来的值完全一样。也就是说模型本身不知道谁在前谁在后。这对语言、时序任务来说是致命的。解决办法就是位置编码。最经典的是正弦位置编码def sinusoidal_position_encoding(seq_len, d_model): pe torch.zeros(seq_len, d_model) position torch.arange(0, seq_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe # (seq_len, d_model)正弦编码的好处是它可以通过线性变换表示相对位置关系理论上能泛化到训练时没见过的更长序列。现在也有不少模型用可学习的位置嵌入如 BERT或相对位置编码如 T5、Swin各有取舍。注意位置编码是加到输入嵌入上的不是拼接到一起。拼接会让维度翻倍还得额外处理加的方案更简洁。3.3 实操中最容易出错的几个地方我在帮别人看代码时下面这几个错误出现的频率最高mask 方向反了。上面说的 scores 矩阵行是 query列是 key。要做因果 mask不能让位置 $i$ 看到 $j i$应该把上三角部分置为 $-\infty$。很多人从别的代码里抄了一个下三角的 mask一跑起来效果还行其实是在偷看未来。判断方法很简单把 mask 可视化出来看看是不是严格的上三角。padding mask 和 causal mask 混用没求和。一个 batch 里句子长度不同pad 的位置要 mask 掉同时语言模型还需要 causal mask。这两个 mask 需要正确地做与运算或广播相加漏掉一个都会出问题。softmax 之前忘了缩放。有些手写实现里直接把 $QK^T$ 丢给 softmax小模型可能看不出问题一旦层数深、维度大梯度就开始作妖。attention 矩阵没删显存爆了。如果只是为了可视化临时存 attention 权重推理时一定要关掉否则 $L^2$ 的矩阵在大序列上是显存杀手。4. 多头注意力机制的原理与工程实现4.1 多头不是在堆参数量很多人第一反应是多头注意力就是把自注意力做几遍再拼起来那不是白白增加了计算量吗其实不是。关键在于多头里的每个头只分到 $d_{model} / h$ 的维度。假设 $d_{model} 512$头数 $h 8$那每个头的 $d_k 64$。8 个头加起来的总计算量和单头用 512 维是基本相当的但表达能力不同了。单头只能在一个投影空间里算注意力多头相当于把特征切到 8 个不同的子空间每个子空间各自算自己的注意力模式最后拼回来再投影一次。这带来什么好处不同的头可以学到不同的东西。有人做过可视化发现有的头专门关注相邻词有的头关注语法依赖比如动词和它的主语有的头关注指代关系。这种分工是单头很难实现的。4.2 工程实现的两种写法多头注意力的实现有两种常见写法一种是显式循环一种是 reshape 批量算。生产代码几乎都用后者因为前者在 Python 层面循环会慢很多。class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0 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) def forward(self, x, maskNone): batch, seq_len, _ x.shape # 投影后拆头(batch, seq_len, d_model) - (batch, heads, seq_len, d_k) Q self.W_q(x).view(batch, seq_len, self.num_heads, self.d_k).transpose(1, 2) K self.W_k(x).view(batch, seq_len, self.num_heads, self.d_k).transpose(1, 2) V self.W_v(x).view(batch, seq_len, self.num_heads, self.d_k).transpose(1, 2) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) out torch.matmul(attn, V) # (batch, heads, seq_len, d_k) # 合并头 out out.transpose(1, 2).contiguous().view(batch, seq_len, self.d_model) return self.W_o(out)这里的viewtranspose是核心技巧。注意transpose之后张量在内存上不再连续后面如果需要view就必须先contiguous()不然会报错。这个坑我踩过不止一次。4.3 头数怎么选有没有经验值头数不是越多越好也不是越少越好。我的经验是头维度 $d_k$ 不要低于 32。低于这个数每个头能表达的东西太少注意力会退化得接近均匀分布。头数一般在 8 到 16 之间。像 BERT-base 用 12 头$d_k 64$GPT-3 这种大模型会用 96 头以上但那是因为 $d_{model}$ 本来就很大。小模型别硬堆头。如果你 $d_{model}$ 只有 128硬上 16 个头每个头就 8 维效果通常不如 4 个头。还有一个实际观察训练完之后很多头的注意力分布其实非常接近均匀也就是摸鱼头。有论文提出可以剪掉一部分头而不掉点这在推理加速里很有用。5. 通道注意力与空间注意力的原理与落地5.1 通道注意力SE 是怎么想的SESqueeze-and-Excitation是通道注意力的经典代表思路很直白先对每个通道做全局平均池化把 $H \times W$ 的空间信息压成一个标量然后通过一个小 MLP 学习通道间的权重最后乘回原特征。class SEBlock(nn.Module): def __init__(self, channels, reduction16): super().__init__() self.avg_pool nn.AdaptiveAvgPool2d(1) self.fc nn.Sequential( nn.Linear(channels, channels // reduction), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels), nn.Sigmoid() ) def forward(self, x): # x: (batch, channels, H, W) b, c, _, _ x.shape y self.avg_pool(x).view(b, c) # (batch, channels) y self.fc(y).view(b, c, 1, 1) # (batch, channels, 1, 1) return x * y # 广播相乘这里的reduction是压缩比控制中间层的宽度。默认 16 是个经验值太小计算量大太大又表达力不足。SE 之所以有效是因为它让网络显式地学习哪些通道对当前任务重要而不是让所有通道平均贡献。CAGrad、ECA 这些后续工作做了改进比如 ECA 直接用一维卷积代替 MLP避免了降维带来的信息损失参数量更小。实际用的时候如果你的 backbone 是 ResNet 系列直接插 SE 基本稳赚不赔。5.2 空间注意力告诉模型看哪里空间注意力关注的是特征图上哪个位置更重要。典型做法是沿着通道维度做池化得到一张 $H \times W$ 的空间图再通过卷积学出空间权重。class SpatialAttention(nn.Module): def __init__(self, kernel_size7): super().__init__() self.conv nn.Conv2d(2, 1, kernel_size, paddingkernel_size // 2) self.sigmoid nn.Sigmoid() def forward(self, x): # 沿通道维做平均和最大池化 avg_out torch.mean(x, dim1, keepdimTrue) # (b, 1, H, W) max_out, _ torch.max(x, dim1, keepdimTrue) # (b, 1, H, W) concat torch.cat([avg_out, max_out], dim1) # (b, 2, H, W) weight self.sigmoid(self.conv(concat)) # (b, 1, H, W) return x * weight关键设计是同时用平均池化和最大池化。平均池化保留了整体响应强度最大池化保留了最显著的特征点两者拼在一起送进卷积比只用一种要更鲁棒。这个思路来自 CBAM实测在细粒度分类和目标检测上都挺有效。5.3 CBAM通道和空间的串联CBAM 就是把上面两个模块串起来先做通道注意力再做空间注意力。顺序是先通道后空间因为通道注意力决定了用哪些特征空间注意力再决定用在哪些位置这个先后逻辑比较自然。class CBAM(nn.Module): def __init__(self, channels, reduction16, kernel_size7): super().__init__() self.channel_att SEBlock(channels, reduction) self.spatial_att SpatialAttention(kernel_size) def forward(self, x): x self.channel_att(x) x self.spatial_att(x) return x实操体会CBAM 插在 backbone 的每个残差块后面效果最好但会带来一定的延迟增加。如果对推理速度敏感可以只在 stage 之间插而不是每个 block 都插。我在一个移动端检测模型上试过只在最后两个 stage 插 CBAMmAP 涨了约 1.2 个点延迟只增加不到 3%性价比不错。5.4 通道-空间协同与自注意力的对比现在很多工作会把通道、空间注意力组合起来比如 BAM、Triplet Attention 这类。它们的共同思路是在不同维度上分别算注意力再融合。有的用加法有的用乘法有的并行算完再拼接。这里容易混淆的一点是空间注意力和图像上的自注意力有什么区别区别在计算复杂度和建模范围。空间注意力通常用卷积实现感受野是局部的、固定的而自注意力是全局的每个位置都能和所有位置交互代价是 $O((HW)^2)$ 的复杂度。Swin Transformer 的窗口注意力就是为了解决这个复杂度问题把全局注意力限制在局部窗口内再做窗口间信息传递。选的时候看你需要多大的感受野如果局部信息够用空间注意力更省如果需要建模远距离依赖比如大目标的整体形状那自注意力或 Swin 这种带层次结构的更合适。6. 常见问题与排查技巧实录6.1 训练不稳定、loss 反复震荡这是最高频的问题。按我的排查顺序一般从这几处入手检查缩放因子。确认 $\sqrt{d_k}$ 用的是 head 维度不是 $d_{model}$也不是 $d_{model}/h$ 之外的东西。搞错了缩放量级梯度会异常。检查初始化。Transformer 对初始化敏感Q、K、V 的权重如果初始方差过大第一层的 attention 就会接近 one-hot梯度很难回传。常用做法是把投影层权重用 Xavier 或小的正态分布初始化并保证残差路径上的初始化方差受控。加 LayerNorm 的位置。Pre-LNLayerNorm 放在子层输入前比 Post-LN 更稳定深层模型基本都用 Pre-LN。如果你从 Post-LN 换到 Pre-LN 发现不收敛可能是 warmup 没做好。warmup 一定要有。前几千步用线性 warmup学习率从小爬到大这一步是 Transformer 训练稳的关键省不得。6.2 注意力分布全是一个样如果可视化 attention 发现所有头、所有位置几乎均匀分布说明模型没学到东西。常见原因学习率太小或太大、mask 把所有位置都遮住了比如全 0 或全 $-\infty$、softmax 温度不合适。先确认 mask 图的正确性再调学习率。6.3 显存不够怎么办序列一长$O(L^2)$ 的 attention 矩阵就爆显存。几个实用手段方法原理适用场景Flash Attention分块计算不显式存完整矩阵训练和推理都能用首选梯度检查点不存中间激活反向时重算训练时省显存速度换空间稀疏注意力只算部分位置对长序列且依赖关系稀疏窗口注意力限制在局部窗口图像、长文本Swin/ Longformer 思路我的建议能上 Flash Attention 就直接上。PyTorch 2.0 之后有torch.nn.functional.scaled_dot_product_attention会自动选择高效实现几行代码就能替换掉手写的注意力速度快、显存省还避免了手写 mask 的常见错误。6.4 常见问题速查表现象可能原因解决方向验证集效果远差于训练集信息泄漏、mask 错误可视化 mask检查因果性loss 不下降学习率不合适、初始化有问题加 warmup改初始化降 lr梯度爆炸没缩放、没 LayerNorm检查缩放因子和归一化位置多卡训练结果不一致随机种子、batch 切分方式固定种子核对各卡数据通道注意力没效果插的位置不对、reduction 太大换位置调 reduction 到 8 或 16空间注意力过拟合模块太多、数据太少减少插入层数加正则7. 我自己踩过的一些坑和长期实践体会在注意力这块我攒了一些文档里不太会写、但实际很影响结果的经验这里一并说说。第一不要迷信堆模块。我早期做检测的时候看到哪个注意力模块涨点就往网络里加结果通道注意力、空间注意力、自注意力全堆上模型参数翻倍速度掉一半mAP 反而没涨。后来想明白注意力本质是一种特征重加权如果原始特征质量本来就差加多少注意力都没用。先把 backbone 和数据弄干净再考虑加模块。第二头维度比头数更重要。前面说过头数不是关键每个头分到的维度才是。我试过在 $d_{model} 256$ 的模型上从 8 头降到 4 头每个头从 32 维变 64 维效果反而更好。所以调多头的时候先保证 $d_k \geq 32$再在此基础上调头数。第三实现注意力之前先把 mask 画出来。这是最省时间的调试习惯。把 $L \times L$ 的 mask 矩阵用 imshow 画一下一眼就能看出是不是上三角、padding 位置对不对。比盯着代码猜快得多。第四如果只是想涨点优先试通道注意力。SE、ECA 这类模块插入简单、几乎不改变网络结构、参数增量小是最低风险的方案。空间注意力和自注意力收益可能更大但需要调参、调位置投入产出比不一定更划算。第五时序任务里要注意注意力对位置信息的依赖。纯正弦位置编码在时序预测里未必是最好的很多时候可学习的位置嵌入配合时间戳特征小时、星期、节假日等效果更好。这块我在电力负荷预测项目里实测过加了时间戳特征后注意力学到的东西明显更合理误差降了将近 8%。第六注意力和卷积不是对立的。现在很多高效模型是卷积和注意力混着用的浅层用卷积提取局部特征深层用注意力建模全局关系这样兼顾了效率和表达力。别一上来就想着全用 Transformer 替换卷积很多时候混合结构才是最优解。最后提醒一句注意力机制这套东西公式看十遍不如自己手写一遍。找一个简单的序列任务从零实现一遍单头自注意力、多头注意力再把 mask 和位置编码加上你对手感的理解会完全不一样。代码写完跑通那一刻之前所有模糊的地方都会清晰起来。
返回列表