1. 为什么你需要关注 segmentation_models_pytorch

如果你正在做图像分割相关的项目,无论是医学影像分析、自动驾驶的感知模块,还是卫星图像的地物识别,我猜你肯定遇到过这几个让人头疼的问题:模型架构选择困难、从头训练收敛慢、以及代码实现复杂。这些问题常常让一个本应聚焦于解决实际业务问题的项目,在前期就耗费大量时间在“造轮子”上。

几年前我做项目时,为了复现一篇论文里的分割网络,光是调试数据加载和模型对齐就花了一两周。直到后来发现了 segmentation_models_pytorch(后面我们简称 SMP)这个库,我才真正体会到什么叫“快速原型开发”。这个由俄罗斯大神 Pavel Yakubovskiy 维护的库,简直是为研究者和开发者量身定做的“瑞士军刀”。它不是什么学术前沿的新模型,而是一个高度工程化、开箱即用的工具集合,能让你在几分钟内就搭出一个性能不俗的分割模型。

简单来说,SMP 的核心价值在于 “统一” 和 “预训练”。它把图像分割领域里最经典、最有效的 9 种模型架构(比如 Unet, Unet++, FPN, DeepLabV3+ 等)和超过 100 种预训练的编码器(也叫骨干网络,如 ResNet, EfficientNet, MobileNet 等)封装成了一个简洁的 API。你不需要再去 GitHub 上到处找各个模型的实现,也不用担心自己写的网络有 bug;更棒的是,你可以直接利用在 ImageNet 上预训练好的编码器权重,这能极大地加速模型收敛,提升最终精度,尤其是在你自己数据集不大的情况下。

我经常跟团队里的新人说,在验证一个分割任务的想法是否可行时,第一件事不是去读最新的论文,而是用 SMP 快速跑一个基线模型出来。它能帮你省下大量环境配置和模型调试的时间,让你把精力真正放在数据、业务逻辑和模型调优上。接下来,我就带你从零开始,手把手走通这个高效的工作流。

2. 十分钟完成环境配置与安装避坑指南

万事开头难,但 SMP 让开头变得异常简单。不过,在“简单”的路上也有一些小坑,我踩过,希望你别再踩。我们一步一步来。

首先,我强烈建议使用 Anaconda 来管理你的 Python 环境。这能避免不同项目间依赖包的冲突。假设你已经装好了 Anaconda,我们打开终端(或 Anaconda Prompt),创建一个新的虚拟环境。我习惯用 Python 3.8,兼容性比较好:

conda create -n smp_env python=3.8
conda activate smp_env

环境创建并激活后,就可以安装 SMP 了。安装命令非常简单:

pip install segmentation-models-pytorch

注意,这里就是第一个关键点! 执行这条命令时,pip 会自动为你安装 SMP 以及它所依赖的 torch 和 torchvision。但是,它安装的是 CPU 版本的最新版 PyTorch。对于图像分割这种计算密集型任务,用 CPU 训练简直是“用自行车拉货”,慢到怀疑人生。我们的目标肯定是使用 GPU 来加速。

所以,更稳妥的流程是:先安装与你的 CUDA 版本匹配的 PyTorch,再安装 SMP。你可以先去 PyTorch 官网 根据你的系统、CUDA 版本,获取正确的安装命令。例如,如果你的 CUDA 版本是 11.3,可以这样安装:

pip install torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cu113

安装完成后,可以在 Python 里验证一下:

import torch
print(torch.__version__)  # 查看 PyTorch 版本
print(torch.cuda.is_available())  # 查看 GPU 是否可用,应该返回 True

确认 GPU 版的 PyTorch 安装成功后,再安装 SMP 及其常用的配套工具库。因为 SMP 本身不包含数据增强等功能,我们通常会搭配 albumentations 这个强大的图像增强库,以及 opencv-python、matplotlib 等。

pip install segmentation-models-pytorch albumentations opencv-python matplotlib scikit-image

至此,核心环境就准备好了。你可以运行 pip list 查看已安装的包。这种先 PyTorch 后 SMP 的顺序,能确保你的环境是基于 GPU 的,为后续的高效训练打下基础。我见过不少朋友因为装错了 PyTorch 版本,训练时才发现 GPU 没用上,白白浪费几天时间,非常可惜。

