资讯动态

MMSegmentation 中的 ResNeSt 分割骨干网络:Split-Attention 原理、源码解析与 Cityscapes/ADE20K 实战配置

发布时间:2026/9/15 17:00:43 来源:尧图企业网站定制
MMSegmentation 中的 ResNeSt 分割骨干网络Split-Attention 原理、源码解析与 Cityscapes/ADE20K 实战配置【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation本文基于 MMSegmentation 仓库 configs/resnest/README.md 与 mmseg/models/backbones/resnest.py 编写。ResNeStSplit-Attention Networks是一种将通道注意力channel-wise attention与多路径表示multi-path representation模块化的骨干网络MMSegmentation 将其作为 FCN、PSPNet、DeepLabV3、DeepLabV3 等多种分割头的高性能 backbone 接入并提供了 Cityscapes 与 ADE20K 上的完整训练配置与预训练权重。读完本文你将掌握 ResNeSt 的 Split-Attention 计算块原理、MMSegmentation 中 ResNeSt 的源码实现细节以及如何基于仓库配置一键复现训练与推理。一、ResNeSt 是什么论文背景与核心思想ResNeStSplit-Attention Networks出自论文ResNeSt: Split-Attention NetworksarXiv:2004.08955Zhang 等人2020。其核心观察是特征图注意力featuremap attention与多路径表示multi-path representation对视觉识别任务至关重要。ResNeSt 将这两者结合提出一个模块化的计算单元——Split-Attention 块在不同的网络分支branches上施加通道维度的注意力channel-wise attention以捕捉跨特征交互cross-feature interactions让不同分支学习多样化的表示diverse representations整个设计最终收敛为一个简单、统一的计算块仅需少量变量即可参数化。原论文报告ResNeSt 在图像分类上以更优的精度/延迟权衡超越 EfficientNet作为骨干网络在多个公开基准上取得了优秀的迁移学习结果并被 COCO-LVIS 挑战赛的获胜方案所采用。论文信息与引用该论文对应的 BibTeX 引用来自 configs/resnest/README.md 的 Citation 小节article{zhang2020resnest, title{ResNeSt: Split-Attention Networks}, author{Zhang, Hang and Wu, Chongruo and Zhang, Zhongyue and Zhu, Yi and Zhang, Zhi and Lin, Haibin and Sun, Yue and He, Tong and Muller, Jonas and Manmatha, R. and Li, Mu and Smola, Alexander}, journal{arXiv preprint arXiv:2004.08955}, year{2020} }二、MMSegmentation 中的 ResNeSt源码级原理剖析ResNeSt 在 MMSegmentation 中的实现位于 mmseg/models/backbones/resnest.py由三个核心组件构成RSoftmax、SplitAttentionConv2d与Bottleneck并通过MODELS.register_module()注册为ResNeSt继承自 ResNet 的 V1d 变体见 mmseg/models/backbones/resnet.py 中ResNetV1d。2.1 RSoftmaxRadix Softmax 模块RSoftmaxresnest.py 第 16-37 行是 Split-Attention 中产生注意力权重的激活函数行为取决于radix参数当radix 1时将输入重塑为(batch, groups, radix, -1)并转置为(batch, radix, groups, -1)在 radix 维度上做softmax实现多个分支间的竞争归一化当radix 1时退化为sigmoid激活等价于 SE-Net 式的通道门控。class RSoftmax(nn.Module): def forward(self, x): batch x.size(0) if self.radix 1: x x.view(batch, self.groups, self.radix, -1).transpose(1, 2) x F.softmax(x, dim1) x x.reshape(batch, -1) else: x torch.sigmoid(x) return x2.2 SplitAttentionConv2dSplit-Attention 卷积单元SplitAttentionConv2dresnest.py 第 40-144 行是 ResNeSt 的核心计算块其构造参数包括参数类型默认值说明in_channelsint必填输入通道数与nn.Conv2d一致channelsint必填输出通道数kernel_sizeint/tuple必填卷积核大小strideint/tuple1步长paddingint/tuple0填充dilationint/tuple1膨胀率groupsint1分组数radixint2SplitAtConv2d 的分支基数reduction_factorint4inter_channels的缩减因子conv_cfgdictNone卷积层配置默认使用普通 conv2dnorm_cfgdictdict(typeBN)归一化层配置dcndictNoneDCN可变形卷积配置其前向过程resnest.py 第 118-144 行完整实现了 Split-Attention 的标准四步流程Split分支卷积先经过groups * radix分组的卷积self.conv输出channels * radix个通道当radix 1时将输出按 radix 拆分为多个分支splits并对所有分支求和得到gapSE 式压缩对gap做全局自适应平均池化F.adaptive_avg_pool2d(gap, 1)通道注意力生成依次经过fc1将通道数压缩到inter_channels、BN ReLU、fc2恢复为channels * radix再经RSoftmax得到注意力权重加权融合Split Attention当radix 1时将注意力权重按 radix 拆分后与各分支逐元素相乘并求和torch.sum(attens * splits, dim1)当radix 1时直接与原始特征相乘。其中inter_channels的计算式为max(in_channels * radix // reduction_factor, 32)即压缩后通道数至少为 32保证小通道数场景下注意力瓶颈不会过窄。此外该模块支持 DCN可变形卷积当传入dcn且未设置fallback_on_stride时会断言conv_cfg必须为 None并将dcn作为卷积配置resnest.py 第 79-84 行。2.3 Bottleneck集成 Split-Attention 的残差块ResNeSt 的Bottleneckresnest.py 第 147-267 行继承自 ResNet 的Bottleneckexpansion 4关键差异在于conv2 被替换为SplitAttentionConv2d并传入groups、radix、reduction_factor等参数resnest.py 第 201-213 行avg_down_stride当启用且conv2_stride 1时stride 不再放在 3x3 卷积中而是在卷积之后插入一个nn.AvgPool2d(3, conv2_stride, padding1)下采样resnest.py 第 185、216-217、241-242 行支持with_cpcheckpoint以节省显存当x.requires_grad时使用torch.utils.checkpoint包装内层前向resnest.py 第 260-263 行。class Bottleneck(_Bottleneck): expansion 4 def __init__(self, inplanes, planes, groups1, base_width4, base_channels64, radix2, reduction_factor4, avg_down_strideTrue, **kwargs): super().__init__(inplanes, planes, **kwargs) # ... self.conv2 SplitAttentionConv2d( width, width, kernel_size3, stride1 if self.avg_down_stride else self.conv2_stride, paddingself.dilation, dilationself.dilation, groupsgroups, radixradix, reduction_factorreduction_factor, conv_cfgself.conv_cfg, norm_cfgself.norm_cfg, dcnself.dcn)2.4 ResNeSt 骨干注册、深度配置与 V1d 继承ResNeSt类resnest.py 第 270-318 行通过MODELS.register_module()注册到 MMSegmentation 的模型注册表中可直接在配置中以typeResNeSt引用。支持的深度arch_settingsresnest.py 第 288-293 行深度各阶段 Bottleneck 数量50(3, 4, 6, 3)101(3, 4, 23, 3)152(3, 8, 36, 3)200(3, 24, 36, 3)构造参数参数默认值说明groups1Bottleneck 中 3x3 卷积的分组数base_width4每组宽度64x4d 表示groups64, width_per_group4radix2SplitAttentionConv2d 的分支基数reduction_factor4SplitAttentionConv2d 中间通道的缩减因子avg_down_strideTrue是否用平均池化实现下采样其余 kwargs-继承自 ResNet/ResNetV1d 的参数depth、in_channels、stem_channels、out_indices、frozen_stages等与 ResNetV1d 的关系ResNeSt继承自ResNetV1dresnet.py 第 703-712 行因此天然具备 V1d 的两个特性deep_stem用三个 3x3 卷积通道为stem_channels//2 → stem_channels//2 → stem_channels替换标准 ResNet 的单个 7x7 卷积resnet.py 第 591-624 行avg_down下采样残差块中先做 2x2 stride2 的平均池化再使用 stride1 的卷积。这解释了为何所有 ResNeSt 配置都将stem_channels设为128ResNeSt 原始设计采用 64→128 的 stem 结构相比 ResNet 的 64 通道 stem 更宽从而与官方预训练权重的结构保持一致。三、仓库配置解析如何在 MMSegmentation 中启用 ResNeSt3.1 配置复用模式configs/resnest/目录下共 8 个训练配置全部采用最小覆盖base继承模式只替换 backbone 为ResNeSt其余分割头、数据集、调度、runtime完全继承自对应算法的标准 ResNet-101 配置。例如resnest_s101-d8_fcn_4xb2-80k_cityscapes-512x1024.py → 继承 configs/fcn/fcn_r101-d8_4xb2-80k_cityscapes-512x1024.pyresnest_s101-d8_pspnet_4xb2-80k_cityscapes512x1024.py → 继承 configs/pspnet/pspnet_r101-d8_4xb2-80k_cityscapes-512x1024.pyresnest_s101-d8_deeplabv3_4xb2-80k_cityscapes-512x1024.py → 继承 configs/deeplabv3/deeplabv3_r101-d8_4xb2-80k_cityscapes-512x1024.pyresnest_s101-d8_deeplabv3plus_4xb2-80k_cityscapes-512x1024.py → 继承 configs/deeplabv3plus/deeplabv3plus_r101-d8_4xb2-80k_cityscapes-512x1024.pyADE20K 下的 4 个配置同理继承各自的_4xb4-160k_ade20k-512x512版本以 FCN Cityscapes 为例完整配置内容为_base_ ../fcn/fcn_r101-d8_4xb2-80k_cityscapes-512x1024.py model dict( pretrainedopen-mmlab://resnest101, backbonedict( typeResNeSt, stem_channels128, radix2, reduction_factor4, avg_down_strideTrue))3.2 关键配置项逐条说明配置项值说明model.pretrainedopen-mmlab://resnest101加载 OpenMMLab 托管的 ResNeSt-101 ImageNet 预训练权重backbone.typeResNeSt注册表中的骨干类型见 resnest.py 第 270 行backbone.stem_channels128输入 stem 通道数与预训练权重结构匹配backbone.radix2Split-Attention 分支基数backbone.reduction_factor4注意力瓶颈缩减因子backbone.avg_down_strideTrue用平均池化实现 stride 下采样由于ResNeSt继承自ResNetV1ddepth、dilationsd8 即输出 stride 8、out_indices等参数均沿用基类默认值配置中的d8表示采用 dilation 策略使输出步长为 8保证高分辨率分割特征。四、实验结果Cityscapes 与 ADE20K 基准以下结果来自 configs/resnest/README.md 的 Results and models 小节训练资源均为 4 块 V100 GPU。mIoU(msflip)表示多尺度 水平翻转测试的指标。4.1 Cityscapes512x102480k 迭代MethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)FCNS-101-D8512x10248000011.42.39V10077.5678.98PSPNetS-101-D8512x10248000011.82.52V10078.5779.19DeepLabV3S-101-D8512x10248000011.91.88V10079.6780.51DeepLabV3S-101-D8512x10248000013.22.36V10079.6280.274.2 ADE20K512x512160k 迭代MethodBackboneCrop SizeLr schdMem (GB)Inf time (fps)DevicemIoUmIoU(msflip)FCNS-101-D8512x51216000014.212.86V10045.6246.16PSPNetS-101-D8512x51216000014.213.02V10045.4446.28DeepLabV3S-101-D8512x51216000014.69.28V10045.7146.59DeepLabV3S-101-D8512x51216000016.211.96V10046.4747.27数据解读在 Cityscapes 上DeepLabV379.67 mIoU与 DeepLabV379.62 mIoU明显优于 FCN77.56 mIoU与 PSPNet78.57 mIoU在 ADE20K 上DeepLabV3 以 46.47 mIoU 领先且四种方法在 msflip 测试下均有约 0.5~1 个点的提升。S-101-D8 指 Split-Attention ResNeSt-101、输出 stride 8 的骨干配置。上述 8 个模型的权重与训练日志地址、批次规模Cityscapes 为 4x28ADE20K 为 4x416等元数据均登记在 configs/resnest/metafile.yaml 中可配合 MMSegmentation 的模型索引机制使用。五、实战训练、测试与推理5.1 单机多卡训练使用 tools/train.py 与仓库提供的 tools/dist_train.sh 启动训练bash tools/dist_train.sh configs/resnest/resnest_s101-d8_deeplabv3_4xb2-80k_cityscapes-512x1024.py 8其中第二个参数为 GPU 数量。训练配置的pretrainedopen-mmlab://resnest101会自动下载 ResNeSt-101 预训练权重用于 backbone 初始化。5.2 测试与指标复现使用 tools/test.py 测试--out保存预测结果--eval mIoU计算 mIoU 指标bash tools/dist_test.sh configs/resnest/resnest_s101-d8_deeplabv3_4xb2-80k_cityscapes-512x1024.py \ work_dirs/resnest_s101-d8_deeplabv3_4xb2-80k_cityscapes-512x1024/latest.pth \ 8 --out results.pkl --eval mIoU如需复现表中的mIoU(msflip)可结合 MMSegmentation 的多尺度 翻转测试评估流程tools/test.py支持的--aug-test选项。5.3 单图推理使用仓库自带的推理脚本 demo/image_demo.py 直接对单张图片进行推理python demo/image_demo.py demo/demo.png \ configs/resnest/resnest_s101-d8_deeplabv3_4xb2-80k_cityscapes-512x1024.py \ /path/to/checkpoint.pth六、单元测试验证ResNeSt 的正确性保障仓库在 tests/test_models/test_backbones/test_resnest.py 中提供了专门的单元测试覆盖两大场景1. Bottleneck 结构与前向test_resnest_bottleneck非法style参数如tensorflow会触发AssertionErrorBottleneckS(64, 256, radix2, reduction_factor4, stride2, stylepytorch)的avd_layer.stride 2验证 avg_down_stride 生效输入(2, 64, 56, 56)经 Bottleneck 后输出形状不变验证残差结构正确。2. 骨干网络整体前向test_resnest_backbone不支持的深度如depth18会抛出KeyError因为arch_settings仅支持 [50, 101, 152, 200]以ResNeSt(depth50, radix2, reduction_factor4, out_indices(0, 1, 2, 3))前向 224x224 输入四个输出阶段的特征形状分别为[2, 256, 56, 56]、[2, 512, 28, 28]、[2, 1024, 14, 14]、[2, 2048, 7, 7]——每个阶段通道数 ×2、分辨率减半验证了 Split-Attention 骨干在多尺度特征提取上的正确性。# tests/test_models/test_backbones/test_resnest.py feat model(imgs) assert feat[0].shape torch.Size([2, 256, 56, 56]) assert feat[1].shape torch.Size([2, 512, 28, 28]) assert feat[2].shape torch.Size([2, 1024, 14, 14]) assert feat[3].shape torch.Size([2, 2048, 7, 7])七、如何将 ResNeSt 扩展到自己的模型由于ResNeSt已注册进 MMSegmentation 的模型注册表任何以 ResNet 为 backbone 的分割配置都可以通过三行改动迁移到 ResNeStmodel dict( pretrainedopen-mmlab://resnest101, backbonedict( typeResNeSt, stem_channels128, radix2, reduction_factor4, avg_down_strideTrue))注意事项如果使用自己训练的 ResNeSt 权重可将pretrained替换为本地权重文件路径或通过init_cfg指定想调整 Split-Attention 的强度可修改radix分支数与reduction_factor注意力瓶颈缩减比radix1时 Split-Attention 退化为 SE 式 sigmoid 门控显存紧张时可在 backbone 中开启with_cpTruecheckpoint以时间换显存所有配置中的stem_channels128必须与预训练权重结构一致否则会因权重 shape 不匹配而加载失败。八、总结ResNeSt 通过 Split-Attention 将通道注意力和多路径表示统一为一个模块化计算块在 MMSegmentation 中作为分割骨干网络提供了从 FCN 到 DeepLabV3 的完整覆盖。本文从论文背景、RSoftmax/SplitAttentionConv2d/Bottleneck的源码实现、8 个仓库配置的复用模式、Cityscapes 与 ADE20K 的基准结果到训练测试推理的完整实战流程系统梳理了 ResNeSt 在 MMSegmentation 中的接入方式。你可以直接参考 configs/resnest/ 下的配置与 configs/resnest/metafile.yaml 中的模型清单将 ResNeSt 应用到自己的分割任务中。【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价