拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

注意力机制PyTorch实现避坑:自注意力、多头、CBAM与Swin窗口

注意力机制PyTorch实现避坑:自注意力、多头、CBAM与Swin窗口

如果把注意力机制理解成一句"给不同的位置分配不同的权重",那基本可以跳过很多教材了。但真正上手写代码、训模型、调参的时候,你会发现坑根本不在这句话上,而在 Q、K、V 到底怎么投影、维度怎么摆、掩码会不会把整行打成 NaN、多头拆开之后显存为什么突然涨了一倍、SE 和 CBAM 到底该插在 backbone 的哪个位置。这篇是我在做 Transformer、注意力机制相关项目时攒下来的笔记整理,从最基础的自注意力机制一路讲到多头注意力机制、通道注意力机制、空间注意力机制,再顺带把时序注意力机制和 Swin 的窗口注意力边界理一遍。适合已经会写 PyTorch、想直接抄作业跑通实现的人,也适合刚学完 Transformer 原理但一写代码就报维度错的朋友。文中所有代码我都跑过,形状注释是真的形状,不是抄来的。

1. 开篇先理清:注意力机制到底解决了什么问题

1.1 从"固定权重"到"按需分配"的直觉过渡

在注意力机制普及之前,序列建模的主流是 RNN 和 CNN。RNN 的问题是信息必须沿着时间步一步步传递,第 100 个词想用上第 1 个词的信息,中间要经过 99 次状态更新,梯度早衰减没了。CNN 用卷积核扩大感受野,但感受野是固定的,堆叠层数越多,远距离依赖才勉强建立起来,而且权重在整个特征图上共享,对"这一帧该看哪一帧"这件事毫无自主性。

注意力机制的核心改动在于:权重要根据输入内容动态算出来,而不是训练好之后固定住。打个比方,RNN 像一条流水线上每个工位都按固定手顺往下传,注意力像开会,谁跟当前议题相关,谁的发言权重就高,相关性是当场根据内容算的,会开完权重就作废,下一句话重新算。这个"当场算"就是 query 和 key 做点积再 softmax 的过程。

所以注意力本质是一个可微的、内容驱动的软寻址机制。它同时解决了两个问题:一是路径长度,任意两个位置之间只需一步就能建立联系,梯度传播路径是 O(1);二是权重动态化,同一套参数在不同输入下能表现出完全不同的关注模式。代价也很直接,两两计算带来了 O(L²) 的复杂度,序列一长显存和算力就顶不住,这才有了后面窗口注意力、线性注意力这些变体。

理解这个取舍关系非常重要。你后面看到的几乎所有注意力变体,本质都是在"更强的表达能力"和"更低的计算成本"之间做平衡。SE 通道注意力机制砍掉空间维度只在通道上算,就是为了便宜;Swin 把全局注意力限制在窗口内,也是为了把 O(L²) 降到 O(L·M²)。抓住这条主线,各路变体就不会看着像一堆互不相关的黑盒。

1.2 一套贯穿全篇的记号约定与维度推演

维度混乱是初学者最大的痛点,所以这里先把记号钉死,后面全是这套。

约定 batch size 为 B,序列长度为 L,模型维度为 D,头数为 H,单头维度为 Dh。输入张量形状统一写成 [B, L, D],这是 PyTorch 里 nn.Linear 最顺手的排布。注意 TensorFlow/Keras 的 MultiHeadAttention 默认吃 [B, L, D],但内部 reshape 逻辑和 PyTorch 手写版不一样,跨框架迁移时最容易在这里翻车。

张量形状含义
x[B, L, D]输入序列
Q / K / V[B, L, D] 或 [B, L, Dh]投影后的查询、键、值
scores[B, H, L, L]注意力打分矩阵
attn[B, H, L, L]softmax 后的权重
out[B, H, L, Dh]加权求和结果
final[B, L, D]拼头并投影后的输出

有一件事必须提前记住:注意力权重矩阵的形状是 [B, H, L, L],它的元素总数是 B·H·L²。这意味着头数 H 增加时,Q/K/V 投影的参数量和主计算量基本不变,但注意力矩阵的显存是线性增长的。B=8、H=8、L=512、fp32 的情况下,单层注意力矩阵约 67MB,反向传播还要再存一份,实际占用接近 200MB。这个数字在 L=2048 时会膨胀到 16 倍,也就是 1GB 上下,一层而已。很多人调参时觉得"头数不影响计算量"就随手把 H 从 8 改成 16,然后显存爆了,问题就出在这里。

