当前位置: 首页 > news >正文

【即插即用完整代码】CVPR 2026新方法归一化空间与通道注意力,无额外参数,轻量且高效,超越CBAM,快速涨点,发表论文!

专栏内提供试读,感兴趣的小伙伴可以订阅一下哈!

适用于所有的CV二维任务:图像分割、超分辨率、目标检测、图像识别、低光增强、遥感检测等

每日分享最新的前沿技术,

助力快速发论文、模型涨点!

摘要

本文提出了一种新型的基于归一化的注意力模块(NAM),通过抑制不重要的权重来提高模型的计算效率,同时保持性能。与SE、BAM、CBAM等其他注意力机制相比,NAM在ResNet和MobileNet上均实现了更高的准确率。代码已公开。

引言

注意力机制是近年来的研究热点,它帮助深度神经网络抑制不重要的像素或通道,从而提高模型性能。然而,现有方法大多关注于捕获显著特征,而忽视了权重对特征的贡献。本文提出了一种基于归一化的注意力模块(NAM),利用批量归一化(Batch Normalization)的缩放因子来衡量权重的重要性,从而进一步抑制不重要的通道或像素。

创新点

  1. 基于归一化的注意力机制:NAM利用批量归一化的缩放因子来衡量权重的重要性,避免了添加全连接层或卷积层,从而减少了参数数量。

  2. 轻量且高效:NAM在不增加额外参数的情况下,通过抑制不重要的权重来提高计算效率。

  3. 超越现有方法:实验表明,NAM在ResNet和MobileNet上均优于SE、BAM、CBAM等现有注意力机制。

主要方法

归一化注意力模块(NAM)

NAM的核心在于利用批量归一化的缩放因子来衡量权重的重要性。具体来说:

  • 通道注意力子模块:通过批量归一化的缩放因子来衡量每个通道的重要性,并生成通道注意力权重。

  • 空间注意力子模块:类似地,通过像素归一化来衡量每个像素的重要性,并生成空间注意力权重。

  • 正则化项:为了进一步抑制不重要的权重,NAM在损失函数中添加了一个L1范数惩罚项,以平衡通道和空间注意力权重的稀疏性。

架构集成

NAM模块被嵌入到每个网络块的末尾,对于残差网络,它被嵌入到残差结构的末尾。这种设计使得NAM可以与现有的网络架构无缝集成,而不需要对网络结构进行重大修改。

实验细节

数据集和模型
  • CIFAR-100:使用ResNet50进行实验,与SE、BAM、CBAM等方法进行比较。

  • ImageNet:使用MobileNet V2进行实验,与SE、BAM、CBAM等方法进行比较。

训练配置
  • CIFAR-100:使用4个Nvidia Tesla V100 GPU进行训练,训练配置与CBAM相同,正则化参数p设置为0.0001。

  • ImageNet:使用4个Nvidia Tesla V100 GPU进行训练,训练配置与CBAM相同,正则化参数p设置为0.001。

实验结果分析

CIFAR-100
  • ResNet50 + NAM(通道注意力):参数数量为23.74M,FLOPs为1.31G,Top-1错误率为19.09%,Top-5错误率为4.5%。

  • ResNet50 + NAM(空间注意力):参数数量为23.71M,FLOPs为1.31G,Top-1错误率为19.38%,Top-5错误率为4.72%。

  • ResNet50 + CBAM:参数数量为26.24M,FLOPs为1.31G,Top-1错误率为19.44%,Top-5错误率为4.66%。

ImageNet
  • MobileNet V2 + NAM:参数数量为3.51M,FLOPs为0.32G,Top-1错误率为29.34%,Top-5错误率为10.18%。

  • MobileNet V2 + CBAM:参数数量为3.54M,FLOPs为0.32G,Top-1错误率为29.74%,Top-5错误率为10.66%。

结论

NAM通过抑制不重要的权重,显著提高了模型的计算效率,同时保持了与现有方法相当的性能。实验结果表明,NAM在ResNet和MobileNet上均优于SE、BAM、CBAM等现有注意力机制。未来,作者计划进一步优化NAM的性能,探索其在其他深度学习架构和应用中的效果。

主要代码展示

