news 2026/9/16 12:39:25

MMSegmentation 数据流全解析:从 DataLoader 到损失回传的格式约定与源码实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MMSegmentation 数据流全解析:从 DataLoader 到损失回传的格式约定与源码实现

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 中的TrainLoopValLoopTestLoop,且没有在自定义模型中覆写train_stepval_steptest_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_pipelinetest_pipeline最后一个组件。有关数据转换 pipeline 的更多信息,可参阅数据转换文档。

在没有任何修改的情况下,PackSegInputs.transform的返回值是一个包含inputsdata_samples的字典。以下伪代码展示了 mmseg 中数据加载器输出的数据类型——它是从数据集中取回的一批数据样本,数据加载器将它们打包成字典;inputs是输入进模型的张量列表,data_samples则包含输入图像的 meta 信息和对应的 ground truth:

dict( inputs=List[torch.Tensor], data_samples=List[SegDataSample] )

PackSegInputs 的打包细节

从 formatting.py 的transform实现可以看到几个关键动作:

  1. 图像处理:若img维度小于 3,先扩展出通道维;随后把图像从 HWC 转置为 CHW,并转换为连续的 torch.Tensor 存入packed_results['inputs']。若图像内存不是 C 连续(C-contiguous),会先np.ascontiguousarray再转换,确保送入模型的数据布局正确。
  2. Ground truth 打包:若结果中包含gt_seg_map,会将其扩展为(1, H, W)并转为int64张量,封装为PixelData后写入data_sample.gt_sem_seg。此外还支持gt_edge_map(边缘图)与gt_depth_map(深度图)的可选打包。
  3. 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_shapepad_shape)与监督信号,这也是数据流格式约定的核心载体。更详细的结构说明参见数据结构文档。

数据预处理器到模型:SegDataPreProcessor 的批处理与归一化

虽然数据流图中将数据预处理器与模型分开绘制,但实际上数据预处理器是模型的一部分BaseSegmentor继承自mmengine.model.BaseModel,预处理器作为其子模块)。MMSegmentation 提供的实现是 SegDataPreProcessor,继承自mmengine.model.BaseDataPreprocessor,在语义分割场景下额外完成了以下工作:

  • Collate 与设备搬移:通过cast_data将数据搬到目标设备;
  • Padding 与堆叠:将 batch 内的图像 pad 到固定sizesize_divisor的整数倍后堆叠成 4D 张量;
  • 颜色空间转换:按需进行 BGR→RGB(bgr_to_rgb)或 RGB→BGR(rgb_to_bgr)通道重排(二者不能同时为 True);
  • 归一化:当且仅当同时指定了meanstd时才启用,否则跳过归一化(这与mmengine.ImgDataPreprocessor的行为不同);
  • 批级增强:训练时支持 mixup / cutmix 等 batch augmentation。

其关键构造参数与语义如下:

参数默认值说明
mean/stdNone各通道像素均值与标准差,必须同时给出才会启用归一化
sizeNone固定的 padding 尺寸(tuple)
size_divisorNonepadding 后尺寸需为该值的整数倍
pad_val0图像 padding 填充值
seg_pad_val255分割图 padding 填充值(255 为忽略类别索引的通用约定)
bgr_to_rgb/rgb_to_bgrFalse是否做通道转换,二者互斥
batch_augmentsNone批级增强配置(如 mixup、cutmix)
test_cfgNone测试时的 padding 配置,支持sizesize_divisor

从 forward 的实现可以看到:训练分支要求data_samples必须存在(training=True时断言非空),随后调用stack_batch完成 padding 与堆叠;测试分支则要求 batch 内图像尺寸一致,若配置了test_cfg则同样执行stack_batch,并把 padding 信息回写进每个data_sample的 metainfo(供后处理去除 padding 区域),否则直接torch.stack堆叠。

数据预处理器的返回值仍是包含inputsdata_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_stepval_step中,推理结果会被传递给Evaluator,关于评估器的更多信息参见评估文档。

在 BaseSegmentor 中,forward作为统一入口按mode分发到losspredict_forward三个抽象方法;这三个方法的具体实现在EncoderDecoder(encoder_decoder.py)等具体 segmentor 中给出:

  • lossextract_feat()提取多级特征 →_decode_head_forward_train()计算主解码头损失 → 存在辅助头时追加_auxiliary_head_forward_train()损失(辅助头仅用于训练阶段的深度监督,推理时被丢弃);
  • predictinference()(整图whole_inference或滑动窗口slide_inference)得到 logits →postprocess_result()打包为SegDataSample列表;
  • _forwardextract_feat()decode_head.forward(),返回未后处理的张量。