1.3 为什么建议先吃透自注意力再碰 Transformer

我见过不少人直接拿 HuggingFace 的 BertModel 跑微调,效果不错,但一被问到位置编码怎么加的、为什么是 pre-LN 而不是 post-LN、attention mask 和 key padding mask 有什么区别,就答不上来。这种状态下遇到自定义需求基本寸步难行,比如要改成流式推理、要加相对位置编码、要做稀疏注意力。

自注意力机制是整个 Transformer 的地基,它就是"Q、K、V 全部来自同一个输入"的最简形式。把它的张量流转、缩放因子、掩码处理这三件事搞透,多头注意力机制无非是加了拆头和拼头两步,交叉注意力无非是把 K、V 换成另一个序列。地基打牢,楼怎么盖都是顺的。

2. 自注意力机制:Q/K/V 从公式到张量形状的完整推演

2.1 Q/K/V 三个投影到底在干什么

标准公式写出来只有一行:Attention(Q,K,V) = softmax(QKᵀ/√d_k)V。但这里每个符号背后都有具体的物理含义,值得掰开讲。

Q 是 query,代表"我在找什么";K 是 key,代表"我有什么可以被匹配的特征";V 是 value,代表"如果匹配上了,我实际贡献什么内容"。Q 和 K 做点积得到相似度分数,softmax 归一化成概率分布,再用这个分布去加权求和 V。整个流程和数据库里的软检索几乎一模一样,唯一区别是这里返回的不是某一条记录,而是所有记录按相关度加权的混合结果,所以它对参数是可导的。

关键点在于,Q、K、V 都不是输入 x 本身,而是 x 经过三个独立的线性投影得到的。为什么必须投影?因为如果直接用 x 当 Q 和 K,那相似度就退化成 x 和 x 自身的点积,语义空间里"用来查询的方向"和"用来被查询的方向"被强行绑在一起,表达能力大打折扣。三个独立投影等于让模型自己学出三套不同的语义子空间,一套负责发问,一套负责应答,一套负责输送内容。这是自注意力机制能work的前提,不是可有可无的装饰。

还要注意投影矩阵 W_q、W_k、W_v 是所有位置共享的。也就是说,第 1 个位置和第 100 个位置用的是同一套投影参数,注意力的"动态"体现在 Q、K 的数值随输入变化,而不是参数随位置变化。这一点常被误解,有人以为注意力是给每个位置单独学一套权重,那就退化成查表了,完全没有泛化能力。

2.2 缩放因子 1/√d_k 的由来与数值实验

分母上那个 √d_k 是最容易被跳过、也最不该被跳过的地方。如果只写 QKᵀ 不缩放,在 d_k 较大时 softmax 会进入饱和区,输出接近 one-hot,梯度几乎为零,训练根本推不动。

推一下为什么。假设 Q 和 K 的每个分量都是独立同分布、均值 0 方差 1 的随机变量,那么点积 Q·K = Σᵢ qᵢkᵢ 是 d_k 个独立同分布乘积之和,其方差是 d_k,标准差是 √d_k。也就是说 d_k 越大,点积的数值范围越宽。softmax 对输入的尺度极度敏感,输入值差个 10,最大项的概率就能到 0.9999 以上,反向传播时 softmax 的雅可比矩阵几乎全是 0。

除以 √d_k 之后,点积的标准差被拉回到 1 附近,softmax 的输入落在一个梯度友好的范围内。我用一段小实验验证过:

import torch torch.manual_seed(0) for d_k in [16, 64, 256, 512]: q = torch.randn(4096, d_k) k = torch.randn(4096, d_k) raw = q @ k.T # 不缩放 scaled = raw / d_k ** 0.5 # 缩放后 p_raw = torch.softmax(raw, dim=-1) p_scaled = torch.softmax(scaled, dim=-1) print( f"d_k={d_k:4d} | raw_std={raw.std():7.3f} | scaled_std={scaled.std():5.3f} " f"| max_prob_raw={p_raw.max(dim=-1).values.mean():.6f} " f"| max_prob_scaled={p_scaled.max(dim=-1).values.mean():.6f}" )

跑出来的结果大致是这样(不同随机种子有小幅波动):

d_k未缩放点积标准差缩放后标准差未缩放最大概率均值缩放后最大概率均值
164.021.000.780.09
648.011.000.980.04
25616.031.001.00000.02
51222.651.001.00000.01

