简介:基于Pytorch与U-Net架构的医学图像分割完整实战项目,适合医学影像分析初学者、算法工程师及科研人员快速开展分割任务训练与推理。压缩包共99个文件,主要包括Python源码(模型构建、数据加载与训练预测)、90张PNG格式图像样本、预训练权重文件、依赖清单及一键执行脚本run.sh,整体约121.88MB,结构清晰便于直接复用。已有587人学习下载。资源提供从数据处理、模型训练到结果预测的完整流程,包含可视化的分割结果图像与README说明,内置的unet_model.pt可加载预训练权重直接测试,配套脚本降低上手门槛,可深入理解U-Net的跳跃连接、Pytorch动态计算图及医学图像标注与评估方法,适合作为算法实战参考或课程设计基础。
1. 医学图像分割用 PyTorch + U-Net:为什么这个组合是多数项目的首选
医学图像分割要处理的不是自然照片里的猫狗,而是CT里的器官、皮肤镜下的病灶、眼底图像里的血管。这类任务的共同特点是标注样本少、目标边缘模糊、类别极度不平衡,而 U-Net 的对称编码器-解码器结构和跳跃连接,恰好让网络在有限的标注下同时保留高层语义和浅层细节,因此在医学分割里几乎是默认基线。配合 PyTorch 的动态图和成熟的生态,从数据加载、模型搭建到训练和预测,都可以在几百行代码内跑通。这篇笔记围绕一个典型的 U-Net 实战项目拆解完整流程,重点放在数据准备、损失函数、一键训练脚本设计和预测阶段的踩坑记录,让新手能照着复现,也让做过几轮训练的人能回头检查自己的实现细节。
2. 医学图像分割项目的数据准备:从标注掩膜到可训练的 Dataset
2.1 项目目录设计:训练脚本、预测脚本和数据集如何拆分
一套能交付的医学分割项目,目录结构通常是这样的:
medseg_project/ ├── data/ │ ├── images/ # 原始图像,png/jpg 均可 │ └── masks/ # 标注掩膜,单通道 PNG,前景为 255 ├── checkpoints/ # 训练权重和日志输出位置 ├── network/ │ ├── __init__.py │ ├── unet_model.py # U-Net 整体结构 │ └── unet_parts.py # 编码、解码、卷积子模块 ├── utils/ │ ├── dataset.py # Dataset 和数据增强 │ ├── loss.py # 损失函数 │ └── metrics.py # Dice、IoU 等评估函数 ├── train.py # 训练入口 ├── predict.py # 预测入口 └── train.sh # 一键执行训练脚本我一般会把网络定义和数据工具严格分离。原因很直接:医学分割项目经常要换数据集、换 backbone,如果把 Dataset、损失函数和网络写在一个文件里,后面想用测试集做一次评估,就得把训练脚本从头到尾读一遍。network/只放模型结构,utils/只放数据加载和评估工具,train.py只负责编排训练流程,这样拆完之后,换数据集只需要改dataset.py,换网络只需要改unet_model.py的 import,训练脚本本身基本不动。
train.sh的作用是把环境激活、依赖安装、训练启动三道命令打包。很多初学者拿到项目后第一件事不是看代码,而是先点开这个脚本看一眼能不能跑。脚本里不写死路径,用相对路径定位项目根目录,这样项目挪到别的机器上,只要目录结构不被破坏,双击就能训练。
2.2 Dataset 类的关键写法:读图、掩膜和归一化
U-Net 训练的第一道门槛就是 Dataset 怎么写。很多人直接用torchvision.datasets.ImageFolder加载图片,但医学分割需要同时读原始图和掩膜,并且保证两者做了相同的随机变换。下面是一个可以直接改用的 Dataset 类:
import os import cv2 import torch from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size=(256, 256), transform=None): self.image_paths = sorted( [os.path.join(image_dir, f) for f in os.listdir(image_dir) if f.endswith(('.png', '.jpg', '.jpeg'))] ) self.mask_paths = sorted( [os.path.join(mask_dir, f) for f in os.listdir(mask_dir) if f.endswith('.png')] ) assert len(self.image_paths) == len(self.mask_paths), \ f"图像和掩膜数量不一致: {len(self.image_paths)} vs {len(self.mask_paths)}" self.image_size = image_size self.transform = transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image = cv2.imread(self.image_paths[idx]) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) # 统一缩放到固定尺寸 image = cv2.resize(image, self.image_size, interpolation=cv2.INTER_LINEAR) mask = cv2.resize(mask, self.image_size, interpolation=cv2.INTER_NEAREST) # 归一化到 [0, 1],注意 mask 要二值化 image = image.astype(np.float32) / 255.0 mask = (mask > 127).astype(np.float32) # 转成 CHW 张量 image = torch.from_numpy(image).permute(2, 0, 1) mask = torch.from_numpy(mask).unsqueeze(0) if self.transform is not None: # transform 需要传入 image 和 mask 的 numpy 数组,返回变换后的结果 transformed = self.transform(image=image.numpy(), mask=mask.numpy()) image = torch.from_numpy(transformed["image"].transpose(2, 0, 1)) mask = torch.from_numpy(transformed["mask"]) return image.float(), mask.float()有几个参数值得重点关注。image_size我通常设为(256, 256),因为大部分医学数据集的原始分辨率不高,256 的输入尺寸能覆盖绝大多数器官和病灶分割,显存占用也可控。mask的插值方式必须是cv2.INTER_NEAREST,不能用线性插值,否则掩膜边缘会出现介于 0 和 255 之间的过渡像素,二值化后反而制造出虚假的细线结构。mask > 127这一步看起来很基础,却是很多项目跑出 NaN 的隐蔽原因——标注文件里可能含有非 0 非 255 的灰色像素,比如标注软件留下的抗锯齿边缘。
另外要注意torch.from_numpy(image).permute(2, 0, 1)这行,OpenCV 读出来是 HWC 排列,PyTorch 的卷积层期望 CHW 排列,不转的话第一个卷积就会报维度错误。如果后续要做数据增强,transform参数建议直接用albumentations库,它天然支持 image 和 mask 同步变换,返回的也是 dict,上面的代码已经兼容这个结构。
2.3 医学数据增强:在线增强和离线增强怎么选
医学图像分割的标注成本很高,一个训练集往往只有几百甚至几十张图,这时数据增强就不再是锦上添花,而是模型能不能收敛的关键。常见做法是使用在线增强,也就是在 Dataset 的__getitem__里对加载出来的样本做随机变换,每轮 epoch 看到的都是不同的图像,相当于变相扩充数据集。
import albumentations as A train_transform = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.RandomBrightnessContrast(p=0.3), A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.05, rotate_limit=15, p=0.5), ])这几个增强操作对医学图像是安全的:翻转和旋转不改变病灶的语义属性,亮度对比度扰动模拟不同设备采集的差异。不要对医学图像做随机裁剪拼图这类增强,比如 MixUp、CutMix,因为病灶区域往往是连续且局部相关的,拼图会破坏解剖结构。也不要使用会影响掩膜语义的操作,比如A.RandomSizedCrop这类涉及缩放和裁剪的组合需要谨慎,裁剪范围如果偏离病灶中心,会让模型学到"病灶永远在图像中心"的错误先验。
离线增强指在训练前把每张图复制多份做变换后存到磁盘。这种方式便于检查增强后的样本质量,但会成倍放大磁盘占用,而且每轮 epoch 看到的增强样本是固定的,数据多样性不如在线增强。我现在的习惯是:先做一组离线增强用于快速检查数据加载逻辑,正式训练全部切换到在线增强。
3. U-Net 网络结构与损失函数:编写核心模型的最小实现
3.1 编码器与解码器:用两个卷积块搭出 U-Net 的骨架
U-Net 的最小实现并不复杂,核心组件是 DoubleConv、Down、Up 三个模块。先看unet_parts.py的定义:
import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), ) def forward(self, x): return self.conv(x) class Down(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.pool = nn.MaxPool2d(2) self.conv = DoubleConv(in_channels, out_channels) def forward(self, x): return self.conv(self.pool(x)) class Up(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() # 上采样方式:转置卷积,或者用双线性插值 + 卷积 self.up = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2) self.conv = DoubleConv(in_channels, out_channels) def forward(self, x1, x2): x1 = self.up(x1) # 如果尺寸不一致,先做边缘裁剪 diffY = x2.size()[2] - x1.size()[2] diffX = x2.size()[3] - x1.size()[3] x1 = nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) x = torch.cat([x2, x1], dim=1) return self.conv(x)DoubleConv是 U-Net 的基本单元,每组包含两个3x3卷积加 BatchNorm 加 ReLU。kernel_size=3配合padding=1保持特征图尺寸不变,只改变通道数,这是后续能拼接的前提。卷积层设置了bias=False,因为后面接了 BatchNorm,BatchNorm 自带了可学习的偏置,如果再保留卷积的 bias,参数冗余且容易产生数值不稳定。
Down先做MaxPool2d(2)把尺寸减半,再做一组 DoubleConv。Up用转置卷积把尺寸翻倍后,与编码器对应层的输出在通道维度拼接。拼接时要注意尺寸对齐:MaxPool 在偶数分辨率下不会出问题,但如果输入尺寸是奇数,转置卷积的输出和编码器输出会有 1 像素的差距,所以forward里做了边缘 pad。这个细节很多人忽略,直接torch.cat会报错,或者是靠修改输入尺寸规避,但根本做法应该是像上面这样兼容不同尺寸。
unet_model.py里把上述模块组装起来:
class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=1, base_channels=64): super().__init__() self.inc = DoubleConv(in_channels, base_channels) # 64 self.down1 = Down(base_channels, base_channels * 2) # 128 self.down2 = Down(base_channels * 2, base_channels * 4) # 256 self.down3 = Down(base_channels * 4, base_channels * 8) # 512 self.down4 = Down(base_channels * 8, base_channels * 16) # 1024 self.up1 = Up(base_channels * 16, base_channels * 8) # 512 self.up2 = Up(base_channels * 8, base_channels * 4) # 256 self.up3 = Up(base_channels * 4, base_channels * 2) # 128 self.up4 = Up(base_channels * 2, base_channels) # 64 self.outc = nn.Conv2d(base_channels, num_classes, kernel_size=1) def forward(self, x): x1 = self.inc(x) x2 = self.down1(x1) x3 = self.down2(x2) x4 = self.down3(x3) x5 = self.down4(x4) x = self.up1(x5, x4) x = self.up2(x, x3) x = self.up3(x, x2) x = self.up4(x, x1) return self.outc(x)base_channels=64是 U-Net 原文的默认值。如果显存紧张,可以改成 32,训练速度会明显提升,但分割精度通常会有小幅下降。in_channels根据图像通道数设置,灰度图设为 1,RGB 图设为 3。num_classes在二分类时设为 1,输出一个通道的概率图;多类别分割时设为类别数(不含背景)。
3.2 跳跃连接为什么在医学分割里这么关键
U-Net 和普通编码器-解码器网络最大的区别就在跳跃连接。编码器下采样 4 次后,特征图从原图尺寸缩小到 1/16,这个层能捕捉到高级语义信息,知道"这里是什么器官",但分辨率太低,无法恢复精细边界。解码器虽然能逐步放大分辨率,但仅靠高层的语义特征很难还原细节。跳跃连接把编码器各层的浅层特征直接拼到解码器对应层,相当于给解码器提供了多尺度的边缘和纹理信息。
拼接和相加是两种主流做法。原始 U-Net 用的是按通道拼接(torch.cat),通道数翻倍,解码器卷积会自动融合这些信息。ResUNet 等变体用相加(torch.add),参数量更少,但相加要求两个特征图的通道数严格一致,灵活性差一些。我在分割任务里默认用拼接,因为医疗数据里的病灶边缘信息非常宝贵,拼接比相加保留了更多的独立通道表达。
需要注意,跳跃连接会引入一个实际工程问题:输入尺寸必须是 16 的倍数。因为网络做了 4 次下采样,每次缩小一半,如果输入尺寸不能整除 16,跳跃连接拼接时就要像前面代码那样做 pad 或 crop。最稳妥的做法是在数据预处理阶段就把所有训练图像统一 resize 到 256×256 或 512×512,让这个问题从根源上消失。
3.3 输出层和损失函数:二分类还是多分类
输出层本身只有一行nn.Conv2d(base_channels, num_classes, kernel_size=1),真正的坑在损失函数的选择。二分类分割里,很多初学者直接用nn.CrossEntropyLoss,但 CrossEntropyLoss 期望的输入是(N, C, H, W)的 logits,且类别通道数至少为 2。如果网络最后只输出 1 个通道,这里就会报维度错误。
二分类分割的标准做法是:输出 1 个通道的 logits,配合nn.BCEWithLogitsLoss,它内部已经把 sigmoid 和交叉熵合并,数值稳定性比手动sigmoid + BCELoss更好。多类别分割则输出 C 个通道,配合nn.CrossEntropyLoss,每个像素的类别由 argmax 决定。
但纯交叉熵在医学分割里有一个突出问题:背景像素占比远大于前景。一张 256×256 的图像里,病灶可能只占 5% 的像素,模型只要把所有像素预测为背景,loss 就已经很低。所以实践中几乎都会引入 Dice Loss 或其变体作为辅助损失:
import torch import torch.nn as nn import torch.nn.functional as F class DiceLoss(nn.Module): def __init__(self, smooth=1.0): super().__init__() self.smooth = smooth def forward(self, logits, targets): probs = torch.sigmoid(logits) # 拉平到 [N, -1] probs = probs.reshape(probs.size(0), -1) targets = targets.reshape(targets.size(0), -1) intersection = (probs * targets).sum(dim=1) union = probs.sum(dim=1) + targets.sum(dim=1) dice = (2.0 * intersection + self.smooth) / (union + self.smooth) return 1.0 - dice.mean()Dice Loss 直接优化像素集合的重叠程度,对前景占比不敏感,非常适合医学分割。smooth是平滑项,防止分子分母同时为 0,一般取 1.0。实际使用中,我习惯把 BCE 和 Dice 加权求和:total_loss = 0.5 * bce_loss + 0.5 * dice_loss。BCE 负责逐像素的分类准确性,Dice 负责整体区域的重叠质量,两者互补。如果只使用 Dice Loss,训练早期梯度会很不稳定,因为 sigmoid 输出的概率与真实掩膜的重叠从 0 开始计算,梯度方向容易抖动。
4. 一键训练脚本设计:让训练与预测跑起来的参数与流程
4.1 训练主循环的最小实现
train.py是项目的核心入口。下面是一个能直接跑通二分类分割训练的循环框架:
import torch import torch.optim as optim from torch.utils.data import DataLoader from network.unet_model import UNet from utils.dataset import SegmentationDataset from utils.loss import DiceLoss def train_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss = 0.0 for images, masks in dataloader: images = images.to(device) masks = masks.to(device) optimizer.zero_grad() logits = model(images) loss = criterion(logits, masks) loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) return total_loss / len(dataloader.dataset) def validate(model, dataloader, device): model.eval() dice_score = 0.0 with torch.no_grad(): for images, masks in dataloader: images = images.to(device) masks = masks.to(device) logits = model(images) probs = torch.sigmoid(logits) preds = (probs > 0.5).float() intersection = (preds * masks).sum(dim=(1, 2, 3)) union = preds.sum(dim=(1, 2, 3)) + masks.sum(dim=(1, 2, 3)) dice_score += (2.0 * intersection / (union + 1e-6)).sum().item() return dice_score / len(dataloader.dataset) def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") train_dataset = SegmentationDataset("data/images", "data/masks", image_size=(256, 256)) val_dataset = SegmentationDataset("data/val_images", "data/val_masks", image_size=(256, 256)) train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=4) val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=4) model = UNet(in_channels=3, num_classes=1).to(device) optimizer = optim.Adam(model.parameters(), lr=1e-3) criterion = lambda logits, targets: 0.5 * torch.nn.functional.binary_cross_entropy_with_logits( logits, targets ) + 0.5 * DiceLoss()(logits, targets) best_dice = 0.0 for epoch in range(100): train_loss = train_epoch(model, train_loader, optimizer, criterion, device) val_dice = validate(model, val_loader, device) print(f"Epoch {epoch+1:03d} | Train Loss: {train_loss:.4f} | Val Dice: {val_dice:.4f}") if val_dice > best_dice: best_dice = val_dice torch.save(model.state_dict(), "checkpoints/best_model.pth") if __name__ == "__main__": main()dataloader的num_workers=4是用多进程预取图像数据,避免 GPU 在等待 CPU 读取图片时空转。如果机器核心数少,可以降到 2;Windows 环境下建议设为 0,否则多进程可能因为 spawn 机制报错。batch_size=8对应 256×256 输入在 8GB 显存上的典型值,如果你的卡只有 4GB,先降到 4 或 2。
保存权重的策略是维护一个best_dice变量,只在验证集 Dice 分数创新高时保存。这比每轮都保存要省心得多。训练的终止条件我一般设两个:一是循环跑满max_epochs;二是在脚本里加一个早停计数器,如果连续 15 个 epoch 验证 Dice 都没有提升,就提前 break。早停能省大量时间,尤其是医学数据集的验证集波动经常很大,模型可能在第 30 个 epoch 就过拟合了。
4.2 关键训练超参数:学习率、batch size 和输入尺寸怎么定
U-Net 训练的超参数相互牵连,直接抄一个固定配置很容易翻车。下面这张表是我在多个医学分割项目里反复调整后的基准值:
| 参数 | 推荐值 | 说明 |
|---|---|---|
| optimizer | Adam | 默认 betas=(0.9, 0.999),上手快 |
| 初始学习率 | 1e-3 | 显存小导致 batch 小,建议降到 5e-4 |
| batch size | 8 | 8GB 显存 + 256×256 输入的极限值 |
| 输入尺寸 | 256×256 | 分辨率不够再上 512,但显存占用翻 4 倍 |
| max_epochs | 100 | 配合早停,一般 50~80 轮能收敛 |
| 早停 patience | 15 | 验证集 Dice 连续 15 轮不提升就停 |
| 学习率调度 | ReduceLROnPlateau | factor=0.5, patience=5 |
学习率是最敏感的参数。Adam 的默认学习率 1e-3 在大部分情况下能直接收敛,但如果 batch size 降到 4 以内,梯度估计的噪声变大,1e-3 可能让 loss 剧烈震荡,此时把学习率降到 5e-4 或 3e-4 更稳妥。反过来,如果用了迁移学习(比如 encoder 部分加载预训练权重),学习率应该整体降到 1e-4,否则微调阶段很容易破坏已经学好的特征。
ReduceLROnPlateau是一个非常实用的调度器:验证集 Dice 连续 5 轮不上升,学习率就乘以 0.5。它的好处是完全不需要预先设定在第几个 epoch 衰减,因为医学数据集的收敛曲线并不平滑,有时候模型会在某个平台期"沉默"十来个 epoch,然后突然继续提升。如果硬编码 StepLR 在固定轮次衰减,很容易错过这个二次上升的窗口。
4.3 一键脚本如何设计:让新环境也能直接开始训练
train.sh是项目里"一键训练"的入口。它的目标不是把训练逻辑藏起来,而是把环境准备和训练启动这两件事自动化,让拿到项目的人不用读 README 就能跑起来:
#!/bin/bash set -e cd "$(dirname "$0")" # 检查虚拟环境,不存在则创建 if [ ! -d "venv" ]; then python3 -m venv venv fi source venv/bin/activate # 安装依赖,requirements.txt 里固定了 torch、torchvision、opencv-python、albumentations pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple # 创建必要的目录 mkdir -p data/images data/masks checkpoints # 启动训练 python train.py --epochs 100 --batch_size 8 --image_size 256脚本开头的set -e表示任何一条命令执行失败就立即退出,避免 pip 安装失败后继续跑训练,最后报一堆看不懂的 traceback。cd "$(dirname "$0")"是关键,它让脚本从自身所在的目录运行,不管用户从哪里调用,都能定位到项目根目录。
Windows 用户可以把这个逻辑做成train.bat:
@echo off cd /d %~dp0 if not exist venv (python -m venv venv) call venv\Scripts\activate.bat pip install -r requirements.txt python train.py --epochs 100 --batch_size 8 --image_size 256 pausetrain.py里的参数全部通过argparse接收,这样一键脚本只是传入了最常用的默认值,用户想改参数时不需要翻代码,直接改脚本里的命令行参数就行。比如想用 GPU 训练但显存不够,可以改成--batch_size 4 --image_size 192,其他部分完全不用动。
5. 训练与预测中的常见坑:现象、原因和解决方法
5.1 训练 loss 在降,但验证集 Dice 一直很低
现象:训练集损失从 0.7 稳步降到 0.2,但验证集的 Dice 分数始终在 0.1 到 0.2 之间徘徊,甚至不升反降。
原因:最常见的是训练集和验证集的数据分布不一致。医学数据的采集设备和标注标准在不同机构间差异极大,如果训练集来自 A 医院的设备,验证集来自 B 医院的公开数据集,模型学到的纹理特征在验证集上完全不适用。另一个隐蔽的原因是训练时做了数据增强,而验证集没有做任何预处理对齐,导致输入分布不同。
解决:先把训练集和验证集按来源做一次分层划分,确保两个集合包含相似比例的样本来源。其次,把数据预处理逻辑抽成一个公共函数,训练和验证都调用同一个函数,不要各写各的。最后,在训练脚本里加一个诊断输出,每 5 个 epoch 把验证集的预测结果和原图、掩膜拼在一起保存成一张图,用肉眼确认模型到底学到了什么。
5.2 预测结果全是黑图或全白图
现象:模型训练时 Dice 表现正常,但单独跑预测脚本时,输出的分割图全黑或全白。
原因:预测阶段缺少 sigmoid 或阈值化操作。训练时 BCEWithLogitsLoss 内部已经做了 sigmoid 计算,所以训练不需要额外处理;但预测阶段拿到的是 logits,必须手动torch.sigmoid(logits)后再与 0.5 比较。另一个原因是预测时输入图像的预处理顺序和训练不一致,比如训练时做了image / 255.0归一化,预测时直接喂原始像素值,模型看到的是完全不同的数值分布。
解决:在预测脚本里,先加载训练时用的同一个归一化函数处理图像,再做前向推理,最后对输出做 sigmoid 和阈值化。建议把训练和预测公用的预处理函数单独放到utils/preprocess.py,两边都从这个模块导入,避免出现一个改了另一个没改的尴尬。
5.3 显存溢出:batch size 和输入尺寸怎么调
现象:训练启动几秒后报CUDA out of memory,或者训练到中途突然显存不足退出。
原因:输入尺寸和 batch size 的组合超出了显存容量。256×256 的输入、batch size 8、base_channels 64 的 U-Net,占用大约 6~7GB 显存,但如果把输入改成 512×512,显存占用直接翻 4 倍,8GB 的卡必然溢出。另一个隐蔽原因是 PyTorch 默认会缓存显存分配器,即使显存显示占满,可能只是没有及时释放碎片。
解决:优先减小 batch size,从 8 降到 4 再降到 2。如果 batch size 已经到 1 还不够,再考虑减小输入尺寸。梯度累积是一个较好的折中方案:
accumulation_steps = 4 optimizer.zero_grad() for step, (images, masks) in enumerate(dataloader): loss = criterion(model(images.to(device)), masks.to(device)) loss = loss / accumulation_steps loss.backward() if (step + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()这里把一次大步长更新拆成 4 个小 batch 累计梯度,等价于 batch size 从 2 扩大到 8,但显存占用只相当于 batch size 2。需要注意,loss要除以accumulation_steps,否则累计梯度会放大 4 倍,导致学习率实际上被放大了,训练容易发散。
5.4 分割边界粗糙、小目标丢失
现象:大块器官分割效果不错,但细小血管、小型病灶的边缘参差不齐,或者小目标完全没有被检测出来。
原因:深层特征的感受野大,擅长识别物体类型,但空间细节丢失严重。跳跃连接虽然能恢复一部分细节,但如果病灶只占整张图像的 2% 以下,BCE 损失对它的贡献极小,模型倾向于把注意力放在背景和大目标上。此外,下采样次数固定为 4 时,1/16 分辨率的特征图上,原本只有 8×8 像素的小病灶可能已经退化成 1~2 个像素点,解码器很难恢复。
解决:损失函数换成 BCE + Dice 的组合,Dice 对前景占比不敏感,能强制模型关注小目标。输入尺寸从 256 提升到 512,相当于让小病灶在特征图中的有效像素翻倍。数据增强里加入小幅度的随机缩放,模拟不同大小病灶的出现。如果项目允许,还可以尝试把深层的base_channels从 64 减小到 32,减少下采样后的信息瓶颈。
5.5 训练时数据增强和预测时预处理不一致
现象:训练时验证集 Dice 达到 0.85,但部署到实际数据上效果很差,边界错位、整体偏移。
原因:训练脚本里的数据增强管道包含随机翻转、旋转、缩放,这些操作在训练时作用于图像和掩膜;但预测脚本只对单张图做 resize 和归一化,没有做任何后处理对齐。更严重的是,如果训练时用albumentations的ShiftScaleRotate做了缩放,模型可能对原始分辨率的图像分布不敏感,一旦预测图不经过同样的归一化,输入分布直接错位。
解决:训练和预测的前处理必须走同一条代码路径。我把归一化、resize 写成一个preprocess_image(image_path, image_size)函数,训练和预测都调用它。数据增强只应用在训练集的 Dataset 内部,预测阶段不调用增强,但预处理函数保持完全一致。这个坑排查起来最费时间,因为模型权重、损失函数、学习率全都看着正常,问题纯粹出在数据流水线的两端不对齐。
6. 预测脚本的最佳实践:从权重到分割图
6.1 预测时加载模型与预处理的一致性
预测脚本比训练脚本更考验工程细节,因为训练过程有验证集实时反馈,出错了能立刻看到;预测时面对的是全新数据,错了可能要到下游分析阶段才暴露。下面是我常用的一个预测脚本核心片段:
import torch import cv2 import numpy as np from network.unet_model import UNet def load_model(weight_path, device, in_channels=3, num_classes=1): model = UNet(in_channels=in_channels, num_classes=num_classes) state_dict = torch.load(weight_path, map_location=device) model.load_state_dict(state_dict) model.to(device) model.eval() return model def predict_image(model, image_path, device, image_size=(256, 256)): # 与训练保持相同的预处理顺序 image = cv2.imread(image_path) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) if image.shape[:2] != image_size: image = cv2.resize(image, image_size, interpolation=cv2.INTER_LINEAR) image = image.astype(np.float32) / 255.0 # 转成 CHW,增加 batch 维度 input_tensor = torch.from_numpy(image).permute(2, 0, 1).unsqueeze(0).to(device) with torch.no_grad(): logits = model(input_tensor) probs = torch.sigmoid(logits) pred = (probs > 0.5).float() # 压缩 batch 和 channel 维度,变成 HxW return pred.squeeze(0).squeeze(0).cpu().numpy() def save_result(pred, save_path): # 二值掩膜转成可保存的 0/255 PNG mask = (pred * 255).astype(np.uint8) cv2.imwrite(save_path, mask)map_location=device这行很关键。在 GPU 上训练的权重文件里含有 CUDA 张量,如果目标机器没有 GPU,直接torch.load会报错,加上map_location就能把权重加载到 CPU,再用to(device)迁移。squeeze(0).squeeze(0)是把(1, 1, H, W)的预测结果还原成(H, W)的二维掩膜。
预测阶段唯一需要注意的输出格式是掩膜保存。医学分割结果通常要求保存为 0 和 255 的单通道 PNG,且尺寸与原图一致。如果预测时做了缩放,还要在保存前把掩膜反缩放到原图尺寸,这里同样用cv2.INTER_NEAREST。
6.2 后处理和验证:连通域过滤与连通域评估
预测出的二值掩膜通常会包含一些零散的噪声小区域。对于医学场景,真实病灶往往是连通的实体,孤立的像素块大概率是误检。常见的后处理是找出所有连通域,滤除面积小于阈值的区域:
import cv2 import numpy as np def remove_small_components(mask, min_area=50): num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8) filtered = np.zeros_like(mask) for label in range(1, num_labels): if stats[label, cv2.CC_STAT_AREA] >= min_area: filtered[labels == label] = 255 return filteredmin_area要根据实际病灶大小标定。皮肤病变分割里,小于 50 像素的区域基本是噪声;血管分割里细小分支可能本身就不足 50 像素,需要降低到 10 或 20。这个参数可以在验证集上统计真实掩膜的连通域面积分布后再定。
最终评估时,除了 Dice 分数,建议加上 IoU 和 Hausdorff 距离两个指标。Dice 和 IoU 衡量区域重叠比例,对整体性能敏感;Hausdorff 距离衡量预测边界与真实边界的最大偏差,能反映边缘质量。如果 Dice 高但 Hausdorff 距离大,说明预测区域整体正确但边缘毛刺严重,此时需要重点检查后处理步骤和损失函数里的边界约束。
我自己的习惯是,每次训练完跑完验证集评估后,把预测结果按"原图、真实掩膜、预测掩膜、叠加图"四联图保存到一个目录里,随机挑 20 张做人工检查。这个习惯帮我发现过好几次问题——有一次模型整体 Dice 到了 0.9,但叠加图里病灶边缘比真实位置整体向外扩了 2~3 个像素,只靠数值指标完全看不出来。希望这套流程和踩坑记录能帮你在自己的分割项目里少走几步弯路。
本文还有配套的精品资源,点击获取