图像分割是计算机视觉里比目标检测更细一档的任务。目标检测给的是“图中哪里有一个物体”图像分割给的是“每个像素属于哪一类”连边界都给你画出来。这次我们要看的 U-Net就是做这件事最常用的基础架构之一。它最早用于医学图像分割后来在遥感、工业质检、路网提取这些场景里也大量出现。重点是它不像一些大模型那样依赖海量数据和高端显卡即使是在一张普通 GPU 甚至 CPU 上也能完成训练和推理所以特别适合作为 PyTorch 实战入门的第一个分割模型。这篇文章会完整走一遍 U-Net 图像分割的实战流程从环境安装、数据集准备、模型搭建到训练验证、指标评估、推理调用和常见报错排查。如果你正在学习 PyTorch或者想把手上的图像分割需求落地成代码这篇可以直接收藏照着做。本文适合这样的读者有一定 Python 基础看过 PyTorch 的 Tensor 和 Dataset 基本概念但还没有完整训练过一个分割模型或者已经在用目标检测想进一步做像素级分割。文章里给出的都是可运行的最小实现不依赖额外的专有框架读懂代码就能迁移到自己任务里。1. 核心能力速览能力项说明项目类型深度学习图像分割实战教程网络架构U-Net编码器-解码器结构 跳跃连接主要功能二分类分割、多类别语义分割、高分辨率掩码预测训练框架PyTorch推荐硬件NVIDIA GPU显存越大能处理的输入尺寸越大CPU 可训练但速度慢启动方式Python 脚本训练训练完成后脚本推理或封装 APIAPI 能力模型训练完成后可自行封装 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 python3.9 conda activate unetPython 版本建议选择 3.8 到 3.10。如果你已经有 Conda 环境也可以直接复用。3.2 安装 PyTorchPyTorch 安装的关键是 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, interpolationImage.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_size3, padding1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) 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_size2, stride2) 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], dim1) return self.conv(x) class UNet(nn.Module): def __init__(self, in_channels3, num_classes2, base_channels64): 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_size1) 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_classes2表示二分类分割num_classes5就表示分割 5 个类别。这里的base_channels是网络基础通道数默认 64显存紧张时可以改成 32 或 16模型体积和显存占用都会明显下降。创建模型并查看参数量model UNet(in_channels3, num_classes2) 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_classes2): 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背景会拉高整体分数掩盖前景目标效果差的问题。工程上更常用 mIoUmean 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_dirdataset/train_images, mask_dirdataset/train_masks, image_size(IMAGE_SIZE, IMAGE_SIZE) ) val_dataset SegmentationDataset( image_dirdataset/val_images, mask_dirdataset/val_masks, image_size(IMAGE_SIZE, IMAGE_SIZE) ) train_loader DataLoader(train_dataset, batch_sizeBATCH_SIZE, shuffleTrue, num_workers2) val_loader DataLoader(val_dataset, batch_sizeBATCH_SIZE, shuffleFalse, num_workers2) # 模型 model UNet(in_channels3, num_classesNUM_CLASSES).to(DEVICE) if NUM_CLASSES 2: criterion nn.BCEWithLogitsLoss() else: criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lrLEARNING_RATE) # 训练循环 best_iou 0.0 for epoch in range(EPOCHS): model.train() train_loss 0.0 for images, masks in tqdm(train_loader, descfEpoch {epoch1}/{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, dim1) 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(fEpoch {epoch1}: 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, deviceDEVICE): 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, dim1).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_channels3, num_classes2) model.load_state_dict(torch.load(best_unet.pth, map_locationDEVICE)) 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_okTrue) 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_channels3, num_classes2) model.load_state_dict(torch.load(best_unet.pth, map_locationDEVICE)) 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, formatPNG) 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)}, timeout60 ) 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调到128x128EPOCHS先调成 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()为 FalseCUDA 版本和 PyTorch 不匹配或驱动过旧执行nvidia-smi查看驱动版本按驱动版本重新安装对应 CUDA 的 PyTorch训练时报 CUDA OutOfMemory输入尺寸过大或 batch 过大查看报错信息中涉及哪一层降低 batch_size、图片尺寸、base_channels损失不下降学习率过高或过低、数据标签错误打印前几个 batch 的标签分布调整学习率检查掩码是否对齐预测结果全黑权重加载路径错误、预处理不一致单独 print 模型输出值范围检查map_location和图像归一化预测结果全白二分类逻辑反了检查掩码中前景像素值是否为 255确认掩码是否除以 255或调整阈值方向数据集读取慢每张图实时 resize 造成瓶颈观察 CPU 使用率预先裁剪成固定尺寸、增加num_workersWindows 下 DataLoader 报错num_workers设置过高查看完整报错栈把num_workers设为 0 或 2验证集 mIoU 高但实际效果差数据泄漏或验证集过小检查验证集是否参与训练重新划分数据扩大验证集接口服务请求超时推理耗时过长测量单张推理耗时减小输入尺寸或给接口加异步任务队列几个重要排错技巧训练前先跑一个 batch 的前向确认模型输出尺寸是[B, num_classes, H, W]。用tensorboard或matplotlib可视化几张预测结果不要只看指标。如果 loss 出现 NaN多半是学习率太大或数据里有异常值先把学习率降到1e-5试。加载权重时记得加map_locationDEVICE否则在无 GPU 机器上会报错。11. 最佳实践与使用建议第一次跑通时不要追求最好的精度先确保整套流程能完整走通。建议先设IMAGE_SIZE128、EPOCHS10、BATCH_SIZE2几分钟内看到 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 之后再去看这些工具会容易理解得多。