d_k=256 以上时,未缩放的最大概率已经贴着 1.0000 了,这就是典型的梯度消失现场。缩放后的分布保持在合理范围,模型才有得学。所以这个 √d_k 不是玄学常数,是有明确统计学依据的。

提示:如果你的模型是自己搭的,d_k 又设得特别大,比如 128 以上,一定要确认缩放这一步没漏。漏掉它最典型的症状是 loss 一开始就卡住不动,而且注意力图上几乎每行都是一个接近 one-hot 的尖峰。

2.3 一份可以直接跑的自注意力实现

下面这段是我平时用的最小实现,形状注释都是真实值,可以直接复制进 notebook 验证。

import math import torch import torch.nn as nn import torch.nn.functional as F class SelfAttention(nn.Module): def __init__(self, d_model, d_k=None, d_v=None, bias=True): super().__init__() d_k = d_k or d_model d_v = d_v or d_model self.d_k = d_k self.W_q = nn.Linear(d_model, d_k, bias=bias) self.W_k = nn.Linear(d_model, d_k, bias=bias) self.W_v = nn.Linear(d_model, d_v, bias=bias) self.W_o = nn.Linear(d_v, d_model, bias=bias) def forward(self, x, pad_mask=None): # x: [B, L, D] 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_v] scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k) # [B, L, L] if pad_mask is not None: # pad_mask: [B, L] 的 bool,True 表示该位置是 padding scores = scores.masked_fill(pad_mask[:, None, :], float("-inf")) attn = F.softmax(scores, dim=-1) # 整行都是 padding 时 softmax 会输出 NaN,这里做一次保护 attn = torch.nan_to_num(attn, nan=0.0) out = attn @ V # [B, L, d_v] return self.W_o(out), attn # 自测 x = torch.randn(2, 6, 32) sa = SelfAttention(32) y, a = sa(x) print(y.shape, a.shape) # torch.Size([2, 6, 32]) torch.Size([2, 6, 6]) print(a.sum(-1)) # 每行和为 1

几个容易被忽视的细节。首先是transpose(-2, -1)而不是transpose(1, 2),用负索引写更通用,序列维在前还是 batch 维在前都能兼容。其次是masked_fill用 float("-inf") 之后必须处理全掩码行,否则 softmax 的分母是 0,结果就是 NaN,而且这个 NaN 会顺着梯度污染整个 batch。torch.nan_to_num是最省事的兜底方案,也可以在 mask 时给一个极小的负数而不是负无穷。

还有一个常见困惑:为什么最后还有一个 W_o 输出投影?因为 V 投影之后各位置加权求和,得到的结果仍然在 V 的子空间里,需要再映射回原始 D 维语义空间,才能和残差连接相加。这个输出投影不是多余的,去掉它残差路径的语义就对不齐了。

2.4 复杂度、显存和几个必踩的坑

自注意力的时间和空间复杂度都是 O(L²·D),其中 L² 来自注意力矩阵。L=512 时还好,L=4096 时单是这一个矩阵就够呛。实际部署中,L 的平方开销往往是瓶颈所在,不是参数量。

踩过的坑按频率排序大概是这几个。第一个是 dtype 混用,模型 fp16 而 mask 用 fp32,导致masked_fill时隐式类型提升,速度掉一半。第二个是忘记.contiguous(),多头场景下 transpose 之后直接 view 会报错,这个下一节细说。第三个是 padding mask 和 causal mask 混着用的时候搞反了维度,结果是左边信息泄露或者整句被全屏蔽掉。

注意:fp16 训练时慎用 float("-inf") 做 mask。某些算子实现下负无穷加上后续运算会溢出成 NaN,改用torch.finfo(scores.dtype).min或者一个足够小的常数(比如 -1e4)更稳,实测下来训练曲线也更平顺。

3. 多头注意力机制:拆头、拼头与工程实现里的坑

3.1 为什么要多头:从"多角度观察"到子空间划分

单头注意力只有一套 W_q、W_k、W_v,意味着它只能用一种相似度度量去决定关注谁。但语言里的关系是多种多样的,有语法上的主谓一致,有指代关系,有语义搭配,还有单纯的位置邻近。一套投影很难同时把这些关系都编码进去。

多头注意力机制的做法是把 D 维的表示切成 H 份,每一份单独做一次自注意力,最后拼回来。每个头有自己独立的投影参数,因此可以学到不同的关注模式。有些头专门盯相邻位置,有些头专门跟踪指代,有些头看起来像在关注标点和句尾。研究里管这个叫"头 specialization",实际训练中确实能观察到部分头会退化或者变得相似,但整体上多头带来的表达能力提升是实打实的。

