MMSegmentation 模型体系全解:分割器架构、核心接口与数据预处理器原理
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
本文以 MMSegmentation 的模型设计为核心,系统梳理"分割器(Segmentor)—主干网络(Backbone)—颈部(Neck)—解码头(Decode Head)—辅助头(Auxiliary Head)"的组件化架构,深入剖析forward、train_step、val_step、test_step四大核心接口的调用语义,并结合源码讲解SegDataPreProcessor数据预处理器与model.test_cfg的推理模式控制。读完本文,你将掌握 MMSegmentation 模型的配置编写方法、训练/验证/推理阶段的完整数据流,以及whole与slide两种推理模式的选取依据。
模型在 MMSegmentation 中的定位
在 MMSegmentation 中,深度学习任务中的神经网络被统一定义为模型(Model),而模型即算法的核心。MMSegmentation 基于 MMEngine 抽象出的统一模型基类BaseModel构建,BaseModel将训练、测试等过程标准化,使不同算法共享同一套生命周期管理。
MMSegmentation 实现的所有模型都继承自BaseModel,并在其基础上实现了前向传播逻辑,为语义分割算法添加了特有的功能。这种设计带来的直接好处是:无论你使用 PSPNet、DeepLabV3、UPerNet 还是最新的 Segmenter、SAN,模型在 Runner 中的训练、验证、测试流程完全一致,差异仅体现在模型组件的配置上。
常用组件:分割器的零件库
MMSegmentation 将网络架构抽象为分割器(Segmentor)——一个包含网络所有组件的模型。目前仓库实现了两种分割器:
- EncoderDecoder(编码器解码器):见 mmseg/models/segmentors/encoder_decoder.py,通常由数据预处理器、骨干网络、解码头和(可选的)辅助头组成;
- CascadeEncoderDecoder(级联编码器解码器):见 mmseg/models/segmentors/cascade_encoder_decoder.py,与前者的差异在于解码器是级联的——前一个解码头(decode_head)的输出会作为后一个解码头的输入,典型应用如 PointRend、K-Net。
两种分割器通常由以下组件拼装而成:
| 组件 | 作用 | 典型实现 |
|---|---|---|
| 数据预处理器(Data Preprocessor) | 将数据复制到目标设备,并预处理为模型输入格式 | SegDataPreProcessor |
| 主干网络(Backbone) | 将图像转换为特征图 | 去掉最后全连接层的 ResNet-50 |
| 颈部(Neck) | 连接主干网络与头,对原始特征图做改进或重新配置 | Feature Pyramid Network(FPN) |
| 解码头(Decode Head) | 将特征图转换为分割掩膜 | PSPNet 的PSPHead、DeepLabV3 的ASPPHead |
| 辅助头(Auxiliary Head) | 可选组件,仅用于计算辅助损失的分割掩膜,推理时可丢弃 | FCNHead |
关于辅助头,源码中有明确注释:辅助头只用于训练期间的深度监督(deep supervision),推理阶段会被丢弃(见 encoder_decoder.py)。在配置中,它对应model.auxiliary_head字段,通过loss_weight控制辅助损失在总损失中的权重。
从源码结构看,BaseSegmentor(mmseg/models/segmentors/base.py)还提供了三个便捷属性,用于判断分割器是否包含对应组件:
with_neck:是否配置了颈部;with_auxiliary_head:是否配置了辅助头;with_decode_head:是否配置了解码头(EncoderDecoder构造时断言必须有解码头)。
基本接口:forward / train_step / val_step / test_step
MMSegmentation 封装BaseModel并实现了BaseSegmentor类,主要提供forward、train_step、val_step和test_step四个接口。其中train_step、val_step、test_step由BaseModel定义标准流程,forward是自定义的核心前向入口。
forward:统一的前向入口
forward方法是训练、验证、测试和简单推理的统一前向入口,返回损失或预测结果。它必须支持三种模式(base.py):
- "tensor":前向推理整个网络并返回张量或张量数组,不做任何后处理,行为与常见
nn.Module一致; - "predict":前向推理并返回预测值,预测结果会被完整后处理为
SegDataSample列表; - "loss":前向推理并根据给定输入和数据样本返回损失的字典。
若传入不支持的模式,forward会抛出RuntimeError(仅支持 loss、predict、tensor 三种模式)。
注意:forward方法不处理反向传播与优化器更新,这两者在train_step方法中完成。
参数说明:
inputs(torch.Tensor):输入张量,通常形状为 (N, C, ...);data_sample(list[SegDataSample]):分割数据样本,通常包含metainfo和gt_sem_seg等信息,默认为 None;mode(str):决定返回值类型,默认为'tensor'。
返回值说明:
- 若
mode == "loss",返回用于反向过程和日志记录的损失张量字典; - 若
mode == "predict",返回SegDataSample的列表,推理结果会被递增地添加到传入的data_sample参数中。每个SegDataSample包含以下关键词:pred_sem_seg(PixelData):语义分割的预测结果;seg_logits(PixelData):标准化前语义分割的预测 logits;
- 若
mode == "tensor",返回张量或张量数组的字典,供自定义使用。
SegDataSample是 MMSegmentation 的数据结构接口,实现自mmengine.structures.BaseDataElement(见 mmseg/structures/seg_data_sample.py),用作不同组件之间的接口。从源码看,它对外暴露gt_sem_seg、pred_sem_seg、seg_logits三个PixelData类型的属性字段。
预测模式:whole_inference 与 slide_inference
模型配置的字段在配置文档中有简要描述,这里重点展开model.test_cfg字段。model.test_cfg用于控制前向行为,"predict"模式下的forward方法可以在两种模式下运行(实现见 encoder_decoder.py):
whole_inference(整图推理):当
cfg.model.test_cfg.mode == 'whole'时,模型使用完整图像进行推理,EncoderDecoder.whole_inference直接对整图调用encode_decode得到 seg_logits。配置示例:model = dict( type='EncoderDecoder' ... test_cfg=dict(mode='whole') )这是绝大多数配置的默认选择。例如 deeplabv3_r50-d8_4xb2-40k_cityscapes-512x1024.py 继承的基础模型中即写有
test_cfg=dict(mode='whole')。slide_inference(滑动窗口推理):当
cfg.model.test_cfg.mode == 'slide'时,模型通过滑动窗口进行推理。注意:选择slide模式时,还必须指定cfg.model.test_cfg.stride和cfg.model.test_cfg.crop_size。配置示例:model = dict( type='EncoderDecoder' ... test_cfg=dict(mode='slide', crop_size=256, stride=170) )
从slide_inference的实现(encoder_decoder.py)可以看到其工作原理:按stride在图像上划出h_grids × w_grids个重叠窗口,逐窗口调用encode_decode得到局部 seg_logits,通过F.pad将每个窗口的 logits 累积到整图坐标上,并用count_mat记录每个像素被覆盖的次数,最后以preds / count_mat取平均,从而消除窗口重叠区域的边界伪影。
一个真实的 slide 配置示例是 pspnet_r50-d8_4xb2-40k_cityscapes-769x769.py:
crop_size = (769, 769) data_preprocessor = dict(size=crop_size) model = dict( data_preprocessor=data_preprocessor, decode_head=dict(align_corners=True), auxiliary_head=dict(align_corners=True), test_cfg=dict(mode='slide', crop_size=(769, 769), stride=(513, 513)))这里crop_size=(769, 769)与stride=(513, 513)的搭配使相邻窗口有约 1/3 的重叠,兼顾了推理质量与计算量。同时,滑窗裁剪得到的局部 patch 会通过predict_by_feat中的img_shape判断(见 decode_head.py)被双线性插值回对应尺寸。
train_step:训练数据流
train_step方法调用loss模式的前向接口以获得损失字典。BaseModel类实现了默认的模型训练过程,包括预处理、模型前向传播、损失计算、优化和反向传播。
参数:
data(dict or tuple or list):从数据集采样的数据。在 MMSegmentation 中,数据字典包含inputs和data_samples两个字段;optim_wrapper(OptimWrapper):用于更新模型参数的 OptimWrapper 实例。OptimWrapper提供了更新参数的通用接口,统一了 PyTorch 优化器的使用方式。
返回值:
- Dict[str, torch.Tensor]:用于记录日志的张量字典。
以EncoderDecoder为例,其loss方法的调用链(见 encoder_decoder.py)为:
loss(): extract_feat() -> _decode_head_forward_train() -> _auxiliary_head_forward_train()(可选) _decode_head_forward_train(): decode_head.loss() _auxiliary_head_forward_train(): auxiliary_head.loss()(可选)其中extract_feat依次执行backbone(inputs)与(若存在)neck(x)(encoder_decoder.py);解码头/辅助头各自的loss方法在 decode_head.py 中实现为forward() -> loss_by_feat()两步,loss_by_feat会先按gt_sem_seg尺寸 resize logits,再计算损失与acc_seg像素准确率。所有子模块的损失通过add_prefix加上decode、aux(级联时为decode_0、decode_1…)前缀后合并进同一个损失字典。
val_step 与 test_step:验证/测试数据流
val_step方法调用predict模式的前向接口并返回预测结果,预测结果将进一步被传递给评测器的进程接口和钩子的after_val_inter接口。
参数:
data(dict or tuple or list):从数据集中采样的数据,数据字典同样包含inputs和data_samples两个字段。
返回值:
list:给定数据的预测结果。
BaseModel中test_step与val_step的实现相同,因此二者的数据流完全一致。
以EncoderDecoder为例,其predict方法的调用链(见 encoder_decoder.py)为:
predict(): inference() -> postprocess_result() inference(): whole_inference()/slide_inference() whole_inference()/slide_inference(): encode_decode() encode_decode(): extract_feat() -> decode_head.predict()推理得到的 seg_logits 会交给postprocess_result(base.py)做最终后处理:根据metainfo中的padding_size裁剪掉 padding 区域、按flip与flip_direction还原翻转、用双线性插值 resize 回ori_shape原始尺寸;当类别数 C > 1 时用argmax得到pred_sem_seg,当 C == 1(二分类)时用sigmoid配合decode_head.threshold阈值生成二值掩膜,最后将seg_logits与pred_sem_seg写入SegDataSample返回。
数据预处理器(SegDataPreProcessor)详解
MMSegmentation 实现的SegDataPreProcessor(mmseg/models/data_preprocessor.py)继承自 MMEngine 的BaseDataPreprocessor,提供数据预处理和将数据复制到目标设备的功能。
设备迁移时机:Runner 在构建阶段将模型传送到指定设备,而SegDataPreProcessor在train_step、val_step和test_step中将数据传送到指定设备,之后处理后的数据才被进一步传递给模型。
构造函数参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
mean | Sequence[Number] | None | R、G、B 通道的像素平均值 |
std | Sequence[Number] | None | R、G、B 通道的像素标准差 |
size | tuple | None | 固定的填充大小 |
size_divisor | int | None | 填充尺寸的除法因子(填充后尺寸为 divisor 的整数倍) |
pad_val | float | 0 | 图像填充值 |
seg_pad_val | float | 255 | 分割图的填充值(255 在语义分割中约定为 ignore index) |
bgr_to_rgb | bool | False | 是否将图像从 BGR 转换为 RGB |
rgb_to_bgr | bool | False | 是否将图像从 RGB 转换为 BGR |
batch_augments | list[dict] | None | 批量级数据增强配置 |
从源码实现看,有几个值得注意的行为(data_preprocessor.py):
bgr_to_rgb与rgb_to_bgr互斥:二者不能同时为 True,否则触发断言;- 归一化可选:仅当同时指定
mean和std时才启用归一化(_enable_normalize = True),若只给mean不给std会直接断言报错;归一化在堆叠成 batch 之后进行; - 额外的
test_cfg参数:支持在测试阶段单独指定size或size_divisor来控制 padding 方式。
数据处理流程
数据按如下方式处理(与源码 docstring 一致):
- 收集数据并将其移动到目标设备(
cast_data); - 用定义的
pad_val将输入填充到目标尺寸,并用定义的seg_pad_val填充分割图(由stack_batch完成); - 将输入堆叠为
batch_inputs; - 如果输入形状为 (3, H, W),则将输入从 BGR 转换为 RGB(
channel_conversion开启时); - 使用定义的
std和mean标准化图像; - 在训练期间进行 Mixup、Cutmix 等批量级数据增强(
batch_augments)。
forward 方法
参数:
data(dict):从数据加载器采样的数据;training(bool):是否启用训练时数据增强。
返回值:
- Dict:与模型输入格式相同的数据。
训练与测试分支的行为不同(data_preprocessor.py):训练时必须存在data_samples,经过stack_batch后若配置了batch_augments再执行批量增强;测试时则要求 batch 内图像尺寸一致,若配置了test_cfg则按其中size/size_divisor做 padding,并把 padding 信息通过set_metainfo写回data_samples(这正是postprocess_result裁剪 padding 区域所需的信息来源),否则直接torch.stack成 batch。
配置示例
在真实配置中,数据预处理器通常写在模型基础配置里。以 deeplabv3_r50-d8.py 为例:
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='ASPPHead', in_channels=2048, in_index=3, channels=512, dilations=(1, 12, 24, 36), 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'))该配置同时演示了:data_preprocessor使用 ImageNet 统计的 mean/std 并开启 BGR→RGB 转换(bgr_to_rgb=True);decode_head使用ASPPHead并以loss_weight=1.0的交叉熵作为主损失;auxiliary_head使用FCNHead且loss_weight=0.4作为辅助损失;test_cfg默认mode='whole'。若需要固定输入尺寸(如 Cityscapes 的 512×1024),可在具体实验配置中覆盖data_preprocessor:
crop_size = (512, 1024) data_preprocessor = dict(size=crop_size) model = dict(data_preprocessor=data_preprocessor)总结
MMSegmentation 的模型设计遵循"一切皆组件、一切皆可配置"的原则:BaseModel提供统一生命周期,BaseSegmentor抽象分割器接口,EncoderDecoder/CascadeEncoderDecoder提供两种可组合的架构范式,SegDataPreProcessor屏蔽了设备迁移、padding、归一化与批量增强等重复劳动。理解forward的三种模式与test_cfg的推理模式,是正确编写模型配置、排查推理性能问题的关键。后续可进一步阅读配置文档了解完整的字段体系,或在 configs 目录中对照各算法的真实配置加深理解。
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考