3. 核心API解析:两行代码构建你的第一个分割模型

环境搞定,我们来感受一下 SMP 的魔力。它的 API 设计哲学是极简主义,创建模型通常只需要一行或两行代码。我们以最经典的 U-Net 架构为例。

import segmentation_models_pytorch as smp

model = smp.Unet(
    encoder_name="resnet34",        # 选择编码器(骨干网络)
    encoder_weights="imagenet",     # 使用在 ImageNet 上预训练的权重
    in_channels=3,                  # 输入图像的通道数,RGB 图为 3
    classes=1,                      # 输出类别数,二分类问题设为 1
)

看,就这么简单!smp.Unet() 这一行调用,就实例化了一个完整的 U-Net 模型。我们来拆解一下这几个关键参数:

  • encoder_name: 这是 SMP 的精华所在。除了 resnet34,你还可以尝试 efficientnet-b7、mobilenet_v2、timm-efficientnet-b0 等上百种选择。编码器负责从图像中提取多层次的特征,不同的编码器在速度和精度上有不同的权衡。
  • encoder_weights: 设置为 "imagenet" 意味着加载预训练权重。这相当于让模型从一个“见过世面”的起点开始学习,比随机初始化权重收敛快得多,效果也通常更好。如果你不想用预训练权重,设为 None 即可。
  • in_channels: 对于标准的彩色图像,这里是 3。如果你的数据是医学影像(如 MRI 的单个序列)或灰度图,可以设为 1。
  • classes: 你需要分割的类别总数。对于只是区分前景和背景的二分类任务(比如从背景中分割出猫),设为 1,模型会输出一个单通道的概率图。对于多分类任务(比如分割出道路、车辆、行人),这里就设为类别数。

想换模型?易如反掌。想把 U-Net 换成更先进的 U-Net++ 或 DeepLabV3+?只需要改一下函数名:

model_unetplusplus = smp.UnetPlusPlus(encoder_name="resnet34", encoder_weights="imagenet", in_channels=3, classes=1)
model_deeplabv3plus = smp.DeepLabV3Plus(encoder_name="resnet34", encoder_weights="imagenet", in_channels=3, classes=1)

SMP 内部已经帮你处理好了所有复杂的网络连接细节。创建好的 model 就是一个标准的 PyTorch nn.Module,你可以像使用任何自定义模型一样,查看它的结构、计算参数量,或者进行微调。

print(model)  # 打印模型结构
total_params = sum(p.numel() for p in model.parameters())
print(f"Total parameters: {total_params}")  # 计算参数量

这种设计让模型切换和对比实验变得极其方便。你可以在几分钟内,用相同的编码器,对比 U-Net、FPN、PSPNet 等不同解码器架构在你数据集上的表现,快速找到最适合的模型,而不是花几天时间去分别实现它们。

4. 数据准备与预处理:让模型“吃”好才能学好

模型架子搭好了,接下来就得准备“食材”——数据。模型训练的好坏,一半取决于数据。SMP 虽然不强制要求固定的数据格式,但它与预训练编码器的特性,要求我们对数据进行与权重训练时相同的预处理。别担心,SMP 也提供了现成的工具函数。

首先,我们需要理解一个概念:大多数预训练编码器(如 ResNet)在 ImageNet 上训练时,都对输入图像进行了标准化(Normalization)。具体来说,就是用特定的均值和标准差对每个通道的像素值进行缩放。为了让预训练特征提取器能正常工作,我们最好对输入数据施加同样的变换。

SMP 提供了一个便捷函数来获取对应编码器的预处理函数:

from segmentation_models_pytorch.encoders import get_preprocessing_fn

preprocessing_fn = get_preprocessing_fn('resnet34', pretrained='imagenet')

这个 preprocessing_fn 就是一个函数,你传入一张 RGB 图像(值范围 0-255),它会返回一个经过标准化处理(例如,转换为特定均值和标准差)的图像。这个函数通常需要和 albumentations 库配合使用,构建完整的数据处理流水线。

接下来,我们定义一个 PyTorch 的 Dataset 类。这里我以一个更通用的场景为例,假设你的数据集结构是:一个文件夹放原图(images),另一个文件夹放对应的掩码图(masks),掩码图是单通道的 PNG,像素值代表了类别索引(例如,0是背景,1是猫)。

