MMsegmentation 中的 DANet:双注意力机制场景分割算法解析与源码级实践指南
发布时间:2026/9/16 8:31:37来源:尧图网络
MMsegmentation 中的 DANet双注意力机制场景分割算法解析与源码级实践指南【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation导读DANetDual Attention Network是面向场景分割的经典注意力算法通过在传统空洞卷积 FCN 之上叠加位置注意力PAM与通道注意力CAM两个模块分别建模空间与通道维度的全局依赖从而显著提升分割精度。本文以 MMsegmentation 仓库中 configs/danet/README.md 为主线结合 da_head.py、self_attention_block.py 的源码实现与 configs/danet 目录下的完整配置文件系统讲解 DANet 的算法原理、模块级实现、三路监督损失机制、配置参数含义以及训练/测试/推理的完整实操流程。读完本文你将掌握在 MMsegmentation 中复现 DANet、改造双注意力解码头并迁移到新数据集的方法。一、算法背景从多尺度融合到注意力建模在 DANet 提出之前场景分割任务主要依赖两种思路捕获上下文一是扩大感受野空洞卷积二是多尺度特征融合如 PSPNet 的金字塔池化、DeepLab 的 ASPP。但这些方法都基于局部感受野的堆叠难以显式建模远距离像素之间的语义关联。DANet 的核心主张见 configs/danet/README.md 的 Abstract是利用自注意力机制自适应地将局部特征与其全局依赖相融合。具体做法是在传统空洞 FCN 之上并联两个注意力分支位置注意力模块Position Attention Module, PAM在空间维度上对任意位置的特征用所有位置特征的加权和来增强——相似的特征无论相距多远都会被关联起来解决同一类别因尺度、视角差异导致的特征不一致问题通道注意力模块Channel Attention Module, CAM在通道维度上通过整合所有通道图之间的关联来强调相互依赖的通道映射使网络自动聚焦于对判别最有利的语义通道。两个分支的输出相加进一步增强特征表示从而获得更精细的分割结果。原始论文在 Cityscapes、PASCAL Context 与 COCO Stuff 三个挑战性数据集上取得了当时的领先精度Cityscapes 测试集在未使用粗标注数据的情况下达到 81.5% Mean IoU。本文档对应的实现属于 MMsegmentation 官方算法集合metafile.yaml 中标注 License 为 Apache License 2.0Framework 为 PyTorch。二、源码级实现解析PAM、CAM 与 DAHeadDANet 在 MMsegmentation 中的完整实现集中在 mmseg/models/decode_heads/da_head.py共包含三个关键类PAM、CAM与注册为DAHead的解码头。此外PAM复用了通用的自注意力基类 mmseg/models/utils/self_attention_block.py 中的SelfAttentionBlock。2.1 SelfAttentionBlock通用的非局部/自注意力基类SelfAttentionBlockself_attention_block.py实现了标准的 key/query/value 自注意力流程其构造参数高度可配置key_in_channels/query_in_channelskey 与 query 投影的输入通道数channelskey/query 投影后的输出通道数out_channels最终输出通道数share_key_querykey 与 query 是否共享投影权重query_downsample/key_downsample对 query/key 特征的下采样模块key_query_num_convs/value_out_num_convs投影使用的卷积层数matmul_norm注意力图是否除以通道数的平方根做归一化with_out是否使用输出投影层。其forwardself_attention_block.py流程为分别对 query/key/value 做 1×1 卷积投影 → 展平为向量序列 → 计算相似度矩阵sim_map query key^T→ softmax 归一化 → 与 value 相乘得到聚合后的 context 特征。2.2 PAM位置注意力模块PAMda_head.py继承自SelfAttentionBlock在实例化时做了针对性配置class PAM(_SelfAttentionBlock): def __init__(self, in_channels, channels): super().__init__( key_in_channelsin_channels, query_in_channelsin_channels, channelschannels, out_channelsin_channels, share_key_queryFalse, query_downsampleNone, key_downsampleNone, key_query_num_convs1, key_query_normFalse, value_out_num_convs1, value_out_normFalse, matmul_normFalse, with_outFalse, conv_cfgNone, norm_cfgNone, act_cfgNone) self.gamma Scale(0)关键点在于注意力图不做matmul_norm缩放也不使用输出投影with_outFalse结构与经典 non-local 块一致引入可学习的缩放标量self.gamma Scale(0)初始化为 0保证训练初期注意力分支输出为 0、网络退化为普通 FCN从而保持训练稳定前向计算out self.gamma(attn_out) x即残差式融合——每个位置的特征等于自身特征加上所有位置特征的加权和权重即 softmax 后的空间相似度。2.3 CAM通道注意力模块CAMda_head.py是独立实现的轻量模块直接在原始特征上计算通道间的相似度class CAM(nn.Module): def __init__(self): super().__init__() self.gamma Scale(0) def forward(self, x): batch_size, channels, height, width x.size() proj_query x.view(batch_size, channels, -1) proj_key x.view(batch_size, channels, -1).permute(0, 2, 1) energy torch.bmm(proj_query, proj_key) energy_new torch.max( energy, -1, keepdimTrue)[0].expand_as(energy) - energy attention F.softmax(energy_new, dim-1) proj_value x.view(batch_size, channels, -1) out torch.bmm(attention, proj_value) out out.view(batch_size, channels, height, width) out self.gamma(out) x return out实现要点将特征展平为C × (H·W)矩阵通过torch.bmm计算C × C的通道相似度矩阵采用max(energy) - energy的最大项消减技巧再做 softmax而非直接对energy做 softmax这一做法与原文一致可以改善数值稳定性并改变注意力的分布形态同样使用Scale(0)初始化的gamma做残差融合out gamma * attn x。2.4 DAHead双注意力解码头整体流程DAHeadda_head.py通过MODELS.register_module()注册可直接在配置中通过typeDAHead引用。其构造参数除继承自BaseDecodeHead的in_channels、channels、num_classes、dropout_ratio、norm_cfg、align_corners、loss_decode等外独有参数为pam_channels (int)PAM 中 key/query 投影的中间通道数默认配置为 64。DAHead内部构建了两条并行分支PAM 分支pam_in_conv3×3 ConvModule→PAM→pam_out_conv→pam_conv_seg1×1 分类卷积CAM 分支cam_in_conv3×3 ConvModule→CAM→cam_out_conv→cam_conv_seg1×1 分类卷积。forwardda_head.py返回一个三元组feat_sum pam_feat cam_feat pam_cam_out self.cls_seg(feat_sum) return pam_cam_out, pam_out, cam_out即三个 logits融合分支pam_cam、PAM 独立分支pam、CAM 独立分支cam。这三个 logits 在训练时全部参与监督详见第四节而推理时仅使用融合分支def predict(self, inputs, batch_img_metas, test_cfg, **kwargs): Forward function for testing, only pam_cam is used. seg_logits self.forward(inputs)[0] return self.predict_by_feat(seg_logits, batch_img_metas, **kwargs)三、配置文件深度解析3.1 模型基类danet_r50-d8.pyDANet 的 R50-D8 模型骨架定义在 configs/base/models/danet_r50-d8.py这是所有 DANet 配置共享的基础核心字段如下norm_cfg dict(typeSyncBN, requires_gradTrue) data_preprocessor dict( typeSegDataPreProcessor, mean[123.675, 116.28, 103.53], std[58.395, 57.12, 57.375], bgr_to_rgbTrue, pad_val0, seg_pad_val255) model dict( typeEncoderDecoder, data_preprocessordata_preprocessor, pretrainedopen-mmlab://resnet50_v1c, backbonedict( typeResNetV1c, depth50, num_stages4, out_indices(0, 1, 2, 3), dilations(1, 1, 2, 4), # 即 D8 的含义第3、4阶段空洞率为2、4 strides(1, 2, 1, 1), norm_cfgnorm_cfg, norm_evalFalse, stylepytorch, contract_dilationTrue), decode_headdict( typeDAHead, in_channels2048, # 主干第4阶段输出通道数 in_index3, # 取 backbone 第4个输出特征 channels512, # PAM/CAM 分支内部工作通道数 pam_channels64, # PAM 中 key/query 投影通道数 dropout_ratio0.1, num_classes19, # Cityscapes 类别数 norm_cfgnorm_cfg, align_cornersFalse, loss_decodedict( typeCrossEntropyLoss, use_sigmoidFalse, loss_weight1.0)), auxiliary_headdict( typeFCNHead, in_channels1024, in_index2, channels256, num_convs1, concat_inputFalse, dropout_ratio0.1, num_classes19, norm_cfgnorm_cfg, align_cornersFalse, loss_decodedict( typeCrossEntropyLoss, use_sigmoidFalse, loss_weight0.4)), train_cfgdict(), test_cfgdict(modewhole))各字段含义与作用backboneResNetV1c带 stem 处 7×7 卷积替换为 3 个 3×3 卷积的变体dilations(1, 1, 2, 4)表示第 3、4 阶段分别使用空洞率 2 和 4保持输出分辨率不下降这是 D8dilation 8的由来decode_head即DAHead。in_channels2048对应 backbone 最高层特征channels512决定 PAM/CAM 分支内部的计算量pam_channels64控制 PAM 的 key/query 投影维度越小越省显存与计算dropout_ratio0.1在分类卷积前施加 dropout 抑制过拟合auxiliary_head在 backbone 第 3 阶段in_index21024 通道额外挂一个 FCN 辅助头做中间层深度监督loss_weight0.4控制辅助损失占比这也是 DANet 精度的重要来源之一test_cfgdict(modewhole)推理时整图一次前向不做滑窗切块切换为modeslide可配合crop_size进行滑窗推理以降低显存。3.2 各数据集入口配置的继承关系configs/danet 目录下的 16 个配置文件全部通过_base_继承上述基类按数据集分为三类Cityscapes4xb2crop 512x1024 或 769x769以 danet_r50-d8_4xb2-40k_cityscapes-512x1024.py 为例_base_ [ ../_base_/models/danet_r50-d8.py, ../_base_/datasets/cityscapes.py, ../_base_/default_runtime.py, ../_base_/schedules/schedule_40k.py ] crop_size (512, 1024) data_preprocessor dict(sizecrop_size) model dict(data_preprocessordata_preprocessor)配置只覆盖了输入裁剪尺寸512×1024类别数 19、训练调度 40k 均继承自基类。切换danet_r101-d8_...文件时仅需改两处pretrainedopen-mmlab://resnet101_v1c与backbonedict(depth101)见 danet_r101-d8_4xb4-160k_ade20k-512x512.py 的写法。ADE20K4xb4crop 512x512以 danet_r50-d8_4xb4-160k_ade20k-512x512.py 为例_base_ [ ../_base_/models/danet_r50-d8.py, ../_base_/datasets/ade20k.py, ../_base_/default_runtime.py, ../_base_/schedules/schedule_160k.py ] crop_size (512, 512) data_preprocessor dict(sizecrop_size) model dict( data_preprocessordata_preprocessor, decode_headdict(num_classes150), # ADE20K 有 150 类 auxiliary_headdict(num_classes150))Pascal VOC 2012 Aug4xb4crop 512x512以 danet_r50-d8_4xb4-20k_voc12aug-512x512.py 为代表数据基类换为voc12aug.py训练迭代数随调度基类20k/40k变化类别数相应覆盖为 21。可以看出从 Cityscapes 迁移到新数据集只需要替换数据集基类、修改num_classes、调整crop_size与调度即可DANet 解码头本身无需改动。3.3 训练调度配置40k 调度的完整定义见 configs/base/schedules/schedule_40k.pyoptimizer dict(typeSGD, lr0.01, momentum0.9, weight_decay0.0005) optim_wrapper dict(typeOptimWrapper, optimizeroptimizer, clip_gradNone) param_scheduler [ dict( typePolyLR, eta_min1e-4, power0.9, begin0, end40000, by_epochFalse) ] train_cfg dict(typeIterBasedTrainLoop, max_iters40000, val_interval4000) val_cfg dict(typeValLoop) test_cfg dict(typeTestLoop) default_hooks dict( timerdict(typeIterTimerHook), loggerdict(typeLoggerHook, interval50, log_metric_by_epochFalse), param_schedulerdict(typeParamSchedulerHook), checkpointdict(typeCheckpointHook, by_epochFalse, interval4000), sampler_seeddict(typeDistSamplerSeedHook), visualizationdict(typeSegVisualizationHook))要点优化器SGD初始学习率 0.01momentum 0.9weight_decay 5e-4多卡 4xb2 时的标准配置若改单卡可参考其他配置按线性缩放学习率学习率策略PolyLRpower0.9最小学习率eta_min1e-4按迭代by_epochFalse衰减训练循环迭代制训练IterBasedTrainLoop每 4000 迭代验证一次并保存一次 checkpointCheckpointHook interval4000共 40000 迭代。80k/160k 配置只需把end、max_iters等比放大。四、三路监督损失机制DANet 的与众不同之处在于训练时同时优化三个分割输出。loss_by_featda_head.py将前向得到的(pam_cam, pam, cam)三元组分别计算交叉熵损失并用add_prefix加上pam_cam_、pam_、cam_前缀写入日志def loss_by_feat(self, seg_logit, batch_data_samples, **kwargs): pam_cam_seg_logit, pam_seg_logit, cam_seg_logit seg_logit loss dict() loss.update(add_prefix( super().loss_by_feat(pam_cam_seg_logit, batch_data_samples), pam_cam)) loss.update(add_prefix( super().loss_by_feat(pam_seg_logit, batch_data_samples), pam)) loss.update(add_prefix( super().loss_by_feat(cam_seg_logit, batch_data_samples), cam)) return loss配合基类配置中的辅助头auxiliary_headFCNHeadloss_weight0.4一次迭代总共包含 4 项损失融合分支权重 1.0、PAM 分支权重 1.0、CAM 分支权重 1.0与辅助头权重 0.4。这种三路解码监督 中间层深度监督的复合监督策略迫使 PAM、CAM 各自学到有判别力的空间/通道注意力再在融合分支中互补。推理时则只走融合分支见前文predict实现因此推理成本与单分支解码头相当。五、训练、测试与推理实操以下命令均在仓库根目录执行多卡命令通过 tools/dist_train.sh 与 tools/dist_test.sh 分发。5.1 单卡训练python tools/train.py configs/danet/danet_r50-d8_4xb2-40k_cityscapes-512x1024.py \ --work-dir work_dirs/danet_r50-d8_4xb2-40k_cityscapes-512x10245.2 多卡分布式训练bash tools/dist_train.sh configs/danet/danet_r50-d8_4xb2-40k_cityscapes-512x1024.py 88为 GPU 卡数与配置中4xb24 卡 × batch 2的总 batch size 8 对应实际卡数请按可用资源调整。5.3 测试与指标评估python tools/test.py configs/danet/danet_r50-d8_4xb2-40k_cityscapes-512x1024.py \ /path/to/checkpoint_file --eval mIoU多卡版本bash tools/dist_test.sh configs/danet/danet_r50-d8_4xb2-40k_cityscapes-512x1024.py \ /path/to/checkpoint_file 8 --eval mIoU--eval mIoU使用 mIoU 评估器 计算 mIoU 与 mAcc 等指标测试时test_cfgdict(modewhole)整图前向。若需多尺度 翻转对应结果表中的 mIoU(msflip) 指标可配合工具脚本进行实际数值可见 metafile.yaml 中各模型的mIoU(msflip)记录。5.4 单张图片推理可视化使用 demo/image_demo.py 即可快速体验 DANet 的分割效果python demo/image_demo.py demo/demo.png \ configs/danet/danet_r50-d8_4xb2-40k_cityscapes-512x1024.py \ /path/to/checkpoint_file --device cuda:0 --out-file result.png预训练权重与训练日志的下载地址统一记录在 configs/danet/metafile.yaml 的Weights与Training log字段中可按模型名如danet_r50-d8_4xb2-40k_cityscapes-512x1024查找对应下载链接。六、基准测试结果以下三张结果表完整继承自 configs/danet/README.md其中 mIoU、内存与推理耗时均为官方在 V100 上的实测记录config 列已转换为仓库内相对路径模型与日志下载地址见 metafile.yaml。6.1 CityscapesMethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)configDANetR-50-D8512x1024400007.42.66V10078.74-configDANetR-101-D8512x10244000010.91.99V10080.52-configDANetR-50-D8769x769400008.81.56V10078.8880.62configDANetR-101-D8769x7694000012.81.07V10079.8881.47configDANetR-50-D8512x102480000--V10079.34-configDANetR-101-D8512x102480000--V10080.41-configDANetR-50-D8769x76980000--V10079.2780.96configDANetR-101-D8769x76980000--V10080.4782.02config6.2 ADE20KMethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)configDANetR-50-D8512x5128000011.521.20V10041.6642.90configDANetR-101-D8512x512800001514.18V10043.6445.19configDANetR-50-D8512x512160000--V10042.4543.25configDANetR-101-D8512x512160000--V10044.1745.02config6.3 Pascal VOC 2012 AugMethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)configDANetR-50-D8512x512200006.520.94V10074.4575.69configDANetR-101-D8512x512200009.913.76V10076.0277.23configDANetR-50-D8512x51240000--V10076.3777.29configDANetR-101-D8512x51240000--V10076.5177.32config从结果可以观察到的规律基于上述官方记录在 Cityscapes 上 512x1024 输入下 R101 比 R50 提升约 1.8 个点同一 backbone 下 769x769 比 512x1024 输入带来 0.1~0.3 个点的提升msflip 多尺度测试普遍比单尺度高 1~1.6 个点ADE20K 与 VOC 场景下将训练迭代从 80k 延长到 160k或 20k 到 40k同样能带来稳定增益。七、复现与扩展建议快速复现下载 metafile.yaml 中对应模型的权重用第五节命令直接测试即可对照上表指标迁移新数据集复制 danet_r50-d8_4xb4-160k_ade20k-512x512.py 式入口配置替换_base_中的数据集基类、修改num_classes与crop_size类别数变化后pam_channels、channels无需调整显存受限场景可降低channels512→256或pam_channels64→32或将test_cfg.mode改为slide配crop_size滑窗推理模块复用PAM、CAM均可在自己的解码头中直接from mmseg.models.decode_heads.da_head import PAM, CAM导入复用PAM底层依赖的 SelfAttentionBlock 也是实现其他 non-local/自注意力模块的通用积木。附录引用规范若在研究中引用 DANet 算法官方推荐引用格式如下来自 configs/danet/README.mdarticle{fu2018dual, title{Dual Attention Network for Scene Segmentation}, author{Jun Fu, Jing Liu, Haijie Tian, Yong Li, Yongjun Bao, Zhiwei Fang,and Hanqing Lu}, booktitle{The IEEE Conference on Computer Vision and Pattern Recognition (CVPR)}, year{2019} }【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网