拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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

MMSegmentation 数据流全解析从 DataLoader 到损失回传的格式约定与源码实现【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation本篇文章围绕 MMSegmentationOpenMMLab 语义分割工具箱中由 MMEngine Runner 调度的完整数据流展开逐段剖析数据加载器、数据预处理器、模型前向与损失计算之间的数据格式约定。读完本文你将掌握PackSegInputs打包规则、SegDataSample结构、SegDataPreProcessor批处理流程、模型三种前向模式及postprocess_result后处理细节并能在训练与推理链路中准确追踪每一份数据的形态变化。数据流概述Runner 如何串联整个训练与评测管线在 MMEngine 的设计中Runner 相当于整个框架的集成器它覆盖了框架的几乎所有方面肩负着组织与调度各模块的责任。因此模块之间的数据流也由 Runner 统一控制——训练循环TrainLoop、验证循环ValLoop与测试循环TestLoop均由 Runner 在合适的时机启动并在每次迭代中驱动「数据加载 → 预处理 → 模型前向 → 优化/评估」这条主链路。MMSegmentation 对 loop 的默认设置是使用IterBasedTrainLoop按迭代数训练模型默认共 20000 次迭代并且每 2000 次迭代后执行一次验证。对应配置如下train_cfg dict(typeIterBasedTrainLoop, max_iters20000, val_interval2000) val_cfg dict(typeValLoop) test_cfg dict(typeTestLoop)需要说明的是上文描述的数据流适用于「用户没有自定义 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 与 SegDataSampleDataLoader 与 PackSegInputs 的分工数据加载器DataLoader是 MMEngine 训练与测试流程中的重要组件它源自 PyTorch 并保持一致的语义从文件系统加载数据原始数据经过数据准备流程pipeline后发送给数据预处理器。MMSegmentation 在 PackSegInputs 中定义了默认的数据格式它是train_pipeline和test_pipeline的最后一个组件。有关数据转换 pipeline 的更多信息可参阅数据转换文档。在没有任何修改的情况下PackSegInputs.transform的返回值是一个包含inputs和data_samples的字典。以下伪代码展示了 mmseg 中数据加载器输出的数据类型——它是从数据集中取回的一批数据样本数据加载器将它们打包成字典inputs是输入进模型的张量列表data_samples则包含输入图像的 meta 信息和对应的 ground truthdict( inputsList[torch.Tensor], data_samplesList[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 truthpred_sem_seg语义分割的预测结果seg_logits预测的 logits归一化前的分割分数。这三个属性通过 property setter 暴露并在内部使用set_field落盘确保类型约束。例如 from mmseg.structures import SegDataSample data_sample SegDataSample() gt_sem_seg_data dict(datatorch.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→RGBbgr_to_rgb或 RGB→BGRrgb_to_bgr通道重排二者不能同时为 True归一化当且仅当同时指定了mean与std时才启用否则跳过归一化这与mmengine.ImgDataPreprocessor的行为不同批级增强训练时支持 mixup / cutmix 等 batch augmentation。其关键构造参数与语义如下参数默认值说明mean/stdNone各通道像素均值与标准差必须同时给出才会启用归一化sizeNone固定的 padding 尺寸tuplesize_divisorNonepadding 后尺寸需为该值的整数倍pad_val0图像 padding 填充值seg_pad_val255分割图 padding 填充值255 为忽略类别索引的通用约定bgr_to_rgb/rgb_to_bgrFalse是否做通道转换二者互斥batch_augmentsNone批级增强配置如 mixup、cutmixtest_cfgNone测试时的 padding 配置支持size或size_divisor键从 forward 的实现可以看到训练分支要求data_samples必须存在trainingTrue时断言非空随后调用stack_batch完成 padding 与堆叠测试分支则要求 batch 内图像尺寸一致若配置了test_cfg则同样执行stack_batch并把 padding 信息回写进每个data_sample的 metainfo供后处理去除 padding 区域否则直接torch.stack堆叠。数据预处理器的返回值仍是包含inputs和data_samples的字典只是inputs升级为批处理图像的 4D 张量data_samples中追加了用于数据预处理的额外元信息。当字典传递给网络时会被解包为两个独立参数dict( inputstorch.Tensor, data_samplesList[SegDataSample] )class Network(BaseSegmentor): def forward(self, inputs: torch.Tensor, data_samples: List[SegDataSample], mode: str): pass模型的前向传播有 3 种模式由入参mode控制详见模型教程tensor整网前向返回无任何后处理的张量行为等同于普通nn.Modulepredict前向并返回完整后处理的预测结果即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三个抽象方法这三个方法的具体实现在EncoderDecoderencoder_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实例。其处理顺序为去除 padding 区域依据data_sample中的padding_size或img_padding_size裁剪掉 padding 边距翻转还原若测试时启用了 flipTTA 场景依据flip_direction对 logits 做水平或垂直翻转还原尺寸还原将 logits 用双线性插值resize回ori_shape原始尺寸类别解码当类别数C 1时对 logits 执行argmax(dim0)得到预测 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并保留输入侧的 metainfoimg_path、ori_shape等便于可视化与评估。在EncoderDecoder中滑动窗口推理slide_inference按test_cfg中的stride与crop_size在图像上滑动裁剪、逐块encode_decode将各块 logits 通过 padding 累加到全图并除以覆盖次数取平均测试模式由test_cfg.modewhole/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(inputsList[Tensor], data_samplesList[SegDataSample])PackSegInputs数据样本载体SegDataSamplegt_sem_seg/pred_sem_seg/seg_logits三个 PixelDataseg_data_sample.py预处理器输出dict(inputs4D Tensor, data_samplesList[SegDataSample])SegDataPreProcessor模型 forward 入口forward(inputs, data_samples, mode)mode ∈ {tensor,predict,loss}base.py推理输出List[SegDataSample]含pred_sem_seg与seg_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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门