基于 PyTorch] 的图像语义分割神经网络 Python 库。
该库的主要特性包括:
- 极其简单的高层 API(仅需两行代码即可创建神经网络)
- 12 种编码器-解码器模型架构(Unet, Unet++, Segformer, DPT, ...)
- 800+ 个预训练的基于卷积和变换的编码器,包括 timm 支持
- 训练流程中常用的指标和损失函数(Dice, Jaccard, Tversky, ...)
- 支持 ONNX 导出,且兼容 torch script/trace/compile
🤝 赞助商:withoutBG
withoutBG 是一款高质量的背景移除工具。他们使用 smp.Unet 构建了其开源图像抠图和精修模型,并自豪地赞助了本项目。
📚 项目文档 📚
访问 Read The Docs 项目页面 或阅读以下 README 以了解更多关于 Segmentation Models Pytorch(简称 SMP)库的信息
📋 Table of content
- Quick start
- Examples
- Models and encoders
- Models API
- Installation
- Competitions won with the library
- Contributing
- Citing
- License
⏳ 快速开始
1. 使用 SMP 创建你的第一个 Segmentation 模型
分割模型只是一个 PyTorch torch.nn.Module,可以像这样轻松创建:
import segmentation_models_pytorch as smp
model = smp.Unet(
encoder_name="resnet34", # choose encoder, e.g. mobilenet_v2 or efficientnet-b7
encoder_weights="imagenet", # use `imagenet` pre-trained weights for encoder initialization
in_channels=1, # model input channels (1 for gray-scale images, 3 for RGB, etc.)
classes=3, # model output channels (number of classes in your dataset)
)
2. 配置数据预处理
所有编码器都具有预训练权重。以与权重预训练相同的方式准备数据可能会带来更好的结果(更高的指标分数和更快的收敛)。如果您训练的是整个模型而不仅仅是解码器,则没有必要这样做。
from segmentation_models_pytorch.encoders import get_preprocessing_fn
preprocess_input = get_preprocessing_fn('resnet18', pretrained='imagenet')
恭喜!你已完成!现在你可以使用你喜欢的框架来训练你的模型!
💡 示例
| 名称 | 链接 | Colab |
|---|---|---|
| 训练 在 OxfordPets 上进行宠物二分类分割 | Notebook | |
| 训练 在 CamVid 上进行汽车二分类分割 | Notebook | |
| 训练 在 CamVid 上进行多类分割 | Notebook | |
| 训练 由 @ternaus 进行的衣物二分类分割 | Repo | |
| 加载和推理 预训练 Segformer | Notebook | |
| 加载和推理 预训练 DPT | Notebook | |
| 加载和推理 预训练 UPerNet | Notebook | |
| 保存和加载 模型到本地 / HuggingFace Hub | Notebook | |
| 导出 训练好的模型到 ONNX | Notebook |
📦 模型和编码器
架构
| 架构 | 论文 | 文档 | 检查点 |
|---|---|---|---|
| Unet | paper | docs | |
| Unet++ | paper | docs | |
| MAnet | paper | docs | |
| Linknet | paper | docs | |
| FPN | paper | docs | |
| PSPNet | paper | docs | |
| PAN | paper | docs | |
| DeepLabV3 | paper | docs | |
| DeepLabV3+ | paper | docs | |
| UPerNet | paper | docs | checkpoints |
| Segformer | paper | docs | checkpoints |
| DPT | paper | docs | checkpoints |
编码器
该库提供了用于分割模型的广泛 预训练 编码器(也称为骨干网络)。我们不是使用分类模型最后一层的特征,而是提取 中间特征 并将其输入解码器以执行分割任务。
所有编码器都附带 预训练权重,这有助于在训练分割模型时实现 更快且更稳定的收敛。
鉴于支持的编码器选择范围广泛,您可以为特定用例选择最佳编码器,例如:
- 用于低延迟应用或边缘设备上的实时推理的 轻量级编码器(mobilenet/mobileone)。
- 用于涉及大量分割类别的复杂任务、提供更高精度的 高容量架构(convnext/swin/mit)。
通过选择合适的编码器,你可以平衡效率、性能和模型复杂度,以满足项目需求。
所有编码器及对应的预训练权重均列于文档中:
🔁 Models API
输入通道
输入通道参数允许你创建一个能够处理具有任意数量通道的张量的模型。 如果你使用来自 ImageNet 的预训练权重,第一个卷积层的权重将被复用:
- 对于 1 通道的情况,它将是第一个卷积层权重的总和。
- 否则,通道将按照
new_weight[:, i] = pretrained_weight[:, i % 3]的方式填充权重,然后使用new_weight * 3 / new_in_channels进行缩放。
model = smp.FPN('resnet34', in_channels=1)
mask = model(torch.ones([1, 1, 64, 64]))
辅助分类输出
所有模型均支持 aux_params 参数,其默认值为 None。
如果 aux_params = None,则不会创建分类辅助输出,否则
模型不仅会生成 mask,还会生成形状为 NC 的 label 输出。
分类头由 GlobalPooling->Dropout(可选)->Linear->Activation(可选) 层组成,可通过 aux_params 进行如下配置:
aux_params=dict(
pooling='avg', # one of 'avg', 'max'
dropout=0.5, # dropout ratio, default is None
activation='sigmoid', # activation function, default is None
classes=4, # define number of output labels
)
model = smp.Unet('resnet34', classes=4, aux_params=aux_params)
mask, label = model(x)
深度
Depth 参数指定编码器中下采样操作的数量,因此如果指定较小的 depth,
可以使模型更轻量。
model = smp.Unet('resnet34', encoder_depth=4)
🛠 安装
PyPI 版本:
$ pip install segmentation-models-pytorch
来自 GitHub 的最新版本:
$ pip install git+https://github.com/qubvel/segmentation_models.pytorch
🏆 使用本库赢得的竞赛
Segmentation Models 包在图像分割竞赛中被广泛使用。
在这里 你可以找到竞赛、获胜者姓名及其解决方案的链接。
🛠 使用 SMP 构建的项目
- withoutBG: 一个使用
smp.Unet进行图像抠图和细化模型的开源背景移除工具。查看 withoutBG Focus on HuggingFace.
🤝 贡献
- 以开发模式安装 SMP
make install_dev # Create .venv, install SMP in dev mode
- 运行测试和代码检查
make test # Run tests suite with pytest
make fixup # Ruff for formatting and lint checks
- 更新表(如果您添加了编码器)
make table # Generates a table with encoders and print to stdout
📝 引用
@misc{Iakubovskii:2019,
Author = {Pavel Iakubovskii},
Title = {Segmentation Models Pytorch},
Year = {2019},
Publisher = {GitHub},
Journal = {GitHub repository},
Howpublished = {\url{https://github.com/qubvel/segmentation_models.pytorch}}
}
🛡️ 许可证
该项目主要依据 MIT License 分发,部分文件受其他许可证约束。请查阅 LICENSES 及各文件中的许可证声明以仔细确认,尤其是用于商业用途时。