postprocess_result:推理结果的后处理打包

推理之后,MMSegmentation 的 postprocess_result 会对分割结果做一系列后处理,将神经网络生成的分割 logits、经过argmax得到的预测 mask 以及 ground truth(若存在)打包进SegDataSample实例。其处理顺序为:

  1. 去除 padding 区域:依据data_sample中的padding_size(或img_padding_size)裁剪掉 padding 边距;
  2. 翻转还原:若测试时启用了 flip(TTA 场景),依据flip_direction对 logits 做水平或垂直翻转还原;
  3. 尺寸还原:将 logits 用双线性插值resizeori_shape原始尺寸;
  4. 类别解码:当类别数C > 1时对 logits 执行argmax(dim=0)得到预测 mask;当C == 1(二分类)时先sigmoid再按decode_head.threshold阈值化;
  5. 结果封装:将seg_logitspred_sem_seg分别写入PixelData并 set 到data_sample

因此,postprocess_result的返回值是SegDataSampleList,每个实例的关键属性为pred_sem_seg(预测 mask)与seg_logits(归一化前的 logits),并保留输入侧的 metainfo(img_pathori_shape等),便于可视化与评估。

EncoderDecoder中,滑动窗口推理(slide_inference)按test_cfg中的stridecrop_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]):分割数据样本,通常包括metainfogt_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
数据样本载体SegDataSamplegt_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_segseg_logitspostprocess_result
训练输出dict[str, Tensor]损失字典(含acc_segloss_by_feat

理解这条数据流,是深入 MMSegmentation 二次开发(自定义数据增强、自研解码头、接入新评估指标)的前提:只要保持PackSegInputs的输出格式与SegDataSample的字段约定,上游的变换与下游的模型、评估器都可以无缝组合。

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

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 12:39:23

C语言基础数据类型与内存操作15个关键点解析

1. 项目背景与核心价值"Back to Base-ics"这个标题乍看像文字游戏,实则精准概括了C语言开发中一个关键痛点——基础数据类型和内存操作的精准控制。在嵌入式开发、系统编程和高性能计算领域,程序员经常需要处理原始字节流、内存对齐和二进制协…

作者头像 李华
网站建设 2026/9/16 12:36:45

Playnite 游戏库便携同步:3 个高频目标一次配齐的实操教程

Playnite 游戏库便携同步:3 个高频目标一次配齐的实操教程 【免费下载链接】Playnite Video game library manager with support for wide range of 3rd party libraries and game emulation support, providing one unified interface for your games. 项目地址:…

作者头像 李华
网站建设 2026/9/16 12:36:29

纯前端基于WebCodecs实现高清录屏并导出MP4的完整方案

最近接了个内部培训系统的需求,要求在浏览器里给学员录制一段操作演示,录完直接下载 MP4,不给装任何插件,也不接受在线录屏工具那种强制水印。最开始想到的方案是 MediaRecorder,毕竟它 API 简单,几行代码就…

作者头像 李华
网站建设 2026/9/16 12:36:19

Java Swing花店管理系统实战:MVC架构与GUI业务流设计

简介:这是一套面向高校Java课程设计的花店管理系统纯GUI实现,专为Java初学者及课设学生打造,聚焦基础Swing组件开发与数据库交互能力训练,完全满足课程验收与答辩需求。资源包共54个文件,含10个核心Java源码&#xff0…

作者头像 李华
网站建设 2026/9/16 12:34:07

纯静态HTML地址发布页:从结构设计到Nginx部署全解析

简介:简洁美观地址发布页HTML源码是一套基于HTMLCSS的轻量级网页源码,适合个人站长、小微企业或个人用户快速搭建地址展示与发布页面。压缩包共6个文件,核心包含一个HTML页面和一份CSS样式表,另附站点图标、文本说明以及两个url快…

作者头像 李华
网站建设 2026/9/16 12:33:58

区块链跨链交易优化与Cber技术架构解析

1. 项目背景与行业痛点当我们在2023年回望数字资产领域的发展历程,会发现一个有趣的现象:尽管区块链技术已经诞生十余年,但绝大多数加密资产的交易模式依然停留在"古典加密时代"。这个术语在业内特指那些依赖中心化交易所、受限于法…

作者头像 李华