尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

MMSegmentation 中的 EMANet:期望最大化注意力语义分割网络的原理、源码解析与 Cityscapes 实战

MMSegmentation 中的 EMANet:期望最大化注意力语义分割网络的原理、源码解析与 Cityscapes 实战 MMSegmentation 中的 EMANet期望最大化注意力语义分割网络的原理、源码解析与 Cityscapes 实战【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation本文围绕 OpenMMLab 语义分割工具箱 MMSegmentation 中对EMANetExpectation-Maximization Attention Networks for Semantic Segmentation的官方实现展开系统讲解其将自注意力机制重写为期望最大化EM迭代的算法思想、EMAHead 源码 的逐层实现细节、4 组 Cityscapes 官方配置的完整参数解读并给出可复现的训练、测试命令。读完本文你将能够理解 EMANet 为何能以紧凑的基bases集合替代全图两两注意力、如何在 MMSegmentation 中直接训练与评估 EMANet以及如何基于现有配置修改num_bases、num_stages等核心超参数。一、EMANet 是什么把注意力机制改写成 EM 迭代传统自注意力机制需要为每个位置计算它与所有其他位置之间的注意力权重虽然能够捕获计算机视觉任务所需的长距离依赖但注意力图的计算与存储开销巨大。EMANet论文标题Expectation-Maximization Attention Networks for Semantic SegmentationICCV 2019的核心思想是不再对全图两两建模而是把注意力机制重新表述为期望最大化Expectation-Maximization, EM过程用迭代方式估计一组紧凑得多的基bases注意力图只在这组基上计算从而通过基的加权求和得到低秩表示能够抑制输入中的噪声信息对输入方差具有更好的鲁棒性在内存和计算上更友好同时设计了基的维护maintenance与归一化normalization方法来稳定训练过程。EMANet 在 PASCAL VOC、PASCAL Context、COCO Stuff 等主流语义分割基准上取得了当时的领先结果详见论文。本文所基于的 MMSegmentation 仓库实现了该论文对应的EMAHead解码头并提供了基于 ResNet-50/ResNet-101D8 空洞策略骨干网络、面向 Cityscapes 数据集的 4 组完整训练配置。官方实现的核心代码位于 ema_head.py其中EMAModule封装了 EM 迭代的核心运算EMAHead则负责与骨干网络输出的衔接和最终分类。二、算法原理期望最大化注意力模块的工作机制EMANet 论文的核心贡献是Expectation-Maximization Attention (EMA) 模块。从算法层面可以将其理解为对输入特征反复执行两个步骤E 步期望步固定当前的基集合计算每个像素特征到每个基的相似度经 softmax 得到注意力图即把每个像素归属到各个基上的概率分布M 步最大化步固定注意力图用所有像素特征的加权平均重新估计基使基更好地代表特征簇的中心。上述两步被迭代执行num_stages次基集合在迭代中不断精化最终用最后一次 E 步得到的注意力图和基通过加权求和重建输入特征。由于基的数量num_bases默认 64远小于像素数重建得到的特征表示是低秩的天然具备去噪与压缩特性。为了保证训练稳定EMANet 还引入了基的维护策略训练阶段以滑动平均momentum的方式将当前 batch 的基更新到全局基缓冲区中使得基的估计在多次迭代之间保持一致。三、源码级解析EMAHead 与 EMAModule 的实现细节MMSegmentation 对 EMANet 的实现位于 mmseg/models/decode_heads/ema_head.py由两个核心类构成。3.1 EMAModuleEM 迭代的核心运算EMAModule的构造函数接收三个关键参数参数含义默认值来自配置channels模块输入特征通道数512ema_channelsnum_bases基的数量64num_stagesEM 迭代次数须 ≥ 13momentum训练时更新全局基的动量0.1初始化时模块会创建形状为[1, channels, num_bases]的基张量用标准差为sqrt(2/num_bases)的正态分布初始化并在通道维做 L2 归一化后注册为 bufferself.bases因此它不参与梯度更新而是作为 EM 迭代的起点与全局维护对象。forward的实现与论文公式严格对应见 ema_head.py将特征[B, C, H, W]展平为[B, C, H*W]并将全局基广播到 batch 维度在torch.no_grad()下迭代num_stages次E 步attention einsum(bcn,bck-bnk, feats, bases)再对基维度做 softmax得到每个像素对每个基的注意力权重[B, H*W, num_bases]对注意力做 L1 归一化后执行 M 步bases einsum(bcn,bnk-bck, feats, attention_normed)即用特征按注意力加权平均更新基随后对基做 L2 归一化用最后一次 E 步得到的注意力图和更新后的基重建特征einsum(bck,bnk-bcn, bases, attention)训练阶段维护全局基取 batch 内基的均值经分布式reduce_meanall-reduce 平均见 reduce_mean 函数后按bases (1 - momentum) * bases momentum * bases滑动更新self.bases。值得注意的实现细节是EM 迭代整体在torch.no_grad()中执行基的更新不参与反向传播反向梯度只经由最后一次重建路径回传这既稳定了训练也节省了显存与计算。3.2 EMAHead解码头的组装与前后衔接EMAHead继承自 BaseDecodeHead构造参数包括ema_channelsEMA 模块的通道数配置中为 512num_bases/num_stages基的数量与 EM 迭代次数64 / 3concat_input分类前是否将输入与输出特征拼接默认Truemomentum基的更新动量默认 0.1。EMAHead的前向流程见 ema_head.py由以下子模块串联ema_in_conv3×3 卷积把骨干输出2048 通道压缩到ema_channels512得到特征feats同时保存一份作为恒等映射identityema_mid_conv1×1 卷积把特征投影到(0, inf)区间使其适合后续的 softmax 注意力计算该卷积参数被requires_grad False冻结见 ema_head.pyema_module执行上述 EM 迭代并输出重建特征对重建特征做 ReLU、经ema_out_conv1×1 卷积后与恒等映射相加并 ReLU形成残差式融合bottleneck3×3 卷积恢复到channels256若concat_inputTrue将ema_in_conv的输入x与输出拼接经conv_cat3×3 卷积融合最后通过基类提供的cls_seg含 Dropout 与 1×1conv_seg卷积输出逐像素分类 logits。从结构可以看出EMA 模块以卷积降维 → EM 重建 → 残差融合 → 分类的方式插入到解码头中既保留了残差学习的稳定性又将全局上下文建模替换为低秩的 EM 表示。3.3 测试用例佐证仓库在 tests/test_models/test_heads/test_ema_head.py 中提供了对EMAHead的单元测试以in_channels4, ema_channels3, channels2, num_stages3, num_bases2构造头部验证ema_mid_conv的所有参数requires_grad均为False并断言ema_module属性存在输入[1, 4, 23, 23]的特征图断言输出形状为(1, num_classes, 23, 23)。该测试从侧面印证了上述实现细节中间投影卷积被冻结、EMA 模块作为独立子模块存在、输出空间分辨率与输入保持一致只改变通道为类别数。四、配置详解基于 ResNet 骨干的 4 组 Cityscapes 配置EMANet 在仓库中的全部配置文件位于 configs/emanet共 4 个配置文件骨干裁剪尺寸emanet_r50-d8_4xb2-80k_cityscapes-512x1024.pyR-50-D8512×1024emanet_r101-d8_4xb2-80k_cityscapes-512x1024.pyR-101-D8512×1024emanet_r50-d8_4xb2-80k_cityscapes-769x769.pyR-50-D8769×769emanet_r101-d8_4xb2-80k_cityscapes-769x769.pyR-101-D8769×769注意原 README 与 metafile.yaml 中 R-50-D8-512×1024 的配置链接存在eemanet_拼写笔误仓库中实际存在的文件名为emanet_r50-d8_4xb2-80k_cityscapes-512x1024.py。4.1 模型骨架模型结构与 EMAHead 参数所有配置均继承自基础模型配置 configs/base/models/emanet_r50-d8.py其核心内容包括数据预处理器SegDataPreProcessorImageNet 风格的均值/方差mean[123.675, 116.28, 103.53]、std[58.395, 57.12, 57.375]、bgr_to_rgbTrue、分割标签填充值seg_pad_val255骨干网络ResNetV1cdepth50、out_indices(0,1,2,3)、空洞率dilations(1,1,2,4)即 D8 策略、contract_dilationTrue预训练权重来自open-mmlab://resnet50_v1c解码头EMAHeaddecode_headin_channels2048、in_index3取骨干第 4 个阶段的输出channels256、ema_channels512num_bases64基的数量决定低秩表示的秩num_stages3EM 迭代次数越大基越精化但计算越多momentum0.1基的滑动更新动量dropout_ratio0.1、num_classes19Cityscapes 类别数、align_cornersFalseloss_decodeCrossEntropyLossloss_weight1.0辅助头FCNHeadauxiliary_head作用于in_channels1024、in_index2的中层特征loss_weight0.4与主头共同监督推理设置test_cfgdict(modewhole)整图推理。4.2 数据集、调度与运行时4 组配置分别继承 configs/base/datasets/cityscapes.py、configs/base/datasets/cityscapes_769x769.py、configs/base/schedules/schedule_80k.py 与 configs/base/default_runtime.py数据管线训练采用RandomResize缩放比 0.5~2.0、keep_ratioTrue→RandomCropcat_max_ratio0.75防止裁剪到单类→RandomFlip→PhotoMetricDistortion→PackSegInputs测试采用固定Resize到(2048, 1024)优化器与调度SGDlr0.01、momentum0.9、weight_decay0.0005配合OptimWrapper学习率采用PolyLRpower0.9、eta_min1e-4、按 iteration 衰减训练循环IterBasedTrainLoopmax_iters80000、val_interval8000Checkpoint 每 8000 iter 保存一次by_epochFalse运行时default_scopemmseg、cudnn_benchmarkTrue、nccl分布式后端、SegLocalVisualizer本地可视化、tta_modelSegTTAModel支持多尺度 翻转 TTA评估IoUMetric指标为mIoU。4.3 512×1024 与 769×769 两组配置的差异以 emanet_r50-d8_4xb2-80k_cityscapes-512x1024.py 为例512×1024 版本仅需在继承骨架后覆盖_base_ [ ../_base_/models/emanet_r50-d8.py, ../_base_/datasets/cityscapes.py, ../_base_/default_runtime.py, ../_base_/schedules/schedule_80k.py ] crop_size (512, 1024) data_preprocessor dict(sizecrop_size) model dict(data_preprocessordata_preprocessor)而 769×769 版本emanet_r50-d8_4xb2-80k_cityscapes-769x769.py除了切换数据集与裁剪尺寸外还做了三处关键调整crop_size (769, 769) model dict( data_preprocessordata_preprocessor, decode_headdict(align_cornersTrue), auxiliary_headdict(align_cornersTrue), test_cfgdict(modeslide, crop_size(769, 769), stride(513, 513)))align_cornersTrue配合 769 分辨率保证双线性插值坐标对齐test_cfg从modewhole切换为modeslide测试时用 769×769 窗口、513×513 步长滑窗推理避免高分辨率整图推理超出显存这也是 769×769 配置显存占用更高、但 mIoU 更高的原因之一。R-101 版本只需将基础模型替换为emanet_r101-d8对应骨架depth101其余结构与上述一致。五、训练、测试与可视化实战在正确安装 MMSegmentation 及其依赖PyTorch、MMEngine、MMCV并准备好 Cityscapes 数据集默认根目录data/cityscapes/图像位于leftImg8bit/、标注位于gtFine/可参考仓库文档目录下的数据集准备说明后即可开始训练与评估。5.1 单机训练python tools/train.py configs/emanet/emanet_r50-d8_4xb2-80k_cityscapes-512x1024.py5.2 多卡分布式训练仓库提供了 dist_train.sh 脚本训练资源为 4×V100 GPU见下文结果表bash tools/dist_train.sh configs/emanet/emanet_r50-d8_4xb2-80k_cityscapes-512x1024.py 4训练过程会按schedule_80k每 8000 iter 执行一次验证并保存 checkpoint日志与权重默认输出到work_dirs/下以配置名命名的目录中。5.3 测试与评估使用 tools/test.py 加载训练好的权重评估 mIoUpython tools/test.py configs/emanet/emanet_r50-d8_4xb2-80k_cityscapes-512x1024.py \ /path/to/checkpoint.pth --out results.pkl或通过--tta开启多尺度 水平翻转的测试时增强配置中已定义tta_pipeline与img_ratios[0.5, 0.75, 1.0, 1.25, 1.5, 1.75]对应结果表中的 mIoU(msflip) 指标。5.4 单图推理演示仓库 demo/image_demo.py 支持对任意图片执行推理并叠加可视化分割结果可结合训练好的权重直接观察 EMANet 在城市街景上的分割效果。六、Cityscapes 官方结果与复现参考以下结果摘自 configs/emanet/README.md 并可由 metafile.yaml 交叉印证Batch Size 均为 8即 4 卡 × 2 样本/卡学习率调度均为 80k iterationsMethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)EMANetR-50-D8512×1024800005.44.58V10077.5979.44EMANetR-101-D8512×1024800006.22.87V10079.1081.21EMANetR-50-D8769×769800008.91.97V10079.3380.49EMANetR-101-D8769×7698000010.11.22V10079.6281.00解读要点骨干越强、分辨率越高mIoU 越高R-101-D8 769×769 组合达到 79.62msflip 后 81.00分辨率提升带来显存与推理时间代价769×769 相比 512×1024显存约增加 65%75%推理 fps 明显下降R-50 从 4.58 降到 1.97R-101 从 2.87 降到 1.22多尺度 翻转msflip普遍带来 1.42.1 个点的提升。这些数字可作为复现实验的对齐基准如果实际训练结果与表中数值偏差较大可优先核对数据管线RandomResize 缩放比、RandomCrop 的cat_max_ratio、PolyLR 超参以及align_corners/test_cfg的设置是否与对应配置一致。七、基于现有配置的定制方向结合 configs/base/models/emanet_r50-d8.py 中的参数可以低成本地开展 EMANet 消融实验调整基的数量修改decode_head中的num_bases如 32/64/128观察低秩表示的容量与精度、显存之间的权衡调整 EM 迭代次数修改num_stages如 1/3/5验证更多迭代对基精化的收益更换骨干或迁移数据集将backbone替换为仓库 configs/base/models 下其他骨架如 R-101、HRNet并把num_classes、数据集配置切换到目标数据集关闭输入拼接设置concat_inputFalse验证残差拼接分支对精度的贡献。由于 EM 迭代在no_grad下进行、基的维护使用动量更新改动momentum会直接影响全局基的稳定性是研究训练收敛行为时的关键旋钮。八、引用若在研究中使用了 EMANet 或本仓库实现可参考如下 BibTeX来自 configs/emanet/README.mdinproceedings{li2019expectation, title{Expectation-maximization attention networks for semantic segmentation}, author{Li, Xia and Zhong, Zhisheng and Wu, Jianlong and Yang, Yibo and Lin, Zhouchen and Liu, Hong}, booktitle{Proceedings of the IEEE International Conference on Computer Vision}, pages{9167--9176}, year{2019} }小结EMANet 用 EM 迭代把昂贵的全图自注意力替换为对紧凑基集合的估计与重建在保持长距离上下文建模能力的同时显著降低了内存与计算开销。在 MMSegmentation 中它的完整实现收敛于 ema_head.py 一个文件之内配以 configs/emanet 下 4 组开箱即用的 Cityscapes 配置即可直接复现 77.5979.62 的 mIoUmsflip 后最高 81.00。无论是学习注意力机制的低秩近似思想还是作为语义分割 baseline 开展实验EMANet 都是值得深入研究的对象。【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表