有个容易误解的点:多头不是把一个完整注意力算 H 遍,而是把维度切开分别算。所以当 d_model 固定时,参数量和主 FLOPs 与头数基本无关。真正的差异在于每个头的维度变小了,单个头的表达容量下降,但视角数量增加。这个取舍需要靠实验决定,没有万能公式。

3.2 两种拆头写法与等价性验证

拆头有两种常见写法,我强烈推荐第一种。

第一种是合并 QKV 投影再 chunk。用一个nn.Linear(D, 3D)一次性算出 Q、K、V,然后沿最后一维切成三块。好处是只有一次矩阵乘,GPU 利用率更高,而且权重初始化时三个投影的分布是一致的。

第二种是三个独立的nn.Linear(D, D),可读性好但多两次 kernel 启动,小 batch 下差距明显。

拆头的过程是 reshape + transpose,这里有个必须注意的顺序问题:

# q: [B, L, D],H 个头,Dh = D // H q = q.reshape(B, L, H, Dh) # [B, L, H, Dh] q = q.transpose(1, 2) # [B, H, L, Dh] # 计算完之后 out = out.transpose(1, 2).reshape(B, L, D)

为什么先 reshape 成 [B, L, H, Dh] 而不是 [B, H, L, Dh]?因为 D 维在内存里是连续的,reshape 成 [B, L, H, Dh] 保证每个头的维度 Dh 是连续的一段,拆出来才对。如果直接 reshape 成 [B, H, L, Dh],那分片的逻辑就变成了按位置切而不是按通道切,语义完全错乱。这个错误很隐蔽,模型照样能训,只是效果差一截,很难通过报错发现。

transpose 之后张量在内存里不再连续,所以往回 reshape 之前必须先contiguous(),或者直接用reshape()让它内部自动处理。我一般写.transpose(1, 2).contiguous().reshape(B, L, D),意图更清晰。

3.3 完整的 MultiHeadAttention 代码

class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1, bias=True): super().__init__() assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除" self.d_model = d_model self.n_heads = n_heads self.d_head = d_model // n_heads self.qkv = nn.Linear(d_model, 3 * d_model, bias=bias) self.out_proj = nn.Linear(d_model, d_model, bias=bias) self.dropout = nn.Dropout(dropout) self.scale = self.d_head ** -0.5 def forward(self, x, attn_mask=None, key_padding_mask=None): B, L, _ = x.shape H, Dh = self.n_heads, self.d_head qkv = self.qkv(x) # [B, L, 3D] q, k, v = qkv.chunk(3, dim=-1) # 各 [B, L, D] # 拆头 q = q.reshape(B, L, H, Dh).transpose(1, 2) # [B, H, L, Dh] k = k.reshape(B, L, H, Dh).transpose(1, 2) v = v.reshape(B, L, H, Dh).transpose(1, 2) scores = (q @ k.transpose(-2, -1)) * self.scale # [B, H, L, L] if attn_mask is not None: # 因果掩码之类,形状可广播到 [B, H, L, L] scores = scores + attn_mask if key_padding_mask is not None: # [B, L],True 表示该 key 位置无效 scores = scores.masked_fill( key_padding_mask[:, None, None, :], float("-inf") ) attn = F.softmax(scores, dim=-1) attn = torch.nan_to_num(attn, nan=0.0) attn = self.dropout(attn) out = attn @ v # [B, H, L, Dh] out = out.transpose(1, 2).contiguous().reshape(B, L, self.d_model) return self.out_proj(out), attn

参数量的算法可以随手估一下:d_model=512、H=8 时,qkv 层是 512×1536+1536 = 787968,out_proj 是 512×512+512 = 262656,合计约 105 万。如果把 H 改成 16,参数量完全不变,因为 3D 的宽度没变。但注意力矩阵从 [B, 8, L, L] 变成 [B, 16, L, L],显存直接翻倍。

3.4 头数、头维度与推理速度的取舍

选头数没有理论最优解,但有几条经验可以少走弯路。

头数 H单头维度 Dh(D=512)特点适用场景
4128单头容量大,视角少小数据量,任务单一
864经典配置,平衡通用首选
1632视角多,单头偏弱大数据量,长文本
3216单头过窄,易退化一般不推荐

经验上 Dh 低于 32 之后收益就开始递减了,因为单个头能表达的语义太窄,很多头会退化成近似相同的行为。BERT-base 用 12 头、Dh=64,GPT-3 这类大模型反而把头维度压到 128 左右但层数堆得很高,思路是用深度换宽度。

