拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
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-Attention):query、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、V:

import 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, mask=None): # 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, mask=None): 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)

这里的view+transpose是核心技巧。注意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 是怎么想的

SE(Squeeze-and-Excitation)是通道注意力的经典代表,思路很直白:先对每个通道做全局平均池化,把 $H \times W$ 的空间信息压成一个标量,然后通过一个小 MLP 学习通道间的权重,最后乘回原特征。

class SEBlock(nn.Module): def __init__(self, channels, reduction=16): super().__init__() self.avg_pool = nn.AdaptiveAvgPool2d(1) self.fc = nn.Sequential( nn.Linear(channels, channels // reduction), nn.ReLU(inplace=True), 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_size=7): super().__init__() self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2) self.sigmoid = nn.Sigmoid() def forward(self, x): # 沿通道维做平均和最大池化 avg_out = torch.mean(x, dim=1, keepdim=True) # (b, 1, H, W) max_out, _ = torch.max(x, dim=1, keepdim=True) # (b, 1, H, W) concat = torch.cat([avg_out, max_out], dim=1) # (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, reduction=16, kernel_size=7): 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 插 CBAM,mAP 涨了约 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-LN(LayerNorm 放在子层输入前)比 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 和位置编码加上,你对手感的理解会完全不一样。代码写完跑通那一刻,之前所有模糊的地方都会清晰起来。

返回列表