import os
import cv2
import numpy as np
from torch.utils.data import Dataset
import albumentations as A
from albumentations.pytorch import ToTensorV2

class SegmentationDataset(Dataset):
    def __init__(self, images_dir, masks_dir, augmentation=None, preprocessing=None):
        self.images_dir = images_dir
        self.masks_dir = masks_dir
        # 获取所有图像文件名
        self.ids = [f for f in os.listdir(images_dir) if f.endswith('.png') or f.endswith('.jpg')]
        self.augmentation = augmentation
        self.preprocessing = preprocessing

    def __getitem__(self, idx):
        # 读取图像和掩码
        img_name = self.ids[idx]
        img_path = os.path.join(self.images_dir, img_name)
        mask_path = os.path.join(self.masks_dir, img_name)  # 假设同名

        image = cv2.imread(img_path)
        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  # OpenCV 读入是 BGR,转为 RGB
        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)  # 以灰度模式读取掩码

        # 确保掩码是整数类型,且类别值正确
        mask = mask.astype(np.int64)

        # 应用数据增强(训练时用,验证时可能不用)
        if self.augmentation:
            transformed = self.augmentation(image=image, mask=mask)
            image, mask = transformed['image'], transformed['mask']

        # 应用预处理(主要是标准化)
        if self.preprocessing:
            transformed = self.preprocessing(image=image, mask=mask)
            image, mask = transformed['image'], transformed['mask']
        else:
            # 如果没有提供预处理,至少将图像转为Tensor并归一化到[0,1]
            image = image.transpose(2, 0, 1).astype('float32') / 255.0
            mask = mask.astype('float32')
            image = torch.from_numpy(image)
            mask = torch.from_numpy(mask)

        return image, mask

    def __len__(self):
        return len(self.ids)

然后,我们分别定义训练和验证时的数据增强与预处理流水线。数据增强是防止过拟合、提升模型泛化能力的利器,尤其是在数据量不足时。

def get_training_augmentation():
    train_transform = [
        A.HorizontalFlip(p=0.5),
        A.VerticalFlip(p=0.5),
        A.RandomRotate90(p=0.5),
        A.ShiftScaleRotate(scale_limit=0.1, rotate_limit=15, shift_limit=0.1, p=0.5, border_mode=0),
        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),
        A.GaussNoise(p=0.2),
        A.Blur(blur_limit=3, p=0.1),
    ]
    return A.Compose(train_transform)

def get_validation_augmentation():
    """验证阶段通常只需要确保尺寸合适,不做随机增强"""
    test_transform = [
        A.PadIfNeeded(min_height=384, min_width=384, border_mode=0, value=0, mask_value=0),
        # 确保图像尺寸能被32整除(很多编码器的下采样倍数)
    ]
    return A.Compose(test_transform)

def get_preprocessing(preprocessing_fn):
    """构建包含标准化和转Tensor的预处理流水线"""
    _transform = [
        A.Lambda(image=preprocessing_fn),  # 应用标准化
        A.Lambda(image=ToTensorV2(), mask=ToTensorV2()),  # 转为PyTorch Tensor
    ]
    return A.Compose(_transform)

最后,在创建数据集时,将预处理函数传入:

preprocessing_fn = get_preprocessing_fn('resnet34', pretrained='imagenet')

train_dataset = SegmentationDataset(
    images_dir='path/to/train/images',
    masks_dir='path/to/train/masks',
    augmentation=get_training_augmentation(),
    preprocessing=get_preprocessing(preprocessing_fn),
)

val_dataset = SegmentationDataset(
    images_dir='path/to/val/images',
    masks_dir='path/to/val/masks',
    augmentation=get_validation_augmentation(), # 验证集也可以用简单的增强
    preprocessing=get_preprocessing(preprocessing_fn),
)

这样,我们就得到了一个标准化的、支持增强的数据管道。数据加载是训练流程中容易被忽视但至关重要的一环,处理好这里,能避免很多后续训练中的诡异问题。

5. 训练流程实战:从损失函数选择到模型保存

数据准备好了,模型也定义好了,激动人心的训练环节就要开始了。SMP 不仅提供了模型,还贴心地封装了一些训练工具,但核心的训练循环我们依然可以自己掌控,这样更灵活。