还有一个实际部署时才暴露的问题:多头注意力的计算是 [B, H, L, L] 的矩阵乘,H 增大时 kernel 的并行度提高但单次矩阵乘的规模变小,GPU 利用率反而可能下降。我实测过同一模型在 H=8 和 H=16 下的推理延迟,H=16 虽然 FLOPs 一样,但延迟高了大概 8%,主要是 kernel 启动和小矩阵乘的效率损失。

实操心得:如果你的序列长度超过 1024,优先考虑的不是调头数,而是先用torch.nn.functional.scaled_dot_product_attention,它会自动选择 FlashAttention 之类的融合实现,显存和速度都比手写版好一大截。手写版的价值在于理解原理和方便魔改,生产环境没必要硬扛。

4. 通道注意力与空间注意力:CNN 骨干网里怎么加才有效

4.1 SE 通道注意力:Squeeze-Excitation 的完整流程

通道注意力机制的代表作是 SE(Squeeze-and-Excitation),思路极其简单:每个通道的重要性不一样,那就让网络自己学一组通道权重。

Squeeze 阶段用全局平均池化把 [B, C, H, W] 压成 [B, C, 1, 1],每个通道得到一个标量,代表这个通道在整个空间上的响应强度。Excitation 阶段接两层全连接(中间有降维),过 sigmoid 得到 0~1 之间的通道权重,最后逐通道乘法回原特征图。

