图像分割是计算机视觉里比目标检测更细一档的任务。目标检测给的是“图中哪里有一个物体”,图像分割给的是“每个像素属于哪一类”,连边界都给你画出来。这次我们要看的 U-Net,就是做这件事最常用的基础架构之一。它最早用于医学图像分割,后来在遥感、工业质检、路网提取这些场景里也大量出现。重点是,它不像一些大模型那样依赖海量数据和高端显卡,即使是在一张普通 GPU 甚至 CPU 上也能完成训练和推理,所以特别适合作为 PyTorch 实战入门的第一个分割模型。
这篇文章会完整走一遍 U-Net 图像分割的实战流程:从环境安装、数据集准备、模型搭建,到训练验证、指标评估、推理调用和常见报错排查。如果你正在学习 PyTorch,或者想把手上的图像分割需求落地成代码,这篇可以直接收藏照着做。
本文适合这样的读者:有一定 Python 基础,看过 PyTorch 的 Tensor 和 Dataset 基本概念,但还没有完整训练过一个分割模型;或者已经在用目标检测,想进一步做像素级分割。文章里给出的都是可运行的最小实现,不依赖额外的专有框架,读懂代码就能迁移到自己任务里。
1. 核心能力速览
| 能力项 | 说明 |
|---|---|
| 项目类型 | 深度学习图像分割实战教程 |
| 网络架构 | U-Net:编码器-解码器结构 + 跳跃连接 |
| 主要功能 | 二分类分割、多类别语义分割、高分辨率掩码预测 |
| 训练框架 | PyTorch |
| 推荐硬件 | NVIDIA GPU,显存越大能处理的输入尺寸越大;CPU 可训练但速度慢 |
| 启动方式 | Python 脚本训练,训练完成后脚本推理或封装 API |
| API 能力 | 模型训练完成后可自行封装 HTTP 推理接口 |
| 批量任务 | 支持目录批量预测,可接入循环脚本 |
| 适合场景 | 医学图像分割、遥感目标提取、工业缺陷分割、路网/建筑提取等 |
这里先把结论说清楚:U-Net 不是某个公司开源的商业产品,而是一个 2015 年提出的经典卷积神经网络架构,它的核心设计是一左一右两条路径,左边不断下采样提取特征,右边不断上采样恢复分辨率,中间用跳跃连接把同尺度的细节信息拼起来。正因为跳跃连接的存在,U-Net 对边缘细节的还原能力比普通编码器-解码器网络好很多,也更适合小样本数据。
2. 适用场景与使用边界
U-Net 最典型的应用场景是医学图像分割,比如细胞分割、器官区域提取、病灶区域勾画。因为医学影像数据往往数量有限、标注昂贵,而 U-Net 在少量数据上也能训练出不错的效果。除了医学,遥感领域也常用 U-Net 做建筑物提取、道路分割、水体识别;工业界则用它圈出产品表面的缺陷区域;广告牌监测、道路监控画面里的特定目标区域提取,也属于这类思路。
不适用场景也要说清楚。U-Net 是逐像素全图计算,推理速度比目标检测慢;如果需要实时处理高分辨率视频流,建议先做区域裁剪或者换轻量分割网络。另外,U-Net 本身对输入尺寸比较敏感,如果直接输入超大原图,显存占用会迅速上涨,常规做法是裁剪成 patch 训练。
使用边界方面,必须强调三点。第一,涉及医学影像时,模型结果只能作为科研或辅助参考,不能直接作为临床诊断依据。第二,数据集里的图像和标注必须来源合法、授权明确,尤其是人脸、车牌、医疗影像这类敏感数据,训练前要做好脱敏。第三,分割模型输出的边界不一定是真实边界,落地到质检、测量等场景时,需要人工审核或额外后处理。
3. 环境准备与前置条件
在写代码之前,先把环境装好。这里给出的是通用流程,实际版本以本机情况为准。
3.1 创建 Python 环境
推荐用 Anaconda 管理环境,避免把系统 Python 弄乱。在命令行执行:
conda create -n unet python=3.9 conda activate unetPython 版本建议选择 3.8 到 3.10。如果你已经有 Conda 环境,也可以直接复用。
3.2 安装 PyTorch
PyTorch 安装的关键是 CUDA 版本匹配。先从命令行执行nvidia-smi查看驱动支持的 CUDA 版本,再到 PyTorch 官网选择对应命令。这里以 CUDA 11.8 和 PyTorch 2.x 为例:
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118如果机器没有 NVIDIA GPU,或者只是想先跑通代码,可以安装 CPU 版本:
pip install torch torchvisionCPU 版本能训练小数据集,只是速度会慢不少。安装完成后,验证一下:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"torch.cuda.is_available()输出True,说明 GPU 可用;输出False,说明当前是 CPU 环境或者 CUDA 配置有问题,后面排错章节会展开。
3.3 安装依赖库
除了 PyTorch,还需要一些图像处理和科学计算库:
pip install numpy opencv-python pillow matplotlib tqdm如果后面要做接口服务,再加上:
pip install fastapi uvicorn python-multipart到这里环境就准备好了。磁盘方面,代码本身很小,但数据集图片和模型权重需要空间,建议预留至少 10 到 20 GB,具体看数据量。
4. 数据集准备
图像分割的数据集核心是“原图 + 掩码”配对。掩码就是一张和原图尺寸相同的单通道图,每个像素的灰度值代表该像素的类别编号。二分类任务里,前景像素是 1,背景像素是 0;多分类任务里,每个类别对应一个整数值。
4.1 推荐目录结构
建议把数据整理成下面这种结构:
dataset/ ├── train_images/ │ ├── 001.jpg │ ├── 002.jpg │ └── ... ├── train_masks/ │ ├── 001.png │ ├── 002.png │ └── ... ├── val_images/ │ ├── 001.jpg │ └── ... └── val_masks/ ├── 001.png └── ...文件名一一对应,原图和掩码名称保持一致,代码里会按名字拼接路径。这里有一个非常重要的工程细节:训练集和验证集要分开,而且验证集不能参与训练,否则指标虚高,模型实际效果会大打折扣。
4.2 自定义 Dataset 读取
在 PyTorch 里,数据读取需要继承torch.utils.data.Dataset。最简单的实现如下:
import os from PIL import Image import torch from torch.utils.data import Dataset from torchvision import transforms class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size=(256, 256)): self.image_dir = image_dir self.mask_dir = mask_dir self.image_size = image_size self.image_names = sorted(os.listdir(image_dir)) self.image_transform = transforms.Compose([ transforms.Resize(image_size), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) self.mask_transform = transforms.Compose([ transforms.Resize(image_size, interpolation=Image.NEAREST), transforms.ToTensor() ]) def __len__(self): return len(self.image_names) def __getitem__(self, idx): img_name = self.image_names[idx] img_path = os.path.join(self.image_dir, img_name) mask_path = os.path.join(self.mask_dir, img_name.replace('.jpg', '.png')) image = Image.open(img_path).convert('RGB') mask = Image.open(mask_path).convert('L') image = self.image_transform(image) mask = self.mask_transform(mask) # 掩码像素除以 255,统一到 [0, 1] 或 [0, num_classes-1] mask = (mask * 255).long().squeeze(0) return image, mask这里有几个细节值得注意。第一,掩码用Image.NEAREST最近邻插值缩放,不能用双线性插值,否则类别边界会出现介于两个整数之间的模糊值。第二,掩码要squeeze(0)去掉通道维度,变成[H, W],训练时和模型输出做损失计算更方便。第三,图片做 Normalize 归一化,掩码不能归一化。
如果你的数据是彩色掩码(可视化用的 RGB 标注),需要先做一个颜色到类别 ID 的映射表,再转成单通道。这个转换通常写在数据预处理脚本里,不要在 Dataset 里重复做。
5. U-Net 模型搭建
5.1 网络结构回顾
U-Net 的结构可以拆成三部分:
- 收缩路径(编码器):重复执行两次 3x3 卷积 + ReLU,再接 2x2 最大池化下采样。每下采样一次,特征图尺寸减半,通道数翻倍。
- 扩展路径(解码器):先做一次上采样,让特征图尺寸翻倍,然后把编码器对应层的特征图拼接上来,再做两次 3x3 卷积 + ReLU。
- 输出层:1x1 卷积,把通道数映射成类别数。
跳跃连接是整个网络的精髓:它把浅层的高分辨率细节和深层的语义特征拼在一起,解决了普通网络深层次特征丢失细节的问题。
5.2 完整 PyTorch 实现
下面是适合入门学习和二分类/多分类任务的标准 U-Net 实现,代码里加了基础注释:
import torch import torch.nn as nn class DoubleConv(nn.Module): """两次卷积 + 批归一化 + ReLU""" def __init__(self, in_channels, out_channels): super(DoubleConv, self).__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): """下采样:最大池化 + DoubleConv""" def __init__(self, in_channels, out_channels): super(Down, self).__init__() self.down = nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_channels, out_channels) ) def forward(self, x): return self.down(x) class Up(nn.Module): """上采样:转置卷积 + 跳跃连接拼接 + DoubleConv""" def __init__(self, in_channels, out_channels): super(Up, self).__init__() self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, 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) class UNet(nn.Module): def __init__(self, in_channels=3, num_classes=2, base_channels=64): super(UNet, self).__init__() self.inc = DoubleConv(in_channels, base_channels) self.down1 = Down(base_channels, base_channels * 2) self.down2 = Down(base_channels * 2, base_channels * 4) self.down3 = Down(base_channels * 4, base_channels * 8) self.down4 = Down(base_channels * 8, base_channels * 16) self.up1 = Up(base_channels * 16, base_channels * 8) self.up2 = Up(base_channels * 8, base_channels * 4) self.up3 = Up(base_channels * 4, base_channels * 2) self.up4 = Up(base_channels * 2, base_channels) 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) logits = self.outc(x) return logits这段代码是标准 U-Net 的通用实现。num_classes=2表示二分类分割,num_classes=5就表示分割 5 个类别。这里的base_channels是网络基础通道数,默认 64;显存紧张时可以改成 32 或 16,模型体积和显存占用都会明显下降。
创建模型并查看参数量:
model = UNet(in_channels=3, num_classes=2) print(f"参数量: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M")标准 U-Net 默认参数量在 3100 万左右,也就是约 31M,模型文件大小约 120 MB,不算大。如果是显存非常小的设备,可以把base_channels调小,或者输入图改成128x128。
6. 训练流程与评估指标
6.1 损失函数选择
分割任务最常用的损失函数有两个维度:像素级损失和区域级损失。
- 二分类:
BCEWithLogitsLoss,配合sigmoid输出。 - 多分类:
CrossEntropyLoss,配合softmax输出。
if num_classes == 2: criterion = nn.BCEWithLogitsLoss() else: criterion = nn.CrossEntropyLoss()也可以用 Dice Loss 或结合 CE Loss,对小目标区域更友好。初次跑通时建议先用 CrossEntropyLoss,稳定后再考虑复杂损失。
6.2 评估指标:IoU 和 Dice
图像分割里最常用的指标是 IoU(交并比)和 Dice 系数。它们的核心都是计算预测区域和真实区域的面积重叠程度,取值越接近 1 越好。
def compute_iou(pred_mask, true_mask, num_classes=2): ious = [] for cls in range(num_classes): pred = (pred_mask == cls) true = (true_mask == cls) intersection = (pred & true).sum().item() union = (pred | true).sum().item() if union == 0: continue ious.append(intersection / union) if len(ious) == 0: return 0.0 return sum(ious) / len(ious)注意一个常见问题:背景类占比很大,如果每个类都计入 IoU,背景会拉高整体分数,掩盖前景目标效果差的问题。工程上更常用 mIoU(mean IoU),即每个类单独算 IoU 再取平均,本文上面的函数就是 mIoU 的计算思路。
6.3 训练主脚本
把数据加载、模型、优化器、训练循环串起来。下面是完整的训练脚本框架:
import torch import torch.optim as optim from torch.utils.data import DataLoader from torchvision import transforms from tqdm import tqdm # 超参数 IMAGE_SIZE = 256 BATCH_SIZE = 8 EPOCHS = 50 LEARNING_RATE = 1e-4 NUM_CLASSES = 2 DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 数据 train_dataset = SegmentationDataset( image_dir="dataset/train_images", mask_dir="dataset/train_masks", image_size=(IMAGE_SIZE, IMAGE_SIZE) ) val_dataset = SegmentationDataset( image_dir="dataset/val_images", mask_dir="dataset/val_masks", image_size=(IMAGE_SIZE, IMAGE_SIZE) ) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2) val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2) # 模型 model = UNet(in_channels=3, num_classes=NUM_CLASSES).to(DEVICE) if NUM_CLASSES == 2: criterion = nn.BCEWithLogitsLoss() else: criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE) # 训练循环 best_iou = 0.0 for epoch in range(EPOCHS): model.train() train_loss = 0.0 for images, masks in tqdm(train_loader, desc=f"Epoch {epoch+1}/{EPOCHS}"): images = images.to(DEVICE) masks = masks.to(DEVICE) outputs = model(images) if NUM_CLASSES == 2: loss = criterion(outputs.squeeze(1), masks.float()) else: loss = criterion(outputs, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() * images.size(0) # 验证 model.eval() val_iou = 0.0 val_count = 0 with torch.no_grad(): for images, masks in val_loader: images = images.to(DEVICE) masks = masks.to(DEVICE) outputs = model(images) if NUM_CLASSES == 2: preds = (torch.sigmoid(outputs) > 0.5).long().squeeze(1) else: preds = torch.argmax(outputs, dim=1) for i in range(images.size(0)): val_iou += compute_iou(preds[i], masks[i], NUM_CLASSES) val_count += 1 avg_train_loss = train_loss / len(train_dataset) avg_val_iou = val_iou / val_count print(f"Epoch {epoch+1}: Loss = {avg_train_loss:.4f}, mIoU = {avg_val_iou:.4f}") if avg_val_iou > best_iou: best_iou = avg_val_iou torch.save(model.state_dict(), "best_unet.pth") print("保存最佳模型 best_unet.pth")训练中几个关键点:
- 显存不够时优先减小
BATCH_SIZE,其次减小IMAGE_SIZE。 - 学习率用
1e-4起步,如果损失震荡,降到1e-5。 - 训练过程里只保存验证集 mIoU 最高的权重,避免末轮过拟合导致模型变差。
num_workers先从 0 或 2 开始,调太高在 Windows 上容易出现多进程报错。
6.4 数据增强
图像分割的数据增强有个特殊要求:图像和掩码必须做完全相同的变换。不能只旋转原图不旋转掩码,也不能给掩码做颜色抖动。
class RandomHorizontalFlip: def __call__(self, image, mask): if torch.rand(1) > 0.5: image = torch.flip(image, dims=[2]) mask = torch.flip(mask, dims=[1]) return image, mask更完整的增强可以引入albumentations库,它对分割任务做了专门处理,可以同时保证图像和掩码同步变换。建议至少加入水平翻转、随机旋转、缩放三种增强方式。
7. 推理与结果验证
训练完成后,写一个简单的推理脚本。这里用单张图片做测试,输出叠加可视化结果,并且把预测掩码单独保存成图片。
import torch import numpy as np import cv2 from PIL import Image from torchvision import transforms DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") def predict_single_image(model, image_path, device=DEVICE): image = Image.open(image_path).convert('RGB') original_size = image.size transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) input_tensor = transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): output = model(input_tensor) if output.shape[1] == 2: prob = torch.sigmoid(output) pred = (prob > 0.5).long().squeeze(0).squeeze(0).cpu().numpy() else: pred = torch.argmax(output, dim=1).squeeze(0).cpu().numpy() pred_mask = Image.fromarray((pred * 255).astype(np.uint8)) pred_mask = pred_mask.resize(original_size, Image.NEAREST) return np.array(pred_mask), image model = UNet(in_channels=3, num_classes=2) model.load_state_dict(torch.load("best_unet.pth", map_location=DEVICE)) model.to(DEVICE) mask, original_img = predict_single_image(model, "test.jpg") # 保存结果 Image.fromarray(mask).save("test_mask.png") # 可视化:把掩码叠加到原图上 overlay = np.array(original_img).copy() overlay[mask > 0] = (0, 255, 0) # 绿色标记前景 cv2.imwrite("test_overlay.jpg", overlay)判断推理成功的标准很直观:
- 保存的
test_mask.png中,前景区域是白色,背景是黑色,目标轮廓清晰。 - 叠加图上,目标区域被准确覆盖,没有大面积漏检或误检。
- 边缘位置有一定误差是正常的,但如果整张掩码全是黑的或全是白的,要检查模型权重路径、图像预处理和数据标注。
8. 接口 API 与批量任务
8.1 批量预测
批量预测的核心是遍历目录里的所有图片,逐张调用上面的推理函数。这里给一个通用脚本:
import os from tqdm import tqdm os.makedirs("outputs", exist_ok=True) for image_name in tqdm(os.listdir("test_images")): image_path = os.path.join("test_images", image_name) mask, original_img = predict_single_image(model, image_path) output_name = image_name.replace('.jpg', '_mask.png') output_path = os.path.join("outputs", output_name) Image.fromarray(mask).save(output_path)批量任务一定要做好失败重试和日志记录。建议每处理一张图片就打印或记录文件名、推理耗时、输出路径,处理失败的图片单独放到failed目录,不要中断整个循环。
8.2 封装 FastAPI 接口
模型训练完成后,用 FastAPI 封装成推理接口,方便接到业务系统里。下面是一个通用模板,实际路径和参数按项目调整:
from fastapi import FastAPI, UploadFile, File from PIL import Image import io import numpy as np app = FastAPI() model = UNet(in_channels=3, num_classes=2) model.load_state_dict(torch.load("best_unet.pth", map_location=DEVICE)) model.to(DEVICE) model.eval() @app.post("/predict") async def predict(file: UploadFile = File(...)): image_bytes = await file.read() image = Image.open(io.BytesIO(image_bytes)).convert("RGB") # 复用单张推理函数 mask, _ = predict_single_image(model, image) # 返回前端可渲染的 PNG 二进制 mask_img = Image.fromarray(mask) buf = io.BytesIO() mask_img.save(buf, format="PNG") buf.seek(0) return {"mask": buf.getvalue().hex()}启动接口服务:
uvicorn server:app --host 127.0.0.1 --port 8000调用方式用 Python 请求:
import requests response = requests.post( "http://127.0.0.1:8000/predict", files={"file": open("test.jpg", "rb")}, timeout=60 ) print(response.json())接口返回值这里做成了 PNG 的十六进制字符串,实际项目里可以直接改返回FileResponse或 base64 编码,前端拿到底图后展示掩码。服务启动后建议先用 curl 做一次连通性测试,再接入业务。
9. 资源占用与性能观察
9.1 显存占用如何观察
训练时开启另一个终端,执行下面命令实时观察显存:
watch -n 1 nvidia-smiWindows 上可以用nvidia-smi手动刷新,或者在任务管理器里看 GPU 显存曲线。
显存占用主要由四个因素决定:
- 输入图片尺寸:
256x256和512x512的显存差距接近 4 倍。 - 批量大小
batch_size:每增加一倍的 batch,显存基本翻倍。 - 网络基础通道数
base_channels:从 64 降到 32,显存和参数都会明显下降。 - 是否使用混合精度:开启 AMP 通常能减少约一半显存占用。
9.2 CPU 推理与 GPU 推理的差异
CPU 也能跑 U-Net,小尺寸图片可以完成推理,但训练时差距非常明显。如果本机没有 GPU,建议把IMAGE_SIZE调到128x128,EPOCHS先调成 10,验证整套流程没问题后,再放到 GPU 机器上正式训练。
9.3 如何降低显存占用
- 把
IMAGE_SIZE从 256 降到 128。 - 把
BATCH_SIZE从 8 降到 2。 - 把
base_channels从 64 改成 32。 - 开启 PyTorch 的自动混合精度,训练脚本里加
torch.cuda.amp.autocast()和GradScaler()。 - 使用梯度累积,让小 batch 模拟大 batch 效果。
实际显存占用多少与模型版本、输入尺寸、batch size 强相关,不能一概而论。最稳妥的做法是先跑一个 batch,看nvidia-smi里的实际占用,再决定是否加大尺寸。
10. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
torch.cuda.is_available()为 False | CUDA 版本和 PyTorch 不匹配,或驱动过旧 | 执行nvidia-smi查看驱动版本 | 按驱动版本重新安装对应 CUDA 的 PyTorch |
| 训练时报 CUDA OutOfMemory | 输入尺寸过大或 batch 过大 | 查看报错信息中涉及哪一层 | 降低 batch_size、图片尺寸、base_channels |
| 损失不下降 | 学习率过高或过低、数据标签错误 | 打印前几个 batch 的标签分布 | 调整学习率,检查掩码是否对齐 |
| 预测结果全黑 | 权重加载路径错误、预处理不一致 | 单独 print 模型输出值范围 | 检查map_location和图像归一化 |
| 预测结果全白 | 二分类逻辑反了 | 检查掩码中前景像素值是否为 255 | 确认掩码是否除以 255,或调整阈值方向 |
| 数据集读取慢 | 每张图实时 resize 造成瓶颈 | 观察 CPU 使用率 | 预先裁剪成固定尺寸、增加num_workers |
| Windows 下 DataLoader 报错 | num_workers设置过高 | 查看完整报错栈 | 把num_workers设为 0 或 2 |
| 验证集 mIoU 高但实际效果差 | 数据泄漏或验证集过小 | 检查验证集是否参与训练 | 重新划分数据,扩大验证集 |
| 接口服务请求超时 | 推理耗时过长 | 测量单张推理耗时 | 减小输入尺寸,或给接口加异步任务队列 |
几个重要排错技巧:
- 训练前先跑一个 batch 的前向,确认模型输出尺寸是
[B, num_classes, H, W]。 - 用
tensorboard或matplotlib可视化几张预测结果,不要只看指标。 - 如果 loss 出现 NaN,多半是学习率太大或数据里有异常值,先把学习率降到
1e-5试。 - 加载权重时记得加
map_location=DEVICE,否则在无 GPU 机器上会报错。
11. 最佳实践与使用建议
第一次跑通时,不要追求最好的精度,先确保整套流程能完整走通。建议先设IMAGE_SIZE=128、EPOCHS=10、BATCH_SIZE=2,几分钟内看到 loss 下降和 mIoU 变化,确认流程没问题再上大参数。
工程上给你几个实用的建议。
第一,保持一套最小可运行配置。把训练脚本、预测脚本、模型文件分开存放,数据目录固定好,不要散落在一堆实验文件夹里。模型权重、输入素材、输出结果分目录管理。
第二,批量任务要加日志和失败重试。处理大量图片时分批执行,每处理完一批就记录进度,中断后可以从断点继续,不用全部重跑。
第三,固定随机种子。训练脚本开头设置torch.manual_seed(42),这样每次复现相同结果,方便调试和对比实验。如果实验对比需要严格复现,还需要固定numpy随机种子和数据加载顺序。
第四,注意类别不均衡问题。很多分割任务里前景区域只占画面的百分之几,如果直接训练,模型会倾向预测背景。处理办法有几种:使用加权损失函数、采用 Dice Loss、对前景区域做重点采样,或者对图像做裁剪增强,让前景出现在更多训练样本中。
第五,合规使用数据。涉及人脸、车牌、医疗影像、私人场所画面的数据集,要确认授权范围,训练和推理阶段都要注意隐私保护。发布模型或商用前,要对输出结果做效果复核,特别是错误分割可能带来安全风险的场景,例如医疗辅助分析、自动驾驶感知。
12. 总结与下一步
U-Net 这个架构最值得尝试的点是:它不挑数据规模,小数据集就能训出可用的分割效果;结构清晰,代码实现起来不复杂;后续迁移到其他分割任务,只需要改数据路径和类别数。
建议先验证三件事:第一,把 U-Net 跑通训练流程并保存模型权重;第二,用一张没参与训练的真实图片做推理,看掩码质量;第三,在验证集上统计 mIoU 指标,形成第一次实验的基准结果。
最容易踩的坑集中在三个地方:掩码和数据路径不对齐、CUDA 版本和 PyTorch 不匹配、掩码预处理时插值方式用错导致边界混乱。文章里已经把这三类问题对应的排查方法列出来了,遇到直接对照处理。
后续可以继续扩展的方向包括:用 Dice Loss 优化小目标分割效果、引入注意力机制(Attention U-Net)、把编码器换成预训练的 ResNet 或 EfficientNet 做迁移学习、用 U-Net++ 或 DeepLabV3+ 对比精度、把完成训练的模型封装成 Docker 服务接入业务系统。PyTorch 的生态里还有mmsegmentation这类成熟分割工具库,跑通基础版本的 U-Net 之后,再去看这些工具会容易理解得多。