首先,我们需要定义损失函数和评估指标。图像分割任务常用的损失函数有 Dice Loss、BCEWithLogitsLoss、Jaccard Loss 以及它们的组合。SMP 内置了这些损失函数,直接调用即可。对于二分类任务,我通常喜欢用 DiceLoss 或 DiceBCELoss,因为它们能直接优化我们关心的 IoU(交并比)指标。

import torch.nn as nn
import segmentation_models_pytorch as smp

# 选择损失函数
loss = smp.utils.losses.DiceLoss()  # 或者 smp.utils.losses.DiceBCELoss()

# 选择评估指标
metrics = [
    smp.utils.metrics.IoU(threshold=0.5),  # IoU
    smp.utils.metrics.Fscore(threshold=0.5), # F1-score
]

接下来,设置优化器。Adam 是常用的默认选择,学习率通常设得小一点,因为我们的编码器是预训练的,微调时不宜变化太快。

import torch.optim as optim

optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)
# 使用 AdamW 并加入权重衰减,有助于防止过拟合

我们还可以使用学习率调度器,比如在训练过程中动态降低学习率。

scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=5, verbose=True)
# 当验证集指标不再提升时,降低学习率。`mode='max'` 是因为我们监控的 IoU 是越大越好。

现在,创建数据加载器,并开始编写训练循环。

from torch.utils.data import DataLoader

train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=2, pin_memory=True)
val_loader = DataLoader(val_dataset, batch_size=2, shuffle=False, num_workers=2, pin_memory=True)

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.to(device)

num_epochs = 50
best_iou = 0.0

for epoch in range(num_epochs):
    # 训练阶段
    model.train()
    train_loss = 0.0
    for images, masks in train_loader:
        images, masks = images.to(device), masks.to(device)

        optimizer.zero_grad()
        outputs = model(images)
        loss_value = loss(outputs, masks)
        loss_value.backward()
        optimizer.step()

        train_loss += loss_value.item()

    avg_train_loss = train_loss / len(train_loader)

    # 验证阶段
    model.eval()
    val_loss = 0.0
    iou_scores = []
    with torch.no_grad():
        for images, masks in val_loader:
            images, masks = images.to(device), masks.to(device)
            outputs = model(images)
            loss_value = loss(outputs, masks)
            val_loss += loss_value.item()

            # 计算 IoU (需要将输出转换为二值掩码)
            preds = (torch.sigmoid(outputs) > 0.5).float()
            iou = smp.utils.metrics.iou_score(preds, masks, threshold=0.5)
            iou_scores.append(iou.cpu().numpy())

    avg_val_loss = val_loss / len(val_loader)
    avg_val_iou = np.mean(iou_scores)

    # 学习率调度
    scheduler.step(avg_val_iou)

    print(f"Epoch {epoch+1}/{num_epochs} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f} | Val IoU: {avg_val_iou:.4f}")

    # 保存最佳模型
    if avg_val_iou > best_iou:
        best_iou = avg_val_iou
        torch.save({
            'epoch': epoch,
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'best_iou': best_iou,
        }, 'best_model.pth')
        print(f"Best model saved with IoU: {best_iou:.4f}")

这个训练循环包含了标准的训练/验证流程、损失计算、指标评估、学习率调整以及模型保存。我建议在训练初期多观察损失和指标的变化。如果训练损失下降但验证损失上升,可能是过拟合了,需要加强数据增强或调整正则化参数。如果两者都下降得很慢,可能需要检查学习率是否合适,或者数据预处理是否有问题。

模型保存时,我习惯保存一个检查点字典,里面包含模型参数、优化器状态和当前最佳指标。这样,即使训练中途中断,也可以从这个检查点恢复训练,非常方便。

6. 模型评估、预测与可视化:看看模型到底学得怎么样

训练完成后,我们得到了一个“最佳模型”。是骡子是马,得拉出来溜溜。评估和可视化不仅能让我们对模型性能有个直观认识,也是发现问题和改进方向的关键。

首先,加载我们保存的最佳模型。

checkpoint = torch.load('best_model.pth', map_location=device)
model.load_state_dict(checkpoint['model_state_dict'])
print(f"Loaded model from epoch {checkpoint['epoch']}, best IoU: {checkpoint['best_iou']:.4f}")
model.eval()  # 切换到评估模式