class Channel_Att(nn.Module): def __init__(self, channels, t=16): super(Channel_Att, self).__init__() self.channels = channels self.bn2 = nn.BatchNorm2d(self.channels, affine=True) def forward(self, x): residual = x x = self.bn2(x) weight_bn = self.bn2.weight.data.abs() / torch.sum(self.bn2.weight.data.abs()) x = x.permute(0, 2, 3, 1).contiguous() x = torch.mul(weight_bn, x) x = x.permute(0, 3, 1, 2).contiguous() x = torch.sigmoid(x) * residual return x def __init__(self, kernel_size=7): super(Spatial_Att, self).__init__() self.kernel_size = kernel_size self.softmax = nn.Softmax(dim=-1) def forward(self, x): # 输入特征:x,形状为 (B, C, H, W) residual = x # 计算每个像素的权重 (Pixel Normalization) pixel_weight = x.mean(dim=1, keepdim=True) # 按通道求均值,形状为 (B, 1, H, W) normalized_weight = pixel_weight / pixel_weight.sum(dim=(2, 3), keepdim=True) # 像素归一化 # 加权输入特征 x = x * normalized_weight # 按像素位置加权 x = torch.sigmoid(x) * residual # 与原输入相乘 return x class NAM(nn.Module): def __init__(self, channels): super(NAM, self).__init__() self.Channel_Att = Channel_Att(channels) self.Spatial_Att = Spatial_Att() def forward(self, x): x_out1 = self.Channel_Att(x) x_out2 = self.Spatial_Att(x_out1) return x_out2 if __name__ == '__main__': device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") nam = NAM(channels=32).to(device) # 将模型移动到设备 input = torch.rand(1, 32, 256, 256).to(device) # 将输入数据移动到设备 # # 参数量计算,可注释 # print("Model Summary:") # print(summary(nam, input_size=(1, 32, 256, 256), device=device.type)) # # 计算量计算,可注释 # flops = FlopCountAnalysis(nam, input) # print("\nFlop Count Table:") # print(flop_count_table(flops)) output = nam(input) print(f"\nInput shape: {input.shape}") print(f"Output shape: {output.shape}")

运行结果展示

每日分享最新的前沿技术,

助力快速发论文、模型涨点!

欢迎点赞关注,评论转发

添加下方个人微信

http://www.cnnetsun.cn/news/1351503.html

相关文章:

  • BetterNCM 插件导致网易云音乐启动失败问题分析
  • Camera:实时监控与数据交互的智能设备服务
  • 事件日志清理与痕迹擦除:wmiexec-Pro的eventlog模块深度应用
  • 告别厂商限制:Tuya TH05Z温湿度传感器接入Zigbee2MQTT完全指南
  • 复购率不理想如何用产品线组合提升长期价值
  • SimpleMem快速上手指南:5分钟搭建LLM智能体记忆系统
  • 基于LangChain的RAG与Agent智能体开发 - 使用LangChain调用大语言模型
  • MiniChain与Hugging Face Datasets:实现高效文档嵌入与相似度搜索的完整指南
  • Deepagents数据可视化:展示AI代理工作成果的终极指南
  • LDNetDiagnoService_IOS架构解析:Ping与Traceroute原理在iOS中的实现
  • 突破Ebitengine着色器限制:多重赋值问题的优雅解决方案
  • 告别繁琐切换!Micro编辑器智能文件类型识别功能详解:让编码效率提升300%的秘诀
  • STK信号处理秘籍:BiQuad滤波器与Chorus效果的应用技巧
  • Python入门:Python3 Pickle模块全面学习教程
  • Deepagents边境安全:监控边境活动的AI代理
  • 打造个性化观影系统:embyToLocalPlayer高级设置与自定义技巧
  • Performer-PyTorch核心组件解析:FastAttention如何实现O(n)复杂度
  • 提升家庭影院体验:embyToLocalPlayer与Kodi、Plex联动使用指南
  • cp-ddd-framework扩展机制详解:@Extension注解让业务逻辑灵活扩展
  • cp-ddd-framework架构演进:如何支撑业务系统从单体到微服务
  • 新手必读:Awesome Maintainers项目中的贡献指南与最佳实践
  • 解决Vim用户痛点:vim-quickui让命令交互变得简单直观的5个案例
  • 为什么选择RSpec-Mocks?探索Ruby测试框架中的强大测试替身解决方案
  • SideMenuController:打造iOS完美侧边菜单的终极Swift框架
  • 随机生成功能大揭秘:用ComfyUI Portrait Master探索无限创意可能性
  • 如何利用Browserify实现高效前端模块化开发:提升代码可维护性的完整指南
  • 如何参与Zellij路线图社区投票:决定终端工作区的未来功能优先级
  • 掌握Devise核心配置:从ALL到STRATEGIES的终极指南
  • 如何用WaveFunctionCollapse算法让孩子轻松理解概率与约束:从像素到城堡的神奇之旅
  • 如何使用Sails.js数据库迁移工具sails-migrations:完整指南