MMSegmentation 数据流全解析:从 DataLoader 到损失回传的格式约定与源码实现
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
本篇文章围绕 MMSegmentation(OpenMMLab 语义分割工具箱)中由 MMEngine Runner 调度的完整数据流展开,逐段剖析数据加载器、数据预处理器、模型前向与损失计算之间的数据格式约定。读完本文,你将掌握PackSegInputs打包规则、SegDataSample结构、SegDataPreProcessor批处理流程、模型三种前向模式及postprocess_result后处理细节,并能在训练与推理链路中准确追踪每一份数据的形态变化。
数据流概述:Runner 如何串联整个训练与评测管线
在 MMEngine 的设计中,Runner 相当于整个框架的"集成器":它覆盖了框架的几乎所有方面,肩负着组织与调度各模块的责任。因此,模块之间的数据流也由 Runner 统一控制——训练循环(TrainLoop)、验证循环(ValLoop)与测试循环(TestLoop)均由 Runner 在合适的时机启动,并在每次迭代中驱动「数据加载 → 预处理 → 模型前向 → 优化/评估」这条主链路。
MMSegmentation 对 loop 的默认设置是:使用IterBasedTrainLoop按迭代数训练模型,默认共 20000 次迭代,并且每 2000 次迭代后执行一次验证。对应配置如下:
train_cfg = dict(type='IterBasedTrainLoop', max_iters=20000, val_interval=2000) val_cfg = dict(type='ValLoop') test_cfg = dict(type='TestLoop')需要说明的是,上文描述的数据流适用于「用户没有自定义 Runner 中的TrainLoop、ValLoop、TestLoop,且没有在自定义模型中覆写train_step、val_step、test_step方法」的默认场景。由于 MMEngine 与 MMSegmentation 具有极高的灵活性和可扩展性,这些基类方法均可以被继承与覆写,从而自定义数据流向。
整条数据流可以用下图概括(虚线框表示数据格式,实线框表示模块或方法):
- 红色主线(train_step):每次训练迭代中,数据加载器从存储中加载图像并传给数据预处理器;预处理器将图像放到指定设备、把数据堆叠成批(batch);模型接受批处理数据作为输入,最后把输出交给优化器(optimizer)完成权重更新。
- 蓝色主线(val_step / test_step):流程与
train_step基本一致,区别仅在于模型输出不同——评估时模型参数被冻结,模型的输出会被传递给 Evaluator 来计算指标(如 mIoU)。
数据加载器到数据预处理器:PackSegInputs 与 SegDataSample
DataLoader 与 PackSegInputs 的分工
数据加载器(DataLoader)是 MMEngine 训练与测试流程中的重要组件,它源自 PyTorch 并保持一致的语义:从文件系统加载数据,原始数据经过数据准备流程(pipeline)后发送给数据预处理器。
MMSegmentation 在 PackSegInputs 中定义了默认的数据格式,它是train_pipeline和test_pipeline的最后一个组件。有关数据转换 pipeline 的更多信息,可参阅数据转换文档。
在没有任何修改的情况下,PackSegInputs.transform的返回值是一个包含inputs和data_samples的字典。以下伪代码展示了 mmseg 中数据加载器输出的数据类型——它是从数据集中取回的一批数据样本,数据加载器将它们打包成字典;inputs是输入进模型的张量列表,data_samples则包含输入图像的 meta 信息和对应的 ground truth:
dict( inputs=List[torch.Tensor], data_samples=List[SegDataSample] )PackSegInputs 的打包细节
从 formatting.py 的transform实现可以看到几个关键动作:
- 图像处理:若
img维度小于 3,先扩展出通道维;随后把图像从 HWC 转置为 CHW,并转换为连续的 torch.Tensor 存入packed_results['inputs']。若图像内存不是 C 连续(C-contiguous),会先np.ascontiguousarray再转换,确保送入模型的数据布局正确。 - Ground truth 打包:若结果中包含
gt_seg_map,会将其扩展为(1, H, W)并转为int64张量,封装为PixelData后写入data_sample.gt_sem_seg。此外还支持gt_edge_map(边缘图)与gt_depth_map(深度图)的可选打包。 - meta 信息收集:
img_meta字典的内容由meta_keys决定,默认包含('img_path', 'seg_map_path', 'ori_shape', 'img_shape', 'pad_shape', 'scale_factor', 'flip', 'flip_direction', 'reduce_zero_label'),并通过data_sample.set_metainfo(img_meta)写入SegDataSample。这些键描述了图像的原始尺寸、padding 后的尺寸、缩放因子与翻转状态,是后处理阶段还原预测结果的关键依据。
仓库中 tests/test_datasets/test_formatting.py 对PackSegInputs的输入输出格式与__repr__输出做了单元测试,可作为自定义 pipeline 时的参考样板。
SegDataSample:连接各组件的统一数据结构
SegDataSample 是 MMSegmentation 的数据结构接口,用于连接不同组件。它实现了抽象数据元素mmengine.structures.BaseDataElement,并将属性划分为三类,全部是PixelData类型:
gt_sem_seg:语义分割的 ground truth;pred_sem_seg:语义分割的预测结果;seg_logits:预测的 logits(归一化前的分割分数)。
这三个属性通过 property + setter 暴露,并在内部使用set_field落盘,确保类型约束。例如:
>>> from mmseg.structures import SegDataSample >>> data_sample = SegDataSample() >>> gt_sem_seg_data = dict(data=torch.randint(0, 2, (1, 4, 4))) >>> data_sample.gt_sem_seg = PixelData(**gt_sem_seg_data) >>> assert 'gt_sem_seg' in data_sample由于SegDataSample同时携带 meta 信息与像素级数据,它可以在「数据加载器 → 预处理器 → 模型 → 评估器」整条链路上无歧义地传递图像的几何信息(如ori_shape、pad_shape)与监督信号,这也是数据流格式约定的核心载体。更详细的结构说明参见数据结构文档。
数据预处理器到模型:SegDataPreProcessor 的批处理与归一化
虽然数据流图中将数据预处理器与模型分开绘制,但实际上数据预处理器是模型的一部分(BaseSegmentor继承自mmengine.model.BaseModel,预处理器作为其子模块)。MMSegmentation 提供的实现是 SegDataPreProcessor,继承自mmengine.model.BaseDataPreprocessor,在语义分割场景下额外完成了以下工作:
- Collate 与设备搬移:通过
cast_data将数据搬到目标设备; - Padding 与堆叠:将 batch 内的图像 pad 到固定
size或size_divisor的整数倍后堆叠成 4D 张量; - 颜色空间转换:按需进行 BGR→RGB(
bgr_to_rgb)或 RGB→BGR(rgb_to_bgr)通道重排(二者不能同时为 True); - 归一化:当且仅当同时指定了
mean与std时才启用,否则跳过归一化(这与mmengine.ImgDataPreprocessor的行为不同); - 批级增强:训练时支持 mixup / cutmix 等 batch augmentation。
其关键构造参数与语义如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
mean/std | None | 各通道像素均值与标准差,必须同时给出才会启用归一化 |
size | None | 固定的 padding 尺寸(tuple) |
size_divisor | None | padding 后尺寸需为该值的整数倍 |
pad_val | 0 | 图像 padding 填充值 |
seg_pad_val | 255 | 分割图 padding 填充值(255 为忽略类别索引的通用约定) |
bgr_to_rgb/rgb_to_bgr | False | 是否做通道转换,二者互斥 |
batch_augments | None | 批级增强配置(如 mixup、cutmix) |
test_cfg | None | 测试时的 padding 配置,支持size或size_divisor键 |
从 forward 的实现可以看到:训练分支要求data_samples必须存在(training=True时断言非空),随后调用stack_batch完成 padding 与堆叠;测试分支则要求 batch 内图像尺寸一致,若配置了test_cfg则同样执行stack_batch,并把 padding 信息回写进每个data_sample的 metainfo(供后处理去除 padding 区域),否则直接torch.stack堆叠。
数据预处理器的返回值仍是包含inputs和data_samples的字典,只是inputs升级为批处理图像的 4D 张量,data_samples中追加了用于数据预处理的额外元信息。当字典传递给网络时,会被解包为两个独立参数:
dict( inputs=torch.Tensor, data_samples=List[SegDataSample] )class Network(BaseSegmentor): def forward(self, inputs: torch.Tensor, data_samples: List[SegDataSample], mode: str): pass模型的前向传播有 3 种模式,由入参mode控制,详见模型教程:
'tensor':整网前向,返回无任何后处理的张量,行为等同于普通nn.Module;'predict':前向并返回完整后处理的预测结果,即SegDataSample列表;'loss':前向并返回由输入与数据样本计算得到的损失字典。
仓库中 tests/test_models/test_data_preprocessor.py 对SegDataPreProcessor的归一化开关、通道转换互斥断言、padding 行为等进行了覆盖测试,是理解预处理器语义的第一手资料。
模型输出:从 logits 到 SegDataSample 的后处理与损失计算
三种前向模式对应三种输出
如模型教程所述,模型三种前向模式对应三种输出:train_step调用'loss'模式,输出损失字典;test_step/val_step调用'predict'模式,输出预测结果。在test_step或val_step中,推理结果会被传递给Evaluator,关于评估器的更多信息参见评估文档。
在 BaseSegmentor 中,forward作为统一入口按mode分发到loss、predict、_forward三个抽象方法;这三个方法的具体实现在EncoderDecoder(encoder_decoder.py)等具体 segmentor 中给出:
- loss:
extract_feat()提取多级特征 →_decode_head_forward_train()计算主解码头损失 → 存在辅助头时追加_auxiliary_head_forward_train()损失(辅助头仅用于训练阶段的深度监督,推理时被丢弃); - predict:
inference()(整图whole_inference或滑动窗口slide_inference)得到 logits →postprocess_result()打包为SegDataSample列表; - _forward:
extract_feat()→decode_head.forward(),返回未后处理的张量。
postprocess_result:推理结果的后处理打包
推理之后,MMSegmentation 的 postprocess_result 会对分割结果做一系列后处理,将神经网络生成的分割 logits、经过argmax得到的预测 mask 以及 ground truth(若存在)打包进SegDataSample实例。其处理顺序为:
- 去除 padding 区域:依据
data_sample中的padding_size(或img_padding_size)裁剪掉 padding 边距; - 翻转还原:若测试时启用了 flip(TTA 场景),依据
flip_direction对 logits 做水平或垂直翻转还原; - 尺寸还原:将 logits 用双线性插值
resize回ori_shape原始尺寸; - 类别解码:当类别数
C > 1时对 logits 执行argmax(dim=0)得到预测 mask;当C == 1(二分类)时先sigmoid再按decode_head.threshold阈值化; - 结果封装:将
seg_logits与pred_sem_seg分别写入PixelData并 set 到data_sample。
因此,postprocess_result的返回值是SegDataSample的List,每个实例的关键属性为pred_sem_seg(预测 mask)与seg_logits(归一化前的 logits),并保留输入侧的 metainfo(img_path、ori_shape等),便于可视化与评估。
在EncoderDecoder中,滑动窗口推理(slide_inference)按test_cfg中的stride与crop_size在图像上滑动裁剪、逐块encode_decode,将各块 logits 通过 padding 累加到全图并除以覆盖次数取平均;测试模式由test_cfg.mode('whole'/'slide')控制。这部分行为同样在 tests/test_models/test_segmentors/test_encoder_decoder.py 中通过构造seg_logits直接调用postprocess_result得到验证。
loss_by_feat:解码头统一的损失计算接口
与数据预处理器一致,损失函数也是模型的一部分——它是解码头(decode head)的属性之一。在 MMSegmentation 中,decode_head的 loss_by_feat 方法是计算损失的统一接口。
参数:
seg_logits(Tensor):解码头前向函数的输出;batch_data_samples(List[SegDataSample]):分割数据样本,通常包括metainfo与gt_sem_seg等信息。
返回值:
dict[str, Tensor]:损失组件的字典。
从实现看,loss_by_feat的内部流程是:先用_stack_batch_gt把 batch 内各样本的gt_sem_seg.data堆叠为(N, 1, H, W)的标签张量;将seg_logits双线性 resize 到与标签一致的分辨率(对齐方式由align_corners控制);若配置了sampler(如 OHEM 采样器)则采样得到逐像素权重seg_weight;随后遍历loss_decode(单个或ModuleList)计算各类损失(如 CrossEntropyLoss),并对同名损失累加;最后额外返回acc_seg像素准确率作为训练监控指标。
注意:train_step会将'loss'模式返回的损失字典传递给 OptimWrapper,以完成梯度计算与模型权重更新,更多细节参见模型教程中的 train_step 章节。loss_by_feat的行为在 tests/test_models/test_heads/test_decode_head.py 中针对不同 head 配置与损失类型有系统性覆盖。
数据流格式约定速查
| 阶段 | 数据形态 | 关键实现 |
|---|---|---|
| DataLoader 输出 | dict(inputs=List[Tensor], data_samples=List[SegDataSample]) | PackSegInputs |
| 数据样本载体 | SegDataSample(gt_sem_seg/pred_sem_seg/seg_logits三个 PixelData) | seg_data_sample.py |
| 预处理器输出 | dict(inputs=4D Tensor, data_samples=List[SegDataSample]) | SegDataPreProcessor |
| 模型 forward 入口 | forward(inputs, data_samples, mode),mode ∈ {tensor,predict,loss} | base.py |
| 推理输出 | List[SegDataSample],含pred_sem_seg与seg_logits | postprocess_result |
| 训练输出 | dict[str, Tensor]损失字典(含acc_seg) | loss_by_feat |
理解这条数据流,是深入 MMSegmentation 二次开发(自定义数据增强、自研解码头、接入新评估指标)的前提:只要保持PackSegInputs的输出格式与SegDataSample的字段约定,上游的变换与下游的模型、评估器都可以无缝组合。
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考