然后,我们可以在整个测试集上进行定量评估。除了 IoU,还可以计算精确率、召回率、F1-score 等。

from sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score
import numpy as np

all_preds = []
all_masks = []

with torch.no_grad():
    for images, masks in val_loader:  # 这里用测试集 test_loader 更合适
        images = images.to(device)
        outputs = model(images)
        preds = (torch.sigmoid(outputs) > 0.5).cpu().numpy().astype(np.uint8).flatten()
        masks = masks.cpu().numpy().astype(np.uint8).flatten()

        all_preds.extend(preds)
        all_masks.extend(masks)

all_preds = np.array(all_preds)
all_masks = np.array(all_masks)

# 计算各项指标
precision = precision_score(all_masks, all_preds)
recall = recall_score(all_masks, all_preds)
f1 = f1_score(all_masks, all_preds)
iou = smp.utils.metrics.iou_score(torch.from_numpy(all_preds), torch.from_numpy(all_masks), threshold=0.5)

print(f"Test Set Metrics:")
print(f"  Precision: {precision:.4f}")
print(f"  Recall:    {recall:.4f}")
print(f"  F1-Score:  {f1:.4f}")
print(f"  IoU:       {iou:.4f}")

定量指标很重要,但可视化更能给我们直观的感受。我们可以随机挑选几张测试图片,将原图、真实掩码和预测掩码放在一起对比。

import matplotlib.pyplot as plt

def visualize_prediction(image, true_mask, pred_mask):
    fig, axes = plt.subplots(1, 3, figsize=(15, 5))

    axes[0].imshow(image.permute(1, 2, 0).cpu().numpy())  # 原图 (C,H,W) -> (H,W,C)
    axes[0].set_title('Input Image')
    axes[0].axis('off')

    axes[1].imshow(true_mask.squeeze().cpu().numpy(), cmap='gray')  # 真实掩码
    axes[1].set_title('Ground Truth')
    axes[1].axis('off')

    axes[2].imshow(pred_mask.squeeze().cpu().numpy(), cmap='gray')  # 预测掩码
    axes[2].set_title('Prediction')
    axes[2].axis('off')

    plt.tight_layout()
    plt.show()

# 从数据集中取一个批次进行可视化
sample_images, sample_masks = next(iter(val_loader))
sample_images = sample_images.to(device)

with torch.no_grad():
    sample_preds = model(sample_images)
    sample_preds = (torch.sigmoid(sample_preds) > 0.5).float()

for i in range(min(3, len(sample_images))):  # 可视化前3张
    visualize_prediction(sample_images[i], sample_masks[i], sample_preds[i])

通过可视化,你可以清晰地看到模型在哪里分割得准,哪里出现了误判(比如把阴影当成了目标,或者漏掉了小目标)。这些观察是后续进行针对性改进(例如增加困难样本、调整损失函数权重、尝试不同模型架构)的直接依据。

最后,我们可以写一个简单的预测函数,用于对单张新图片进行分割。

def predict_single_image(image_path, model, preprocessing_fn, device):
    """对单张图像进行预测"""
    # 读取和预处理图像
    image = cv2.imread(image_path)
    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
    original_size = image.shape[:2]  # 保存原始尺寸

    # 应用预处理(这里假设预处理函数能处理单张图)
    # 通常需要将图像resize到模型输入尺寸,并做标准化
    transform = get_preprocessing(preprocessing_fn)  # 需要定义一个用于单图的预处理
    sample = transform(image=image)
    image_tensor = sample['image'].unsqueeze(0).to(device)  # 增加batch维度

    # 预测
    model.eval()
    with torch.no_grad():
        prediction = model(image_tensor)
        prediction = torch.sigmoid(prediction)
        prediction = (prediction > 0.5).float()

    # 将预测结果转换回原图尺寸
    prediction = prediction.squeeze().cpu().numpy()  # (H, W)
    prediction = cv2.resize(prediction, (original_size[1], original_size[0]), interpolation=cv2.INTER_NEAREST)

    return prediction

这个流程走下来,你就完成了一个完整的图像分割项目从搭建到评估的全过程。SMP 库在其中扮演了“加速器”的角色,让你免于重复造轮子,快速验证想法,把创造力集中在解决真正的业务挑战上。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