MMSegmentation 模型体系全解:分割器架构、核心接口与数据预处理器原理
2026/9/16 10:19:23 网站建设 项目流程

MMSegmentation 模型体系全解:分割器架构、核心接口与数据预处理器原理

【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation

本文以 MMSegmentation 的模型设计为核心,系统梳理"分割器(Segmentor)—主干网络(Backbone)—颈部(Neck)—解码头(Decode Head)—辅助头(Auxiliary Head)"的组件化架构,深入剖析forwardtrain_stepval_steptest_step四大核心接口的调用语义,并结合源码讲解SegDataPreProcessor数据预处理器与model.test_cfg的推理模式控制。读完本文,你将掌握 MMSegmentation 模型的配置编写方法、训练/验证/推理阶段的完整数据流,以及wholeslide两种推理模式的选取依据。

模型在 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类,主要提供forwardtrain_stepval_steptest_step四个接口。其中train_stepval_steptest_stepBaseModel定义标准流程,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]):分割数据样本,通常包含metainfogt_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_segpred_sem_segseg_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.stridecfg.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 中,数据字典包含inputsdata_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加上decodeaux(级联时为decode_0decode_1…)前缀后合并进同一个损失字典。

val_step 与 test_step:验证/测试数据流

val_step方法调用predict模式的前向接口并返回预测结果,预测结果将进一步被传递给评测器的进程接口和钩子的after_val_inter接口。

参数:

  • data(dict or tuple or list):从数据集中采样的数据,数据字典同样包含inputsdata_samples两个字段。

返回值:

  • list:给定数据的预测结果。

BaseModeltest_stepval_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 区域、按flipflip_direction还原翻转、用双线性插值 resize 回ori_shape原始尺寸;当类别数 C > 1 时用argmax得到pred_sem_seg,当 C == 1(二分类)时用sigmoid配合decode_head.threshold阈值生成二值掩膜,最后将seg_logitspred_sem_seg写入SegDataSample返回。

数据预处理器(SegDataPreProcessor)详解

MMSegmentation 实现的SegDataPreProcessor(mmseg/models/data_preprocessor.py)继承自 MMEngine 的BaseDataPreprocessor,提供数据预处理和将数据复制到目标设备的功能。

设备迁移时机:Runner 在构建阶段将模型传送到指定设备,而SegDataPreProcessortrain_stepval_steptest_step中将数据传送到指定设备,之后处理后的数据才被进一步传递给模型。

构造函数参数

参数类型默认值说明
meanSequence[Number]NoneR、G、B 通道的像素平均值
stdSequence[Number]NoneR、G、B 通道的像素标准差
sizetupleNone固定的填充大小
size_divisorintNone填充尺寸的除法因子(填充后尺寸为 divisor 的整数倍)
pad_valfloat0图像填充值
seg_pad_valfloat255分割图的填充值(255 在语义分割中约定为 ignore index)
bgr_to_rgbboolFalse是否将图像从 BGR 转换为 RGB
rgb_to_bgrboolFalse是否将图像从 RGB 转换为 BGR
batch_augmentslist[dict]None批量级数据增强配置

从源码实现看,有几个值得注意的行为(data_preprocessor.py):

  • bgr_to_rgbrgb_to_bgr互斥:二者不能同时为 True,否则触发断言;
  • 归一化可选:仅当同时指定meanstd时才启用归一化(_enable_normalize = True),若只给mean不给std会直接断言报错;归一化在堆叠成 batch 之后进行;
  • 额外的test_cfg参数:支持在测试阶段单独指定sizesize_divisor来控制 padding 方式。

数据处理流程

数据按如下方式处理(与源码 docstring 一致):

  1. 收集数据并将其移动到目标设备(cast_data);
  2. 用定义的pad_val将输入填充到目标尺寸,并用定义的seg_pad_val填充分割图(由stack_batch完成);
  3. 将输入堆叠为batch_inputs
  4. 如果输入形状为 (3, H, W),则将输入从 BGR 转换为 RGB(channel_conversion开启时);
  5. 使用定义的stdmean标准化图像;
  6. 在训练期间进行 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使用FCNHeadloss_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询