MMSegmentation 中的 CCNet:十字交叉注意力语义分割的实现原理、配置解析与训练实践
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
导读
CCNet(Criss-Cross Network)是一种以十字交叉注意力(Criss-Cross Attention)为核心的语义分割网络,它以极低的显存与计算开销获取全图上下文信息,是 OpenMMLab 语义分割工具箱 MMSegmentation 中官方收录并长期维护的经典算法之一。本文以 configs/ccnet/README.md 为骨架,结合仓库中CCHead的实现源码与configs/ccnet/下的全套训练配置,系统讲解 CCNet 的算法原理、MMSegmentation 中的代码实现、配置文件各字段含义,以及基于 Cityscapes、ADE20K、Pascal VOC 的复现结果与训练、测试命令,帮助你快速上手并在自己的数据集上复现或迁移这一算法。
CCNet 算法核心:以十字交叉注意力高效建模全图上下文
问题动机:上下文信息对语义分割至关重要
上下文信息(Contextual Information)在语义分割、目标检测等视觉理解任务中起着决定性作用。一个像素的分类往往不能只看自身局部的颜色与纹理,还需要借助其周围乃至全图的语义线索。早期方法中,non-local 类模块虽然能建模全图依赖,但其注意力图在空间维度上是 O(HW × HW) 的规模,显存占用与计算量都随特征图尺寸急剧膨胀,难以在分割任务常见的大分辨率输入上使用。
CCNet 提出了一种全新的思路:不直接计算全图任意两像素之间的两两关系,而是先让每个像素只收集其所在"十字路径"(同一行与同一列)上的上下文信息,再通过一次循环(recurrence)操作,让信息沿十字路径在图中传播,最终使每个像素都能捕获到全图范围的依赖关系。
方法要点:循环十字交叉注意力 + 类别一致损失
从 configs/ccnet/README.md 的 Abstract 可以归纳出 CCNet 的三大核心贡献:
- 十字交叉注意力模块(Criss-Cross Attention Module):对每个像素,该模块沿其所在行与列构成的十字路径聚合上下文,单次操作的计算规模从全图 O(HW × HW) 降为 O(HW × (H+W)),显存友好。
- 循环操作(Recurrent Operation):将十字交叉注意力模块串联执行两次(
recurrence=2),信息即可沿十字路径传播到全图,每个像素最终捕获完整图像依赖,而不需要显式地计算全图两两相似度。 - 类别一致损失(Category Consistent Loss):进一步约束十字交叉注意力模块产生更具判别性的特征。
论文报告了两项关键效率收益(见 README Abstract,属论文声称数据):相比 non-local block,循环十字交叉注意力模块的显存占用降低约 11 倍,FLOPs 减少约 85%。这正是 CCNet 被广泛用于高分辨率分割任务的核心原因。
效率与效果的平衡
在 MMSegmentation 收录的配置中,CCNet 的解码头直接在 ResNet 主干输出的高分辨率特征(如 Cityscapes 的 512×1024、769×769)上工作,正是因为十字交叉注意力的线性复杂度(相对特征图边长)使其能够承受大分辨率输入。这一点可以从下方结果表中显存占用(6~12.2 GB)与推理速度(1.01~20.89 fps,V100)中直观看出。
MMSegmentation 中的源码实现:CCHead 与 CrissCrossAttention
CCHead:一个继承了 FCNHead 的解码头
在 MMSegmentation 中,CCNet 的解码头实现在 mmseg/models/decode_heads/cc_head.py,并通过@MODELS.register_module()注册为CCHead。其类结构要点如下:
@MODELS.register_module() class CCHead(FCNHead): """CCNet: Criss-Cross Attention for Semantic Segmentation...""" def __init__(self, recurrence=2, **kwargs): if CrissCrossAttention is None: raise RuntimeError('Please install mmcv-full for ' 'CrissCrossAttention ops') super().__init__(num_convs=2, **kwargs) self.recurrence = recurrence self.cca = CrissCrossAttention(self.channels) def forward(self, inputs): x = self._transform_inputs(inputs) output = self.convs0 for _ in range(self.recurrence): output = self.cca(output) output = self.convs1 if self.concat_input: output = self.conv_cat(torch.cat([x, output], dim=1)) output = self.cls_seg(output) return output从源码可以确认以下实现事实:
- CCHead 继承自 FCNHead(见 mmseg/models/decode_heads/fcn_head.py),其内部卷积结构由 FCNHead 以
num_convs=2构建:先通过第一个卷积把主干特征从in_channels(如 2048)降到channels(如 512),再在两层卷积之间循环插入recurrence次十字交叉注意力模块,最后经cls_seg逐像素分类。 - 核心算子来自 mmcv.ops:
CrissCrossAttention由mmcv.ops提供(from mmcv.ops import CrissCrossAttention),并以self.cca = CrissCrossAttention(self.channels)实例化。若环境中缺少带 CUDA 算子的 mmcv,CCHead 会直接抛出RuntimeError('Please install mmcv-full for CrissCrossAttention ops'),这是该模块的硬性安装前提。 recurrence是唯一新增参数,默认值为 2,对应论文中的"一次循环即可捕获全图依赖"的设计,在配置文件中可以直接覆盖。
测试验证
仓库中的单元测试 tests/test_models/test_heads/test_cc_head.py 验证了 CCHead 的结构与前向过程:
def test_cc_head(): head = CCHead(in_channels=16, channels=8, num_classes=19) assert len(head.convs) == 2 assert hasattr(head, 'cca') if not torch.cuda.is_available(): pytest.skip('CCHead requires CUDA') inputs = [torch.randn(1, 16, 23, 23)] head, inputs = to_cuda(head, inputs) outputs = head(inputs)同时,tests/test_models/test_forward.py 中的test_ccnet_forward使用ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py整网前向测试验证配置可加载。注意两个测试都通过pytest.skip('CCNet requires CUDA')跳过无 GPU 环境,从侧面印证:十字交叉注意力算子是 CUDA 实现,训练与推理都需要 GPU 环境。
配置解析:从基础模型到完整训练方案
configs/ccnet/目录共收录 16 个官方训练配置(对应 R-50/R-101 主干、Cityscapes/ADE20K/Pascal VOC 数据集与 20k/40k/80k/160k 迭代数),全部采用_base_继承机制组合而成。下面以最常用的 Cityscapes 配置为例逐层拆解。
顶层训练配置:ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py
完整文件内容见 configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py:
_base_ = [ '../_base_/models/ccnet_r50-d8.py', '../_base_/datasets/cityscapes.py', '../_base_/default_runtime.py', '../_base_/schedules/schedule_40k.py' ] crop_size = (512, 1024) data_preprocessor = dict(size=crop_size) model = dict(data_preprocessor=data_preprocessor)它通过四个_base_分别继承模型结构、数据集配置、运行时配置、训练调度四个维度,随后仅覆盖crop_size与data_preprocessor.size,即可让数据预处理与模型输入尺寸对齐到 512×1024。文件名中各段的含义为:
r50-d8:ResNet-50 主干,dilated 策略为 8(dilations=(1, 1, 2, 4),即 stage3/4 使用膨胀卷积,最终输出 stride 为 8);4xb2:4 张 GPU、每卡 batch size 2(总 batch size 8,见 configs/ccnet/metafile.yaml 的Batch Size: 8);40k:训练 40000 次迭代;512x1024:裁剪尺寸。
模型基础配置:ccnet_r50-d8.py
模型骨架定义在 configs/base/models/ccnet_r50-d8.py,它由主干、解码头、辅助头三部分组成:
norm_cfg = dict(type='SyncBN', requires_grad=True) data_preprocessor = dict( type='SegDataPreProcessor', mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], bgr_to_rgb=True, pad_val=0, seg_pad_val=255) model = dict( type='EncoderDecoder', data_preprocessor=data_preprocessor, pretrained='open-mmlab://resnet50_v1c', backbone=dict( type='ResNetV1c', depth=50, num_stages=4, out_indices=(0, 1, 2, 3), dilations=(1, 1, 2, 4), strides=(1, 2, 1, 1), norm_cfg=norm_cfg, norm_eval=False, style='pytorch', contract_dilation=True), decode_head=dict( type='CCHead', in_channels=2048, in_index=3, channels=512, recurrence=2, dropout_ratio=0.1, num_classes=19, norm_cfg=norm_cfg, align_corners=False, loss_decode=dict( type='CrossEntropyLoss', use_sigmoid=False, loss_weight=1.0)), auxiliary_head=dict( type='FCNHead', in_channels=1024, in_index=2, channels=256, num_convs=1, concat_input=False, dropout_ratio=0.1, num_classes=19, norm_cfg=norm_cfg, align_corners=False, loss_decode=dict( type='CrossEntropyLoss', use_sigmoid=False, loss_weight=0.4)), train_cfg=dict(), test_cfg=dict(mode='whole'))各关键字段的语义与取值如下:
| 配置字段 | 取值 | 含义 |
|---|---|---|
type='CCHead' | — | 使用上面分析的十字交叉注意力解码头 |
in_channels=2048 | — | 主干 stage4 输出通道数(ResNet-50 最后一层) |
in_index=3 | — | 取主干out_indices的第 3 层(stride 8 特征)作为解码头输入 |
channels=512 | — | 注意力模块内部的隐藏通道数,也是CrissCrossAttention的输入通道 |
recurrence=2 | 默认 2 | 十字交叉注意力循环次数,对应论文"循环一次即可覆盖全图" |
dropout_ratio=0.1 | — | 分类层前的 dropout 比例 |
num_classes=19 | Cityscapes 为 19 | 类别数,迁移到 ADE20K 时改为 150 |
auxiliary_head | FCNHead,loss_weight=0.4 | 在主干 stage3(in_index=2,1024 通道)上挂载辅助 FCN 头,辅助损失权重 0.4,主损失权重 1.0 |
test_cfg=dict(mode='whole') | — | 整图推理,不切块 |
pretrained='open-mmlab://resnet50_v1c' | — | 使用 ResNetV1c 的 ImageNet 预训练权重 |
值得注意的工程细节:
- 数据预处理:
SegDataPreProcessor采用 ImageNet 均值/方差(mean=[123.675, 116.28, 103.53]、std=[58.395, 57.12, 57.375])做归一化,且bgr_to_rgb=True表示输入按 BGR 读取后转为 RGB,这些是 MMSegmentation 1.x 统一的预处理约定。 - 膨胀策略 d8:
dilations=(1, 1, 2, 4)、strides=(1, 2, 1, 1)表示 stage3 步长改为 1 并膨胀 2、stage4 膨胀 4,使最终特征图保持输入尺寸的 1/8,从而保留密集预测所需的空间分辨率。
其他数据集的配置差异
以 ADE20K 配置 configs/ccnet/ccnet_r50-d8_4xb4-80k_ade20k-512x512.py 为例,它与 Cityscapes 配置的差异非常直观:
_base_ = [ '../_base_/models/ccnet_r50-d8.py', '../_base_/datasets/ade20k.py', '../_base_/default_runtime.py', '../_base_/schedules/schedule_80k.py' ] crop_size = (512, 512) data_preprocessor = dict(size=crop_size) model = dict( data_preprocessor=data_preprocessor, decode_head=dict(num_classes=150), auxiliary_head=dict(num_classes=150))仅需三处修改即可从 Cityscapes 迁移到 ADE20K:换用ade20k.py数据集配置与schedule_80k.py调度、裁剪尺寸改为 512×512、num_classes改为 150(解码头与辅助头同步修改)。这体现了 MMSegmentation 配置继承体系在算法迁移上的低成本优势。
训练调度:schedule_40k.py
configs/base/schedules/schedule_40k.py 定义了 CCNet 在 Cityscapes 上的训练超参:
- 优化器:SGD,
lr=0.01,momentum=0.9,weight_decay=0.0005; - 学习率策略:
PolyLR(多项式衰减),power=0.9,eta_min=1e-4,按迭代(by_epoch=False)从 40000 衰减; - 训练循环:
IterBasedTrainLoop,max_iters=40000,每 4000 迭代验证一次并保存一次 checkpoint; - 日志钩子:每 50 次迭代输出一次日志(
LoggerHook)。
官方复现结果与模型
configs/ccnet/README.md记录了 CCNet 在三个主流分割基准上的完整复现结果(指标由官方训练日志统计,测试设备为 V100;"ms+flip" 表示多尺度 + 翻转测试)。以下结果表中的配置文件链接均已转换为仓库内相对路径,权重与训练日志的下载地址收录于 configs/ccnet/metafile.yaml,可通过mim search mmsegmentation --model ccnet或模型下载工具按Name字段获取。
Cityscapes(19 类,训练 40000 / 80000 迭代)
| Method | Backbone | Crop Size | Lr schd | Mem (GB) | Inf time (fps) | Device | mIoU | mIoU(ms+flip) | config |
|---|---|---|---|---|---|---|---|---|---|
| CCNet | R-50-D8 | 512x1024 | 40000 | 6.0 | 3.32 | V100 | 77.76 | 78.87 | config |
| CCNet | R-101-D8 | 512x1024 | 40000 | 9.5 | 2.31 | V100 | 76.35 | 78.19 | config |
| CCNet | R-50-D8 | 769x769 | 40000 | 6.8 | 1.43 | V100 | 78.46 | 79.93 | config |
| CCNet | R-101-D8 | 769x769 | 40000 | 10.7 | 1.01 | V100 | 76.94 | 78.62 | config |
| CCNet | R-50-D8 | 512x1024 | 80000 | - | - | V100 | 79.03 | 80.16 | config |
| CCNet | R-101-D8 | 512x1024 | 80000 | - | - | V100 | 78.87 | 79.90 | config |
| CCNet | R-50-D8 | 769x769 | 80000 | - | - | V100 | 79.29 | 81.08 | config |
| CCNet | R-101-D8 | 769x769 | 80000 | - | - | V100 | 79.45 | 80.66 | config |
ADE20K(150 类,训练 80000 / 160000 迭代)
| Method | Backbone | Crop Size | Lr schd | Mem (GB) | Inf time (fps) | Device | mIoU | mIoU(ms+flip) | config |
|---|---|---|---|---|---|---|---|---|---|
| CCNet | R-50-D8 | 512x512 | 80000 | 8.8 | 20.89 | V100 | 41.78 | 42.98 | config |
| CCNet | R-101-D8 | 512x512 | 80000 | 12.2 | 14.11 | V100 | 43.97 | 45.13 | config |
| CCNet | R-50-D8 | 512x512 | 160000 | - | - | V100 | 42.08 | 43.13 | config |
| CCNet | R-101-D8 | 512x512 | 160000 | - | - | V100 | 43.71 | 45.04 | config |
Pascal VOC 2012 + Aug(21 类,训练 20000 / 40000 迭代)
| Method | Backbone | Crop Size | Lr schd | Mem (GB) | Inf time (fps) | Device | mIoU | mIoU(ms+flip) | config |
|---|---|---|---|---|---|---|---|---|---|
| CCNet | R-50-D8 | 512x512 | 20000 | 6.0 | 20.45 | V100 | 76.17 | 77.51 | config |
| CCNet | R-101-D8 | 512x512 | 20000 | 9.5 | 13.64 | V100 | 77.27 | 79.02 | config |
| CCNet | R-50-D8 | 512x512 | 40000 | - | - | V100 | 75.96 | 77.04 | config |
| CCNet | R-101-D8 | 512x512 | 40000 | - | - | V100 | 77.87 | 78.90 | config |
从结果可以归纳出两条可复用的工程经验:
- 长训练更有利:Cityscapes 上 40k → 80k 迭代,R-50-D8 在 512×1024 下 mIoU 从 77.76 提升到 79.03,R-101-D8 在 769×769 下更达到 79.45(ms+flip 81.08),说明 CCNet 在大迭代数下能更充分地发挥循环注意力的建模能力;
- 大裁剪尺寸收益明显:Cityscapes 上 769×769 普遍优于 512×1024(如 R-50-D8 80k 时 79.29 vs 79.03),代价是显存与推理时间上升(40000 迭代下显存 6.0 → 6.8 GB,推理 3.32 → 1.43 fps)。
训练与测试实操
在完成数据准备(Cityscapes、ADE20K、Pascal VOC 的目录组织方式参见 docs/zh_cn/user_guides/2_dataset_prepare.md 或英文版 docs/en/user_guides/2_dataset_prepare.md)并确保环境安装了包含CrissCrossAttentionCUDA 算子的 mmcv 后,即可直接使用仓库提供的脚本训练与评测。
单机多卡训练
bash tools/dist_train.sh configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py 4单卡训练
python tools/train.py configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py测试与精度复现
# 单卡测试,需将 <checkpoint> 替换为 metafile.yaml 中对应模型的权重路径 python tools/test.py configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py <checkpoint> --eval mIoU # 多卡测试 bash tools/dist_test.sh configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py <checkpoint> 4 --eval mIoU若需要复现表中的mIoU(ms+flip),可在测试命令中追加--aug-test开启多尺度 + 翻转测试;tools/test.py与tools/train.py均支持--cfg-options xxx=yyy覆盖配置字段(详见 docs/en/user_guides/1_config.md),例如临时修改类别数或学习率无需改动原配置文件。训练入口tools/train.py会按schedule_40k.py中的CheckpointHook(每 4000 迭代)自动保存权重与日志,便于中途接续训练与监控曲线。
迁移到自定义数据集
基于上述配置继承机制,将 CCNet 迁移到自己的数据集只需三步:
- 准备数据集:按 MMSegmentation 的标准目录结构组织图像与标注,参考 docs/zh_cn/user_guides/2_dataset_prepare.md 或直接仿照 configs/base/datasets/cityscapes.py 新建数据集配置,修改
data_root、metainfo.classes与palette; - 新建模型配置:复制 configs/ccnet/ccnet_r50-d8_4xb2-40k_cityscapes-512x1024.py,将
_base_中的数据集配置替换为自己的数据集,并在decode_head、auxiliary_head中把num_classes改为自己的类别数,同时调整crop_size与data_preprocessor.size; - 调整训练调度:根据数据规模修改
_base_/schedules/中的max_iters、val_interval与学习率,或在命令行用--cfg-options覆盖。
引用
若在论文或项目中使用了 CCNet 或本文介绍的复现配置,可引用以下文献(源自 configs/ccnet/README.md 的 Citation 部分):
@article{huang2018ccnet, title={CCNet: Criss-Cross Attention for Semantic Segmentation}, author={Huang, Zilong and Wang, Xinggang and Huang, Lichao and Huang, Chang and Wei, Yunchao and Liu, Wenyu}, booktitle={ICCV}, year={2019} }小结
CCNet 通过"十字交叉注意力 + 循环传播"这一精巧设计,以接近线性的显存与算力开销逼近全图上下文建模效果,是 MMSegmentation 中兼顾效率与精度的高性价比解码头之一。本文从 configs/ccnet/README.md 出发,结合 mmseg/models/decode_heads/cc_head.py 的实现、configs/base/models/ccnet_r50-d8.py 的配置细节与官方复现结果,完整覆盖了从原理、源码到训练评测的全链路。需要进一步探索时,可以继续阅读仓库的模型配置全集 configs/ccnet/、元信息文件 configs/ccnet/metafile.yaml,以及单测 tests/test_models/test_heads/test_cc_head.py 验证你对模块行为的理解。
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考