ARTICLE DETAIL

资讯详情

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

注意力机制全面解析:从SE到多头注意力及Pytorch实现

注意力机制全面解析:从SE到多头注意力及Pytorch实现 1. 从为什么说起卷积网络到底漏掉了什么先说个可能颠覆新手认知的事实传统的卷积神经网络本质上是一个局部关系提取器而不是一个全局关系提取器。每一层卷积核的感受野再大也只是在窗口内做加权求和信息从底层到高层层层传递靠堆叠层数来间接获得全局视野。这种方式有一个天然短板——它无法主动告诉网络哪部分特征更重要。举一个最直观的例子你要用网络识别一张照片里的一只鸟。卷积层会无差别地把羽毛、树枝、背景云彩全部提取成特征图然后送到后面的分类器。可实际上真正能区分这是鸟的可能只有鸟类特有的喙、翅膀纹理以及它站在树枝上的姿态。背景和树枝对分类结果的贡献是干扰项。如果网络能用某种机制知道自己该把注意力放在喙和翅膀上忽略云彩和树枝识别准确率是不是能提升一个档次这正是注意力机制Attention Mechanism要解决的事情。它模仿了人脑的视觉选择机制在扫视一幅图像时你不会均匀分配视觉资源而是先把目光聚焦到最关键的局部区域再快速掠过其余部分。把这个逻辑塞进神经网络就成了图像处理领域过去几年最重要的技术演进之一。SE模块、CBAM、CA注意力、自注意力、多头注意力这些名词本质上都是让网络学会分配注意力权重这件事的不同实现方案。这一篇我按自己的学习路径来写配合完整可跑的Pytorch代码把图像处理中主流的注意力机制全部梳理一遍。你不需要有很深的数学基础懂一点Pytorch的tensor操作就能看懂关键是跟着代码把维度变化捋清楚——注意力机制代码的绝大部分坑都藏在tensor的维度变换里。2. 通道维度的注意力SE模块是真正的起点2.1 SE模块的设计逻辑和完整实现SENetSqueeze-and-Excitation Networks是2018年ImageNet分类冠军的基石也是我建议每个初学者第一个掌握的注意力模块。它的思想非常朴素既然卷积输出的是一个多通道的特征图形状为[B, C, H, W]B是批次大小C是通道数H和W是特征图高宽那每个通道就代表一种特征响应模式。SE做的就是先全局压缩空间信息再学习每个通道的重要性权重最后把权重乘回原特征图。它的执行流程只有三步Squeeze压缩对特征图做全局平均池化把每个通道的H x W空间信息压缩成一个数值得到[B, C, 1, 1]的向量。这一步相当于统计这个通道在整张图上平均激活了多少。Excitation激励将[B, C, 1, 1]送入两个全连接层中间先降维再升维经过激活函数得到每个通道的权重系数范围在0到1之间表示该通道的重要性。Reweight重标定把通道权重系数逐通道乘回原始特征图。Pytorch代码是这样的import torch import torch.nn as nn class SEBlock(nn.Module): def __init__(self, channels, reduction16): super(SEBlock, self).__init__() # reduction是降维倍数一般是16可以调节 self.squeeze nn.AdaptiveAvgPool2d(1) self.excitation nn.Sequential( nn.Linear(channels, channels // reduction, biasFalse), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels, biasFalse), nn.Sigmoid() ) def forward(self, x): b, c, h, w x.size() # 1. Squeeze: [B, C, H, W] - [B, C, 1, 1] y self.squeeze(x) # 2. 展平: [B, C, 1, 1] - [B, C] y y.view(b, c) # 3. Excitation: [B, C] - [B, C] y self.excitation(y) # 4. 还原为 [B, C, 1, 1] 以便广播乘法 y y.view(b, c, 1, 1) # 5. Reweight: 逐通道乘回原特征图 return x * y.expand_as(x)2.2 代码背后的四个细节第一个容易踩的坑是自适应池化。nn.AdaptiveAvgPool2d(1)不管输入的H x W是多少都能输出1 x 1这让SE模块可以插在任意尺寸的特征图之后。但它直接丢掉了空间位置信息——这也成为后续许多注意力模块改进的切入点。第二个坑是降维倍数reduction的取值。原文用的是16意思是先把C维降到C/16学习完再升回C维。这个设计的目的是为了减少参数量和计算量同时用一个瓶颈结构逼着网络把通道信息压缩到一个更紧凑的表示里。不过降维倍数太大比如64会丢失信息太小比如4又起不到压缩作用实测下来16到24之间是安全区间。第三个值得注意的细节是激活函数的选择。第一层全连接后面用ReLU是为了引入非线性最后一层用Sigmoid而不是ReLU是为了把权重约束在0到1之间。注意Sigmoid输出不是和为1的分布所以SE的权重不叫通道分布而叫通道重要性系数——每个通道被独立地增强或抑制。第四个细节是为什么用全连接层而不是卷积层。在[B, C, 1, 1]的尺度上用一个1x1卷积其实是等价的但全连接在Pytorch里写起来更直接。如果想把SE改成更卷积友好的写法把两个nn.Linear换成nn.Conv2d(channels, channels // reduction, 1)加nn.Conv2d(channels // reduction, channels, 1)也完全没问题实际效果没有本质区别。我强烈建议你亲手打印一下每个步骤的tensor形状。很多人在这一步不看shape直接用y self.excitation(self.squeeze(x))一把梭结果在维度对不上时报错时根本无从下手。矩阵维度的变化是注意力机制代码里唯一真正需要警惕的地方。3. 从通道走向空间CBAM和CA的两种不同解法3.1 CBAM——通道注意力加空间注意力的串联结构SE模块只做了通道维度的注意力对空间一无所知。但很多任务里空间位置同样重要你想检测的目标在图像的左上角那右下角的信息基本是噪音。CBAMConvolutional Block Attention Module就是来补这个短板的它的核心差异是在通道注意力之后又接了一个空间注意力分支。CBAM先把特征图分别做全局平均池化和全局最大池化得到两组通道描述符送入共享的MLP就是SE里的那个两全连接结构输出两个通道权重向量相加后经过Sigmoid得到最终的通道注意力权重。这个思路比SE多用了最大池化这条分支原因是仅仅平均池化会抹平极端激活而最大池化能捕捉到最显著响应的通道一平均一最大互补能更完整地描述通道的统计特性。代码实现如下class CBAMChannelAttention(nn.Module): def __init__(self, channels, reduction16): super(CBAMChannelAttention, self).__init__() self.mlp nn.Sequential( nn.Linear(channels, channels // reduction, biasFalse), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels, biasFalse) ) self.sigmoid nn.Sigmoid() def forward(self, x): b, c, h, w x.size() # 平均池化分支 avg_out torch.mean(x, dim[2, 3], keepdimTrue) # [B, C, 1, 1] # 最大池化分支 max_out torch.max(x, dim2, keepdimTrue)[0] # [B, C, 1, W] max_out torch.max(max_out, dim3, keepdimTrue)[0] # [B, C, 1, 1] avg_out avg_out.view(b, c) max_out max_out.view(b, c) # 两条分支共享MLP avg_out self.mlp(avg_out) max_out self.mlp(max_out) # 相加后过Sigmoid scale self.sigmoid(avg_out max_out).view(b, c, 1, 1) return x * scale.expand_as(x)然后是空间注意力class CBAMSpatialAttention(nn.Module): def __init__(self, kernel_size7): super(CBAMSpatialAttention, self).__init__() self.conv nn.Conv2d(2, 1, kernel_sizekernel_size, paddingkernel_size // 2, biasFalse) 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] # 拼接成两个通道 cat torch.cat([avg_out, max_out], dim1) # [B, 2, H, W] # 用一个7x7卷积把2通道压缩成1通道 out self.conv(cat) return self.sigmoid(out)把两个模块串起来就是完整的CBAM先算通道注意力权重乘上原特征图再把加权后的结果送入空间注意力分支乘上空间权重图。这个串联顺序是论文验证过的——先通道后空间效果最好反过来会略差。CBAM里最需要在意的细节是卷积核大小。原文用的是7x7因为空间注意力本质上要建立空间像素间的局部联系卷积核太大会增加大量参数太小比如1x1又只能做逐像素缩放失去空间上下文建模的能力。你在小特征图上做CBAM时建议把7x7改成3x3或5x5否则padding后的感受野覆盖不全。3.2 CA注意力——把坐标信息当作通道之外的第三维度CBAM虽然引入了空间注意力但它的空间分支是用卷积在局部窗口内建模的对远处的空间依赖依然无能为力。CACoordinate Attention坐标注意力换了个思路我不在空间上做卷积而是把特征图按高度方向和宽度方向分别池化再像处理通道那样给它们也学习一组权重。这样高度方向和宽度方向的坐标信息就被显式编码进了注意力权重里。CA的前向过程细致拆解如下class CoordAttention(nn.Module): def __init__(self, channels, reduction32): super(CoordAttention, self).__init__() # 两个方向的池化平均池化分别在H方向和W方向 self.pool_h nn.AdaptiveAvgPool2d((None, 1)) self.pool_w nn.AdaptiveAvgPool2d((1, None)) # 先降维再升维的卷积 self.conv1 nn.Conv2d(channels, channels // reduction, kernel_size1, biasFalse) self.bn1 nn.BatchNorm2d(channels // reduction) self.act nn.ReLU(inplaceTrue) # 分离成两个方向的权重 self.conv_h nn.Conv2d(channels // reduction, channels, kernel_size1, biasFalse) self.conv_w nn.Conv2d(channels // reduction, channels, kernel_size1, biasFalse) self.sigmoid nn.Sigmoid() def forward(self, x): b, c, h, w x.size() # 1. 高度方向池化每个通道的每一行压缩成一个数 - [B, C, H, 1] x_h self.pool_h(x) # 2. 宽度方向池化每个通道的每一列压缩成一个数 - [B, C, 1, W] x_w self.pool_w(x).permute(0, 1, 3, 2) # 转置成 [B, C, W, 1] # 3. 拼接 [B, C, H, 1] 和 [B, C, W, 1] - [B, C, HW, 1] cat torch.cat([x_h, x_w], dim2) # 4. 1x1卷积降维BNReLU cat self.conv1(cat) cat self.bn1(cat) cat self.act(cat) # 5. 重新分离回两个方向 h_len h cat_h, cat_w torch.split(cat, [h_len, w], dim2) # cat_w 现在是 [B, C, W, 1]转置回 [B, C, 1, W] cat_w cat_w.permute(0, 1, 3, 2) # 6. 分别升维 out_h self.sigmoid(self.conv_h(cat_h)) # [B, C, H, 1] out_w self.sigmoid(self.conv_w(cat_w)) # [B, C, 1, W] # 7. 用外积的方式在原始特征图上广播权重 out x * out_h * out_w return out这里解决的最关键问题是attention不能只关注某个点还要关注整行和整列。CSDN上很多搬运代码都在第2步或第5步翻车因为宽方向池化后[B, C, 1, W]和高度方向[B, C, H, 1]拼接时宽高方向混在一起不统一不转置根本拼接不了。permute和torch.split的维度你要对着注释多捋两遍。CA的定位非常准它比SE多捕捉了位置信息又比CBAM的空间卷积感受野更大计算量还远低于自注意力。在移动端网络或实时语义分割任务中CA几乎是性价比最高的轻量注意力方案。4. 自注意力与多头注意力图像中的全局建模利器4.1 自注意力的本质让像素之间互相投票SE、CBAM、CA都属于轻量级注意力它们在通道或局部空间上做重标定但始终没有回答一个问题怎么让特征图里相距很远的两个像素建立直接联系卷积需要堆很多层才能间接扩大感受野而自注意力Self-Attention可以让任意两个位置一步直达。自注意力的第一步是要把特征图从图像格式展平成序列格式[B, C, H, W]变成[B, N, C]其中N H x W也就是把每个像素或者说每个位置的特征向量当作序列里的一个词。然后每个位置生成三个向量Query查询、Key键、Value值。Query负责问我该关注谁Key负责回答我有什么特征值得被关注Value是真正被加权求和的内容。注意力分数的计算公式是Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * VQ * K^T得到的是一个N x N的矩阵表示每个位置对其他所有位置的关注度除以sqrt(d_k)是为了防止内积过大导致softmax梯度消失。然后过softmax归一化成权重再对Value加权求和。Pytorch实现放在图像分类的骨干网络里class SelfAttention(nn.Module): def __init__(self, in_channels, key_channelsNone): super(SelfAttention, self).__init__() key_channels key_channels or in_channels // 8 self.q_conv nn.Conv2d(in_channels, key_channels, kernel_size1) self.k_conv nn.Conv2d(in_channels, key_channels, kernel_size1) self.v_conv nn.Conv2d(in_channels, in_channels, kernel_size1) self.softmax nn.Softmax(dim-1) self.gamma nn.Parameter(torch.zeros(1)) def forward(self, x): b, c, h, w x.size() # 投影并展平 q self.q_conv(x).view(b, -1, h * w).permute(0, 2, 1) # [B, N, d_k] k self.k_conv(x).view(b, -1, h * w) # [B, d_k, N] v self.v_conv(x).view(b, -1, h * w) # [B, C, N] # 注意力矩阵 attn torch.bmm(q, k) / (q.size(-1) ** 0.5) # [B, N, N] attn self.softmax(attn) # 加权求和 out torch.bmm(v, attn.permute(0, 2, 1)) # [B, C, N] out out.view(b, c, h, w) # 残差连接 可学习的缩放系数 out self.gamma * out x return outself.gamma是一个初始化为0的可学习标量。这招在Self-Attention的多个变体里都能看到目的是让网络先以原特征为主随着训练逐渐增大注意力分支的贡献避免一开始就在随机初始化状态下大幅度扰动特征导致训练不稳定。4.2 多头注意力一个头看全局多个头看不同侧面单头自注意力有一个问题Q * K^T的权重经过softmax归一化后往往被少数几个强相关位置主导其他位置拿到的权重趋近于0。这会让模型陷入只盯着某几个点的困境。多头注意力Multi-Head Attention的解法是把通道维度切成几段每段单独做一次自注意力最后拼回去。这样一来每个头可以学习到不同的关系模式——有的头关注颜色有的头关注边缘有的头关注上下文。代码实现class MultiHeadAttention(nn.Module): def __init__(self, channels, num_heads8, key_channelsNone): super(MultiHeadAttention, self).__init__() self.num_heads num_heads self.key_channels key_channels or channels // num_heads assert channels % num_heads 0 self.q_conv nn.Conv2d(channels, channels, kernel_size1) self.k_conv nn.Conv2d(channels, channels, kernel_size1) self.v_conv nn.Conv2d(channels, channels, kernel_size1) self.out_conv nn.Conv2d(channels, channels, kernel_size1) self.softmax nn.Softmax(dim-1) self.gamma nn.Parameter(torch.zeros(1)) def forward(self, x): b, c, h, w x.size() n h * w # 投影后按头数切分: [B, C, N] - [B, heads, C//heads, N] q self.q_conv(x).view(b, self.num_heads, self.key_channels, n) k self.k_conv(x).view(b, self.num_heads, self.key_channels, n) v self.v_conv(x).view(b, self.num_heads, self.key_channels, n) # [B, heads, N, d_k] 和 [B, heads, d_k, N] q q.permute(0, 1, 3, 2) k k.permute(0, 1, 2, 3) v v.permute(0, 1, 3, 2) attn torch.matmul(q, k) / (self.key_channels ** 0.5) # [B, heads, N, N] attn self.softmax(attn) out torch.matmul(attn, v) # [B, heads, N, d_k] out out.permute(0, 1, 3, 2).reshape(b, c, n) out out.view(b, c, h, w) out self.out_conv(out) return self.gamma * out x多头注意力复杂度高这是它在图像任务里最大的短板。[B, heads, N, N]这个注意力矩阵是平方级的特征图是64x64时N 4096单张图的注意力矩阵就有4096 x 4096 1677万个元素显存直接告急。所以实际使用多头注意力时几乎都要配一个空间降采样操作比如先池化到16x16再做注意力用完再上采样回去这就是后面要讲的嵌套式用法。5. 把注意力模块插进网络完整可跑的Pytorch实战5.1 在分类骨干网络里加SE模块一个注意力模块单独放着没意义关键是知道怎么嵌进现有网络。我拿ResNet最基本的BasicBlock举例。在残差分支相加之前对主分支输出的特征图做一次SE重标定再加回恒等映射class SEBasicBlock(nn.Module): def __init__(self, in_channels, out_channels, stride1): super(SEBasicBlock, self).__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.se SEBlock(out_channels, reduction16) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity self.shortcut(x) out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out self.se(out) out identity out self.relu(out) return out这里你可以尝试一个实验把SE模块挪到self.conv1后面或者挪到残差相加之后再乘一次SE权重对比一下验证集准确率的变化。我的实测结果是放在最后一个卷积之后、残差相加之前效果最好。原因在于残差相加是为了保留原始信息如果先相加再乘SE权重会同时压制了恒等分支里不该被压制的原始特征破坏残差结构的本意。5.2 在UNet里加注意力——医学图像分割的经典方案图像分割任务里UNet是最常用的骨架编码器不断下采样提取语义特征解码器不断上采样恢复空间细节。但下采样会丢失细节上采样又会引入噪声所以一个非常经典的做法是在跳跃连接处加注意力门控Attention Gate强迫解码器只从对应的编码器特征里提取与目标区域相关的特征图。这里给一个简化版实现思路把编码器的特征ggate信号和解码器当前的特征x分别过1x1卷积相加后经过ReLU和另一个1x1卷积再过Sigmoid得到空间权重图最后把权重乘回编码器特征。class AttentionGate(nn.Module): def __init__(self, in_channels_g, in_channels_x, out_channels): super(AttentionGate, self).__init__() self.w_g nn.Conv2d(in_channels_g, out_channels, kernel_size1, biasFalse) self.w_x nn.Conv2d(in_channels_x, out_channels, kernel_size1, biasFalse) self.psi nn.Conv2d(out_channels, 1, kernel_size1, biasFalse) self.relu nn.ReLU(inplaceTrue) self.sigmoid nn.Sigmoid() def forward(self, g, x): # g是来自解码器的上采样特征x是来自编码器的跳跃连接特征 g1 self.w_g(g) x1 self.w_x(x) if g1.size()[2:] ! x1.size()[2:]: # 尺寸不一致时需要插值对齐 x1 nn.functional.interpolate(x1, sizeg1.size()[2:], modebilinear, align_cornersFalse) psi self.relu(g1 x1) psi self.sigmoid(self.psi(psi)) return x * psi这种注意力门控和通道注意力最大的不同是它不是对通道做缩放而是得出一张与特征图同分辨率的空间权重热力图属于像素级注意力。你要可视化它的输出会看到网络自动把响应集中在目标器官或目标物体区域背景区域权重趋近于0。5.3 把多头自注意力包装成即插即用模块如果你不想动整个网络结构只想在某层后面加一个全局建模能力可以把多头注意力包成一个独立模块插入到骨干网络的指定位置。这里我推荐使用带空间降采样的版本防止显存爆炸class DownsampleSelfAttention(nn.Module): def __init__(self, channels, num_heads8, downsample_factor8): super(DownsampleSelfAttention, self).__init__() self.down nn.AdaptiveAvgPool2d((None, None)) # 实际用插值 self.downsample_factor downsample_factor self.attn MultiHeadAttention(channels, num_headsnum_heads) def forward(self, x): b, c, h, w x.size() # 下采样到较小分辨率做全局注意力 small_h, small_w h // self.downsample_factor, w // self.downsample_factor x_small nn.functional.interpolate(x, size(small_h, small_w), modebilinear, align_cornersFalse) x_small self.attn(x_small) # 上采样回原分辨率 x_large nn.functional.interpolate(x_small, size(h, w), modebilinear, align_cornersFalse) # 残差连接 return x x_large实际用的时候多头的num_heads建议设成8key_channels每个头的维度一般为channels // num_heads比如512通道就是64维。如果通道数不能被头数整除先做一次1x1卷积把通道数调整到可整除的数值再送入多头注意力。我见过不少人在512通道上强行设置9个头结果直接报错或者通道分配不均导致效果崩盘。6. 注意力模块选型对照不同场景到底该选哪种按经验帮大家整理了一张选型表这六大类我都在实际项目里用代码跑过对比算是有一定可信度的参考模块核心操作复杂度适合场景注意点SE通道全局池化全连接重标定极低分类网络骨干、轻量模型丢失空间位置信息CBAM通道注意力空间卷积注意力低检测/分割特征图不大时空间卷积核大小要按特征图调CA双向坐标池化权重外积低语义分割、移动端、检测维度变换繁琐注意permute顺序Non-local/自注意力全局QKV点积高平方级需要长距离依赖的中高分辨率特征显存爆炸建议配下采样多头自注意力多头并行QKV很高Transformer类结构、大模型需预训练/大batch小数据集易过拟合Attention Gate双分支注意力门控低UNet类分割网络的跳跃连接需要处理g和x的尺寸对齐几个判断原则你的任务是简单的图像分类数据量不大SE足够它在计算量上几乎可以忽略不计还能稳定涨点1%到2%。你的任务是目标检测或分割空间位置重要CBAM和CA优先。两者的区别在于CBAM更简单、直观CA对远处空间关系的建模更强。你的任务是图像生成、高分辨率图轻量注意力已经不太够用了应该考虑把多头注意力用在低分辨率特征层上高分辨率层继续用卷积。你的数据集只有几千张图不建议直接用自注意力那一套注意力对大数据量的依赖很重小数据上基本会被卷积网络压着打。7. 复现与调试经验我踩过的和你们即将踩的坑7.1 shape不匹配是注意力代码最大的坑我把所有常见报错路线都走了一遍总结出三大高频错误场景场景一SE里全连接层输入输出维度写反。nn.Linear(channels // reduction, channels, biasFalse)很容易在注意力计算时写反成nn.Linear(channels, channels // reduction)然后第二个全连接又跟着错。核对维度时直接把print(y.shape)打在每一行后面一目了然。场景二CBAM的torch.max用了dim1而不是dim[2,3]。在空间注意力里你是在通道维度上做最大池化取每个像素最强的通道响应所以是dim1结果[B, 1, H, W]而在通道注意力里你是把空间维度分别池化得到每个通道的全局响应所以要分开对H和W两次torch.max。同一个torch.max用法在不同模块里维度参数完全不同。场景三多头注意力的permute和reshape组合导致输出特征错乱。尤其注意view是在内存连续的前提下按顺序切分而permute之后tensor的内存布局已经变化必须先用contiguous()再view否则轻则结果错乱重则直接报错说viewrequires tensor to be contiguous。这个是Pytorch新手的经典坑。7.2 注意力模块的位置比模块本身更重要我做过一组消融实验同一个SE模块放在ResNet的conv1前、每个BasicBlock里、以及只在最后一层全局池化前最终精度差别最大能到3%以上。结论是放在每个残差块里收益最稳定这也是论文默认做法。只在最后一层放相当于只做了一次全局重标定前面的特征没有享受到注意力的梯度回传。放在第一层卷积之前几乎等于没用输入图像本身没有通道语义。对CBAM和CA这类带空间建模能力的注意力可以适当减少放置密度比如每两个残差块放一次计算量能省不少精度基本不掉。7.3 注意力模块不是越深越好这是新手最容易误解的一点——以为网络每个block后面都挂一个注意力模块效果一定最好。实际上当注意力模块过多时特征会被反复重标定原始信息被过度挤压导致训练困难甚至精度下降。我的建议是先在倒数第二或第三层加一个注意力模块跑通整个训练流程观察收敛情况。确认有效后再逐步往浅层加每次加一层对比验证集指标。每加一层就检查一次显存占用和训练速度心里有数。7.4 学习率调度和初始化注意力模块里的Sigmoid输出初始值大约在0.5附近这意味着网络一开始激活值会缩水一半如果主分支没有残差连接等于给整个网络加了一个固定衰减训练初期会明显变慢。解决办法就是前面提到过的gamma参数self.gamma nn.Parameter(torch.zeros(1))初始化为0意味着注意力分支的输出是0out gamma * out x就等于恒等输出网络起步时不受到任何扰动。随着训练更新gamma会慢慢变大注意力分支自动逐渐打开。这个技巧在自注意力系列里几乎是标配但很多人抄代码时把它漏了后面训练效果差还找不到原因。7.5 测试模式下的行为差异如果你在训练时开了BatchNorm注意在model.eval()下注意力模块的BN层也会切换到running mean和running var。用预训练模型做推理时一定要确认BN的running统计量已经随训练更新完毕否则注意力权重在推理和训练时的分布会不一致表现出来就是训练准确率正常、验证集准确率忽高忽低。我遇到过最诡异的一次是迁移学习时只冻结了骨干网络忘了冻结注意力模块里的BN层结果BN的running mean被少量新数据疯狂带偏验证集特征图热力图肉眼可见地全乱套了。后来统一把注意力模块里的BN层也冻住问题立刻消失。8. 一张图看懂全局注意力机制在图像任务中的完整体系用文字画一张全局路线图。整个注意力体系就三条主线通道主线SE - CEChannel Attention- 各种通道注意力变体核心是哪些通道重要。空间主线Spatial Attention - Self-Attention - Non-local核心是哪些位置重要。混合主线CBAM通道局部空间- CA通道坐标空间- 自注意力全局空间通道混合建模核心是怎么把两者结合。所有的注意力模块基本都逃不出这个框架。你看任何一个新出的注意力论文第一件事先问自己三个问题它是在哪个维度做的池化或聚合它生成的权重形状是[B, C, 1, 1]、[B, 1, H, W]还是[B, N, N]它最后是怎么乘回原特征图的这三个问题一答模块的本质就清晰了再花哨的结构也无非是这几个维度的组合排列。按照这个思路去读源码你不需要阅读成千上万行代码只需要拿一小块特征图打桩打印跑一遍前向传播观察每个tensor的形状变化比盯着论文里的公式推导快得多。我每次接触一个新注意力变体固定动作就是写一个最小复现脚本随机初始化输入[2, 64, 32, 32]然后逐行打印中间变量shape跑通了再往真实网络里塞。最后给长期主义者的建议注意力机制不是银弹它解决的核心问题是特征被无差别对待但前提是你的骨干网络已经把特征提取得足够好。先把卷积基础打扎实再研究注意力怎么加反过来注意力模块再花哨也救不了一个设计混乱的骨干网络。我见过有人在ResNet18都跑不稳定的情况下直接上Transformer最后模型不收敛还怪数据集这属实是走偏了方向。按部就班把每种注意力的代码和维度变换吃透图像处理的注意力这块就算真正入门了。
返回列表