class SEBlock(nn.Module): def __init__(self, channels, reduction=16): super().__init__() hidden = max(channels // reduction, 4) self.pool = nn.AdaptiveAvgPool2d(1) self.fc = nn.Sequential( nn.Linear(channels, hidden, bias=False), nn.ReLU(inplace=True), nn.Linear(hidden, channels, bias=False), nn.Sigmoid(), ) def forward(self, x): B, C, _, _ = x.shape s = self.pool(x).view(B, C) # [B, C] w = self.fc(s).view(B, C, 1, 1) # [B, C, 1, 1] return x * w

这里 reduction 的比例是关键超参。原论文用 16,但在通道数很少的场景(比如 C=32)下,C//16=2,降维太狠信息损失严重,所以我在代码里加了max(..., 4)的下限保护。反过来,如果通道数上千,reduction=16 中间的隐层还有几十上百维,参数量也不小,可以适当加大比例。

SE 的代码量小,但它有个隐蔽的性能问题:两个全连接层在 GPU 上延迟不小,尤其是 batch 小的时候。有一版改进把第二个 FC 换成 1x1 卷积,或者直接用一次矩阵运算完成,效果差不多但快一些。

4.2 CBAM:通道-空间协同注意力的串行结构

CBAM 是在 SE 基础上的加强版,它把通道注意力和空间注意力串起来用,顺序是"先通道后空间"。这个顺序不是随便定的,原论文做过消融实验,串行优于并行,通道在前优于空间在前。直觉上说得通:先确定哪些通道重要,在已经筛选过的特征上再定位哪些空间位置重要,比反过来更合理。

CBAM 的通道分支和 SE 有个重要区别——它同时用了平均池化和最大池化,两条路走同一个 MLP 然后相加。最大池化能捕捉最显著的响应,平均池化反映整体分布,两者互补。

class ChannelAttention(nn.Module): def __init__(self, channels, reduction=16): super().__init__() hidden = max(channels // reduction, 4) self.avg_pool = nn.AdaptiveAvgPool2d(1) self.max_pool = nn.AdaptiveMaxPool2d(1) # 用 1x1 卷积实现,避免 Linear 在高维下的额外开销 self.mlp = nn.Sequential( nn.Conv2d(channels, hidden, 1, bias=False), nn.ReLU(inplace=True), nn.Conv2d(hidden, channels, 1, bias=False), ) self.sigmoid = nn.Sigmoid() def forward(self, x): a = self.mlp(self.avg_pool(x)) m = self.mlp(self.max_pool(x)) return x * self.sigmoid(a + m) class SpatialAttention(nn.Module): def __init__(self, kernel_size=7): super().__init__() assert kernel_size in (3, 7), "kernel size 一般取 3 或 7" self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False) 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] cat = torch.cat([avg_out, max_out], dim=1) # [B,2,H,W] return x * self.sigmoid(self.conv(cat)) class CBAM(nn.Module): def __init__(self, channels, reduction=16, kernel_size=7): super().__init__() self.ca = ChannelAttention(channels, reduction) self.sa = SpatialAttention(kernel_size) def forward(self, x): x = self.ca(x) x = self.sa(x) return x

空间分支里用 7x7 卷积而不是 3x3,是因为需要在空间上聚合足够大的邻域信息才能判断"哪块区域重要",感受野太小容易只看到局部噪声。但这个 7x7 卷积本身也带来了一些计算量,对分辨率高的浅层特征图不太友好。

4.3 CA 注意力:坐标信息嵌入的轻量化思路

SE 和 CBAM 都有一个共同缺陷:全局池化把空间信息压成一个标量,位置信息彻底丢了。对于需要精确定位的任务,比如关键点检测、细粒度分类,这个损失可能很致命。

CA(Coordinate Attention)的做法是别做全局池化,改成沿 H 方向和 W 方向分别池化,得到两组带方向的位置信息,再融合起来生成注意力权重。

class CoordAttention(nn.Module): def __init__(self, channels, reduction=32): super().__init__() hidden = max(8, channels // reduction) self.pool_h = nn.AdaptiveAvgPool2d((None, 1)) # [B,C,H,1] self.pool_w = nn.AdaptiveAvgPool2d((1, None)) # [B,C,1,W] self.conv1 = nn.Conv2d(channels, hidden, 1, bias=False) self.bn1 = nn.BatchNorm2d(hidden) self.act = nn.ReLU(inplace=True) self.conv_h = nn.Conv2d(hidden, channels, 1, bias=False) self.conv_w = nn.Conv2d(hidden, channels, 1, bias=False) def forward(self, x): B, C, H, W = x.shape x_h = self.pool_h(x) # [B,C,H,1] x_w = self.pool_w(x).permute(0, 1, 3, 2) # [B,C,W,1] y = torch.cat([x_h, x_w], dim=2) # [B,C,H+W,1] y = self.act(self.bn1(self.conv1(y))) y_h, y_w = torch.split(y, [H, W], dim=2) y_w = y_w.permute(0, 1, 3, 2) # [B,C,1,W] a_h = self.conv_h(y_h).sigmoid() # [B,C,H,1] a_w = self.conv_w(y_w).sigmoid() # [B,C,1,W] return x * a_h * a_w

注意最后是两次广播相乘,a_h 沿宽度方向广播,a_w 沿高度方向广播,合起来就得到了每个位置 (i, j) 的独立权重。这样注意力图不再是 SE 那种"一列权重",而是真正有空间分辨能力的二维权重,同时参数增加有限。

通道保留比例 reduction=32 比 SE 的 16 更激进,因为 CA 中间的特征图长度是 H+W,本身就不小,再乘通道数容易超标。代码里用max(8, ...)保底。

4.4 三种注意力的横向对比与插入位置

维度SECBAMCA
关注维度仅通道通道+空间通道+空间(带方向)
池化方式全局平均全局平均+最大沿 H/W 分别池化
位置信息丢失丢失部分保留
额外参数极少少少
计算开销最低中中低
典型用途分类骨干网检测/分割定位敏感任务

插入位置这件事比选哪个模块更影响效果。我的经验是这样几条。

第一,不要在网络的第一个卷积层后面就加。浅层特征分辨率高,通道数少,注意力模块的收益小但计算开销占比大。第二,瓶颈结构(比如 ResNet 的 bottleneck)里应该加在残差相加之前的最后一个卷积之后,这样权重作用在待相加的特征上,不会绕过残差路径。第三,detection 这类任务在 neck 部分加收益往往比 backbone 里加更明显,因为 neck 的特征图分辨率适中,语义也更接近任务目标。

注意:如果你的 backbone 是预训练权重加载的,加了注意力模块之后需要重新训练或者至少用较小的学习率微调较长时间,直接冻结 backbone 只训新模块效果通常不好。

5. 时序注意力与 Swin 窗口注意力:几个容易搞混的边界

5.1 时序注意力机制原理:掩码、位置与自回归

时序注意力机制和普通自注意力的唯一区别就是加了因果掩码,保证位置 t 只能看到 1 到 t,看不到未来。这是自回归生成的基本要求,不管你是在做语言模型还是时序预测。

掩码的构造很简单,一个上三角矩阵,未来位置填负无穷:

def causal_mask(L, device, dtype=torch.float32): # 上三角(不含对角线)为 -inf m = torch.triu(torch.ones(L, L, device=device, dtype=torch.bool), diagonal=1) mask = torch.zeros(L, L, device=device, dtype=dtype) return mask.masked_fill(m, torch.finfo(dtype).min) # 用法:scores = scores + causal_mask(L, x.device) # 可广播到 [B,H,L,L]

这里用了torch.finfo(dtype).min而不是负无穷,理由前面说过,fp16 下更稳。掩码是加在 softmax 之前的 logits 上,不是在权重上加,顺序搞反了结果就完全错了。

时序场景里还有一个常被忽略的细节:如果做的是多变量时序预测,位置编码的选择影响很大。正弦位置编码对固定长度的序列够用,但如果序列长度变化剧烈(比如传感器采样不均),用可学习的位置嵌入配上一个长度上限,或者干脆用相对位置编码,效果更稳。

另外提一下效率。自回归推理时,每一步都会重新计算整个前缀的 K 和 V,重复度极高。标准做法是用 KV Cache,只算当前 token 的 Q、K、V,把 K、V 缓存起来拼接。正确实现的 KV Cache 能把生成速度提升数倍,但缓存管理容易出 bug,最常见的是忘记在生成新序列时清空缓存,导致结果串味。

5.2 Swin 的窗口注意力与移位机制

Swin Transformer 解决的是视觉任务里序列太长的问题。224×224 的图如果按 16×16 的 patch 切,序列长度就是 196,看着还行;但如果是 1024×1024 的图,序列长度直接到 4096,全局注意力的 O(L²) 就顶不住了。

Swin 的思路是把图像划分成固定大小的窗口,窗口内做自注意力,窗口之间不交互。这样复杂度从 O((HW)²) 降到 O(HW·M²),M 是窗口边长,一般取 7。但纯窗口注意力会让不同窗口之间完全隔离,感受野受限,所以 Swin 又加了 shifted window:下一层把窗口划分整体平移半个窗口大小,让原本不在同一个窗口里的 patch 有机会碰面。

窗口划分的核心代码就几行:

def window_partition(x, window_size): # x: [B, H, W, C] B, H, W, C = x.shape x = x.view(B, H // window_size, window_size, W // window_size, window_size, C) windows = x.permute(0, 1, 3, 2, 4, 5).contiguous() windows = windows.view(-1, window_size, window_size, C) return windows # [B * nW, M, M, C] def window_reverse(windows, window_size, H, W): B = int(windows.shape[0] / (H * W / window_size / window_size)) x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) return x

shift 的实现用 torch.roll,分别沿 H 和 W 方向滚动 -M//2,算完注意力再滚回来。要注意的是 H 或 W 不是窗口大小整数倍时需要 pad,算完再去掉 pad。这个 pad 逻辑是 Swin 实现里最容易写错的部分,形状对不上时优先检查这里。

移位窗口带来的一个副作用是,滚回来的窗口里包含了来自不同区域的 patch,注意力计算时跨区域的部分是无效的,需要额外加一个 attention mask 屏蔽掉。这个 mask 的构造在早期的开源实现里出过不少 bug,如果你发现自己的 Swin 训出来比预期差,可以先用简单的窗口注意力(不做 shift)跑一遍作为 baseline,差距太大就说明 mask 有问题。

5.3 PyTorch 与 TensorFlow 落地时的差异

两个框架的多头注意力接口差异比想象中大。

PyTorch 从 2.0 开始提供了torch.nn.functional.scaled_dot_product_attention,参数顺序是 (query, key, value, attn_mask, dropout_p, is_causal),会自动选后端实现。用is_causal=True直接开因果掩码,比手写 mask 又快又省事。但要注意它默认不开 dropout 的 RNG 对齐,训练时如果你想精确复现,得传dropout_p并把模型设成 train 模式。

另一条路是用nn.MultiheadAttention,它的输入默认是 [L, B, D],和你手写的 [B, L, D] 正好反过来,batch_first=True参数可以切换。这个坑我踩过不止一次,形状没报错但结果完全不对,因为 [L, B, D] 和 [B, L, D] 在 L 和 B 数值接近时不会触发维度检查。

TensorFlow/Keras 的tf.keras.layers.MultiHeadAttention接口更"高层",它内部处理了拆头、缩放、拼头,还支持use_causal_mask参数。好处是不容易写错,坏处是魔改困难,比如你想换成自己设计的稀疏注意力模式,就得从tf.keras.layers.Layer重写。

跨框架迁移时最需要注意的是权重排布。PyTorch 的 Linear 权重形状是 [out_features, in_features],TensorFlow 的 Dense 是 [in_features, out_features],转的时候要转置。多头 QKV 合并层的排布在两个框架里也不同,最好写一个小脚本,用同一组随机输入对比两边输出,逐层对齐,比凭记忆猜靠谱得多。

6. 常见问题与排查技巧实录

6.1 训练不收敛、梯度异常、注意力图全灰

这类问题我按症状整理一下排查路径。

症状一:loss 从第一步就卡住不动,梯度范数极小。八成是缩放因子漏了,或者 softmax 的输入被掩码打成了全负无穷。打印一下 scores 的标准差,如果超过 10 就基本确认了。

症状二:loss 是 NaN。按顺序查三处:一是掩码后的全掩码行,二是 fp16 下的负无穷溢出,三是学习率过大。最快的定位方式是在 forward 里加torch.autograd.set_detect_anomaly(True),虽然慢但能直接告诉你哪个算子出的 NaN。

症状三:注意力图几乎全灰,也就是每行权重接近均匀分布。这说明模型没学到任何有区分度的关注模式,可能原因是缩放过度(分母写成了 d_k 而不是 √d_k),或者温度参数设得太高。偶尔也见到把 softmax 的 dim 写错,沿着 L 维以外的维度归一化,那就完全乱套了。

症状四:训练集 loss 正常下降,验证集差距很大。注意力模块对小数据集特别容易过拟合,尤其是加了 CBAM 这类参数模块之后。这时候先减 reduction 比例,或者直接在浅层不加注意力。

6.2 显存与速度优化的实操清单

这一块我做了一轮系统的对比,结论整理成清单更直观。

  • 优先用scaled_dot_product_attention替换手写实现,L≥512 时显存通常能降 30% 到 50%,速度提升 1.5 到 3 倍。
  • 反向传播不需要的中间注意力图用torch.no_grad()包起来,尤其是做可视化的部分,很多人忘了它也会占显存。
  • 头数不要为了"看起来更强"往上堆,注意力矩阵的显存和 H 成正比。
  • 推理阶段一定上 KV Cache,自回归生成能省掉大量重复计算。
  • 如果只做推理,把模型转成 fp16 甚至 int8,注意力部分的收益比全连接层更明显,因为它是访存密集型算子。
  • 长序列场景考虑分块计算注意力,比如按 1024 一块,块间用近似交互,牺牲一点精度换大幅显存下降。

6.3 问题速查表

现象最可能原因快速验证方法处理方式
loss 立刻卡住漏了 √d_k 缩放打印 scores.std()加缩放或调温度
loss 为 NaN全掩码行 softmax检查 mask 是否有整行为 Truenan_to_num 或改掩码值
训练正常推理不对忘了 eval() 或没清 KV Cache关掉 dropout 再测模型切 eval,重置缓存
显存突然翻倍头数增加或 L 变长打印注意力矩阵形状换融合实现或分块
加了注意力反而掉点插入位置不当或过拟合对比不加注意力的 baseline换位置、减 reduction
多头结果和单头一样reshape 顺序写错检查拆头后 Dh 的连续性用 [B,L,H,Dh] 再 transpose
不同框架结果对不上权重排布或输入布局差异同输入逐层对齐输出转置权重、统一 batch_first

实操心得:手写注意力最大的价值是调试友好。我在排查问题时习惯保留一个"慢速参考实现",用 fp32、循环写法,跑一小批数据。当融合算子或者 fp16 版本表现异常时,拿参考实现的输出对比,通常五分钟就能定位是数值精度问题还是逻辑 bug。

7. 我个人在项目里踩出来的几条经验

注意力机制的代码看着短,但每一个维度、每一个掩码、每一个缩放系数都对应着具体的数学含义,改动任何一处都要想清楚它在数值上会发生什么。我最初写多头注意力的时候,拆头顺序写反了,模型照样能训,loss 曲线也降,只是最终精度比参考实现低了两个点,这种"沉默的 bug"最费时间。

另外一条经验是,别一上来就追求最花哨的变体。SE 和 CBAM 这种几行代码的模块,在大多数视觉任务上已经能带来稳定收益,CA 和 Swin 这类复杂结构需要配套的训练策略和足够的数据量才能发挥出来。先在简单模块上把插入位置、reduction 比例这些超参调明白,再上复杂结构,性价比高得多。

最后分享一个我常用的自检手法:写完全力注意力相关的新模块后,先跑三个测试——形状测试(输入输出维度一致)、权重测试(注意力权重每行和为 1)、梯度测试(反向传播后参数有梯度且不是 NaN)。这三条过了,再去跑真实数据。比直接上训练快了不知道多少倍,也省了不少显卡时间。

返回列表