ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

MMSegmentation 新增模块实战指南:从零开发 backbone、head、loss 与 segmentor

2026/9/16 13:49:53 拓冰建站 浏览量
MMSegmentation 新增模块实战指南:从零开发 backbone、head、loss 与 segmentor MMSegmentation 新增模块实战指南从零开发 backbone、head、loss 与 segmentor【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation本指南面向需要在 MMSegmentation 1.x 中扩展自定义算法的开发者。全文以官方文档《新增模块》为骨架围绕主干网络backbone、解码头decode head、损失函数loss、数据预处理器data preprocessor与分割器segmentor五类组件的开发流程展开并逐一对照仓库源码给出底层实现依据帮助你快速掌握基于注册机制Registry的模块化开发范式能够独立为项目添加新组件并在配置文件中即插即用。一、模块化开发的基石注册器Registry机制在动手写代码之前需要先理解 MMSegmentation 组件之所以能通过type字符串在配置文件中被实例化的底层机制——注册器。MMSegmentation 提供了 21 个注册节点registry node全部定义在 mmseg/registry/registry.py 中用于跨项目支持模块复用每个节点都是 MMEngine 根注册器的子节点。例如其中的MODELS注册器管理所有模型相关组件# mmseg/registry/registry.py from mmengine.registry import Registry # manage all kinds of models MODELS Registry(model, parentMMENGINE_MODELS)所有新增组件backbone、head、loss、segmentor、data preprocessor都使用同一个MODELS注册器通过MODELS.register_module()装饰器完成注册。注册之后配置文件中的typeMyComponent就会在构建模型时被自动解析并实例化。二、添加新的主干网络backbone1. 创建 backbone 文件在mmseg/models/backbones/下创建新文件以 MobileNet 为例即mmseg/models/backbones/mobilenet.py所有主干网络都继承自nn.Module并通过装饰器注册import torch.nn as nn from mmseg.registry import MODELS MODELS.register_module() class MobileNet(nn.Module): def __init__(self, arg1, arg2): pass def forward(self, x): # should return a tuple pass def init_weights(self, pretrainedNone): pass三个核心方法的分工__init__接收配置参数并构建网络层注意 backbone 一般需要支持norm_cfg、style如pytorch与caffe的初始化差异等通用参数forward(x)前向传播必须返回一个 tuple即多级特征图列表供后续 neck/head 使用这是与 MMSegmentation 的EncoderDecoder等分割器的约定init_weights(pretrainedNone)加载预训练权重或执行参数初始化。2. 在__init__.py中引入模块打开mmseg/models/backbones/__init__.py追加导入语句from .mobilenet import MobileNet这一步不可省略因为注册器只有在模块被 import 之后才能感知到新注册的类。以当前仓库为例mmseg/models/backbones/init.py 中已经注册了ResNet、ResNetV1c、ResNetV1d、HRNet、MobileNetV2、MobileNetV3、SwinTransformer、MixVisionTransformer等 20 余个 backbone新模块加入后同样会被写入__all__导出列表。3. 在配置文件中使用model dict( ... backbonedict( typeMobileNet, arg1xxx, arg2xxx), ... )构建时MMEngine 的MODELS.build()会根据type查表找到MobileNet类并把arg1、arg2等字段作为关键字参数传入构造函数。参考仓库中真实 backbone 的写法mmseg/models/backbones/resnet.py含ResNet/ResNetV1c/ResNetV1d、mmseg/models/backbones/mobilenet_v2.py、mmseg/models/backbones/swin.py它们共同遵循返回多级特征 tuple 支持预训练权重初始化的约定。三、添加新的头head1. 继承 BaseDecodeHeadMMSegmentation 提供BaseDecodeHead作为所有分割头的基类新实现的解码头都应从它派生。基类定义在 mmseg/models/decode_heads/decode_head.py。以 PSPNetPyramid Scene Parsing Network为例在mmseg/models/decode_heads/psp_head.py中实现新的解码头需要覆盖三个函数from mmseg.registry import MODELS MODELS.register_module() class PSPHead(BaseDecodeHead): def __init__(self, pool_scales(1, 2, 3, 6), **kwargs): super(PSPHead, self).__init__(**kwargs) def init_weights(self): pass def forward(self, inputs): pass2. 理解 BaseDecodeHead 提供的现成能力BaseDecodeHead并非空壳它在__init__中已经完成了大量公共工作见 decode_head.py输入选择与变换通过in_channels、in_index、input_transform三个参数控制从 backbone 的多级特征中选取哪些层。input_transform支持三种取值——None只允许单层特征、resize_concat将多层特征 resize 到同一尺寸后 concat常用于 HRNet 的 FCN head、multiple_select将多层特征打包成 list 传入 head损失构建loss_decode参数dict 或 dict 列表会被MODELS.build()实例化为损失模块并存入self.loss_decode默认是dict(typeCrossEntropyLoss, use_sigmoidFalse, loss_weight1.0)分类层自动创建self.conv_seg nn.Conv2d(channels, out_channels, kernel_size1)并可选nn.Dropout2d(dropout_ratio)loss 与 predict 的统一入口基类实现了loss()内部走forward() - loss_by_feat()和predict()内部走forward() - predict_by_feat()两条完整调用链子类只需实现forward()即可同时获得训练和推理能力。以仓库中真实的 PSPHead 实现mmseg/models/decode_heads/psp_head.py为例其forward通过_forward_feature完成金字塔池化模块PPM的特征融合再调用基类的cls_seg完成逐像素分类def forward(self, inputs): output self._forward_feature(inputs) output self.cls_seg(output) return output3. 注册并配置同样地在mmseg/models/decode_heads/__init__.py中导入PSPHead注册器即可找到它。PSPNet 的完整配置如下来自官方文档与 configs/pspnet/pspnet_r50-d8_4xb2-40k_cityscapes-512x1024.py 中引用的_base_/models/pspnet_r50-d8.py同源norm_cfg dict(typeSyncBN, requires_gradTrue) model dict( typeEncoderDecoder, pretrainedpretrain_model/resnet50_v1c_trick-2cccc1ad.pth, backbonedict( typeResNetV1c, depth50, num_stages4, out_indices(0, 1, 2, 3), dilations(1, 1, 2, 4), strides(1, 2, 1, 1), norm_cfgnorm_cfg, norm_evalFalse, stylepytorch, contract_dilationTrue), decode_headdict( typePSPHead, in_channels2048, in_index3, channels512, pool_scales(1, 2, 3, 6), dropout_ratio0.1, num_classes19, norm_cfgnorm_cfg, align_cornersFalse, loss_decodedict( typeCrossEntropyLoss, use_sigmoidFalse, loss_weight1.0)))关键参数说明in_channels2048ResNetV1c 第 4 阶段输出的通道数对应in_index3in_index3取 backboneout_indices的第 3 层特征pool_scales(1, 2, 3, 6)金字塔池化的四组尺度对应 PSPNet 原文num_classes19Cityscapes 的类别数含 ignore 的语义分割任务通常设为有效类别数loss_decode训练损失配置use_sigmoidFalse表示使用 softmax 交叉熵loss_weight用于多损失加权。从 BaseDecodeHead 的构造代码 还可以看到两点易踩坑的约束当num_classes2时建议显式设置out_channels1并配合threshold默认 0.3做二分类阈值分割若out_channels与num_classes均不等于 1 且二者不相等会直接抛出ValueError。四、添加新的损失函数loss假设要为分割任务新增一个名为MyLoss的损失函数。损失模块位于mmseg/models/losses/目录注意是losses不是loss当前仓库已内置CrossEntropyLoss、DiceLoss、FocalLoss、LovaszLoss、OhemCrossEntropyLoss、TverskyLoss等十余种见 mmseg/models/losses/init.py。1. 实现损失函数在mmseg/models/losses/my_loss.py中实现核心是利用装饰器weighted_loss对逐元素损失进行加权与归约import torch import torch.nn as nn from mmseg.registry import MODELS from .utils import weighted_loss weighted_loss def my_loss(pred, target): assert pred.size() target.size() and target.numel() 0 loss torch.abs(pred - target) return loss MODELS.register_module() class MyLoss(nn.Module): def __init__(self, reductionmean, loss_weight1.0): super(MyLoss, self).__init__() self.reduction reduction self.loss_weight loss_weight def forward(self, pred, target, weightNone, avg_factorNone, reduction_overrideNone): assert reduction_override in (None, none, mean, sum) reduction ( reduction_override if reduction_override else self.reduction) loss self.loss_weight * my_loss( pred, target, weight, reductionreduction, avg_factoravg_factor) return lossweighted_loss装饰器的底层实现在 mmseg/models/losses/utils.py它要求被装饰函数只计算逐元素element-wise损失然后自动为其附加weight、reduction、avg_factor三个参数——weight用于逐元素加权如 OHEM 采样权重reduction支持none/mean/sumavg_factor用于在mean模式下按自定义因子归一内部加了torch.finfo(torch.float32).eps防止除以零。2. 注册损失在mmseg/models/losses/__init__.py中导入from .my_loss import MyLoss, my_loss3. 在配置中使用修改解码头中的loss_decode字段即可启用loss_weight用于平衡多个损失loss_decodedict(typeMyLoss, loss_weight1.0))也可以传入 dict 列表同时使用多个损失例如在 decode_head.py 的注释 中给出的组合loss_decode[ dict(typeCrossEntropyLoss, loss_nameloss_ce), dict(typeDiceLoss, loss_nameloss_dice), ]注意loss_name需要以loss_开头才会被纳入反向传播并在训练日志中显示。五、添加新的数据预处理器data preprocessor在 MMSegmentation 1.x 中SegDataPreProcessor定义在 mmseg/models/data_preprocessor.py负责将数据复制到目标设备并按默认的模型输入格式做归一化、padding 等预处理。自定义预处理器同样走注册 配置三步。1. 创建文件在mmseg/models/下创建my_datapreprocessor.py继承 MMEngine 的BaseDataPreprocessorfrom mmengine.model import BaseDataPreprocessor from mmseg.registry import MODELS MODELS.register_module() class MyDataPreProcessor(BaseDataPreprocessor): def __init__(self, **kwargs): super().__init__(**kwargs) def forward(self, data: dict, training: boolFalse) - Dict[str, Any]: # TODO Define the logic for data pre-processing in the forward method passforward方法接收一个包含inputs与data_samples的 dict需要在其中实现将输入图像搬到self.device、按mean/std归一化、按pad_size_divisor对齐尺寸等逻辑。2. 导入注册在mmseg/models/__init__.py中导入from .my_datapreprocessor import MyDataPreProcessor3. 在配置中使用model dict( data_preprocessordict(typeMyDataPreProcessor), ... )真实项目中data_preprocessor通常在基类配置中定义如_base_/models/pspnet_r50-d8.py中的data_preprocessor配置然后在具体任务配置里覆盖尺寸等参数。例如 configs/pspnet/pspnet_r50-d8_4xb2-40k_cityscapes-512x1024.py 中crop_size (512, 1024) data_preprocessor dict(sizecrop_size) model dict(data_preprocessordata_preprocessor)六、开发新的分割器segmentor分割器是定义backbone 如何与 head 组合、整体算法如何执行的算法架构。MMSegmentation 中所有分割器的基类是BaseSegmentor定义在 mmseg/models/segmentors/base.py它继承自 MMEngine 的BaseModel统一了前向过程的三种模式。1. 理解三种前向模式BaseSegmentor.forward()根据mode参数分发到三个抽象方法见 base.pymodetensor调用_forward()仅做网络前向backbone → neck → head不做任何后处理用于提取特征或 ONNX 导出modepredict调用predict()前向后还要做 resize 回原图尺寸、argmax多类或sigmoid threshold二分类等后处理返回SegDataSample列表modeloss调用loss()根据data_samples中的gt_sem_seg计算损失字典用于训练。开发新分割器时需要重写与这三种模式对应的loss、predict、_forward三个方法。2. 实现新分割器在mmseg/models/segmentors/my_segmentor.py中实现from typing import Dict, Optional, Union import torch from mmseg.registry import MODELS from mmseg.models import BaseSegmentor MODELS.register_module() class MySegmentor(BaseSegmentor): def __init__(self, **kwargs): super().__init__(**kwargs) # TODO users should build components of the network here def loss(self, inputs: Tensor, data_samples: SampleList) - dict: Calculate losses from a batch of inputs and data samples. pass def predict(self, inputs: Tensor, data_samples: OptSampleListNone) - SampleList: Predict results from a batch of inputs and data samples with post- processing. pass def _forward(self, inputs: Tensor, data_samples: OptSampleList None) - Tuple[List[Tensor]]: Network forward process. Usually includes backbone, neck and head forward without any post- processing. pass三个方法的职责边界_forward串联 backbone、neck、head 的纯前向返回中间张量loss在_forward结果之上计算训练损失predict在_forward结果之上完成后处理反 padding、反 flip、resize 到原图、argmax/sigmoid并写入SegDataSample。BaseSegmentor还内置了postprocess_result方法base.py统一处理 padding 裁剪、水平/垂直 flip 恢复、resize 回ori_shape、多类别argmax与二分类sigmoid decode_head.threshold等后处理逻辑新分割器可以直接复用它。3. 注册与使用在mmseg/models/segmentors/__init__.py中导入from .my_segmentor import MySegmentor然后在配置文件中通过type指定model dict( typeMySegmentor ... )仓库中已内置多种分割器实现可供参考mmseg/models/segmentors/encoder_decoder.py最常用的 EncoderDecoder、mmseg/models/segmentors/cascade_encoder_decoder.py级联式、mmseg/models/segmentors/multimodal_encoder_decoder.py多模态以及 mmseg/models/segmentors/seg_tta.py测试时增强 TTA。七、总结新增模块的通用三步法无论新增哪类组件流程都高度统一可以概括为三步法实现并注册在对应目录backbones/、decode_heads/、losses/、segmentors/等创建新文件用MODELS.register_module()装饰类导入导出在所属包的__init__.py中from .xxx import Xxx让注册器感知新模块配置即用在模型配置中通过typeXxx加上对应的参数字段完成实例化无需修改任何框架源码。这套机制的本质是借助 MMEngine 注册表将类与配置字符串解耦——新增算法时零侵入、即插即用这也是 MMSegmentation 能够以统一规范承载 PSPNet、DeepLabV3、SegFormer、Mask2Former 等数十种分割算法配置示例见 configs/ 目录下各算法子目录的架构基础。按照本指南开发的自定义组件同样可以直接参与训练tools/train.py、测试tools/test.py与推理demo/image_demo.py全流程。【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考