ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch从零实现U-Net:图像分割完整实战指南

PyTorch从零实现U-Net:图像分割完整实战指南 图像分割是计算机视觉中比目标检测更进一步的任务。目标检测用边界框回答图片里有什么、在哪个位置而图像分割要求模型对每一个像素给出类别判断。U-Net 是目前图像分割领域最常作为入门和基线使用的深度学习架构最早从医学图像分割中流行起来后来被广泛用于遥感影像、工业质检、广告牌检测等场景。这篇文章用 PyTorch 从零实现一个 U-Net覆盖数据集准备、模型搭建、训练、评估、推理和常见问题排查最终得到一个可以直接启动训练并在测试图片上输出分割掩码的完整项目。读者只需要掌握 Python 基础、PyTorch 张量操作和卷积神经网络的基本概念就可以按下面的顺序把项目跑通。1. 先理解图像分割与 U-Net 的核心机制1.1 图像分割到底解决什么问题通俗地说图像分割就是把图片里的每个像素分类。分类网络告诉整张图是什么目标检测网络用框告诉物体在哪分割网络则细化到像素级橘子属于“橘子”这一类橘子旁边的桌面属于“桌面”这一类背景中的阴影也要被分配一个确定类别。在工程上常见的分割任务可以分成两类语义分割只区分物体类别不区分同一类别的不同个体。例如地图上所有建筑都标成“建筑”。实例分割要区分同一个类别里的不同个体。例如图片里有三个人要把每个人分别标出来。U-Net 最初解决的是语义分割任务尤其适合二分类分割比如“前景/背景”、“病灶/正常组织”、“道路/非道路”。如果要做多类别语义分割只需要把输出通道数改成类别数量并更换损失函数。如果要做实例分割通常要在 U-Net 基础上再加检测分支或者改用 Mask R-CNN 这类专用网络。1.2 U-Net 的编码器-解码器结构与跳跃连接U-Net 名字里的字母 U代表它的网络结构呈现一个对称的 U 形。左边是编码器右边是解码器。编码器做的事情是不断下采样通过卷积提取特征通过最大池化降低分辨率。下采样之后特征图从高分辨率变成低分辨率通道数逐渐增加网络能学到越来越抽象、越来越语义化的信息。但代价是空间细节丢失边缘信息被压缩。解码器做的事情是上采样通过转置卷积把低分辨率特征逐步恢复到原始尺寸。但是只靠解码器自己恢复细节是不够的因为池化操作已经把大量边缘、纹理信息丢掉了。U-Net 的关键设计是跳跃连接编码器每一层的高分辨率特征直接拼接到解码器对应层。解码器在恢复分辨率时不仅能拿到上采样后的特征还能拿到编码器保留的边缘细节。这种设计同时兼顾了语义信息和空间信息非常适合分割任务。这也是 U-Net 和普通自编码器、FCN全卷积网络最大的区别。FCN 也用跳跃连接但 U-Net 的方式更直接直接用通道拼接而不是逐元素相加。1.3 为什么 U-Net 适合小数据集和医学图像U-Net 之所以在医学图像领域流行是因为医学样本数量往往只有几百到几千张标注成本高获取大量数据不现实。U-Net 在这个约束下表现突出原因有三个。第一跳跃连接让梯度可以更直接地从输出回传到浅层训练更稳定。第二对称的编码器-解码器结构参数量适中没有全连接层不容易因为参数过多而过拟合。第三U-Net 对输入尺寸更灵活常见实现支持不同分辨率的输入。这些特性不仅适用于医学图像。工业场景里广告牌区域分割、产品表面缺陷分割、卫星图像建筑提取数据量通常也不大U-Net 依然是值得优先尝试的基线模型。2. 环境准备与项目结构设计2.1 PyTorch 安装与环境检查建议使用 Anaconda 创建独立环境避免不同项目的 PyTorch 版本互相干扰。下面以 Python 3.9 为例。conda create -n unet python3.9 conda activate unet安装 PyTorch 时CPU 版本和 GPU 版本命令不同。CPU 版本可以直接安装运行小数据集没有问题但训练速度慢。GPU 版本需要本机已经正确安装 NVIDIA 驱动并且驱动支持的 CUDA 版本要能覆盖 PyTorch 的 CUDA 要求。# CPU 版本 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # GPU 版本示例为 CUDA 11.8 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118安装完成后用下面这段命令验证环境是否可用。python -c import torch; print(torch.__version__); print(cuda available:, torch.cuda.is_available()); print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else cpu)预期输出中第一行是 PyTorch 版本号第二行如果是 True说明 CUDA 可用第三行会显示显卡名称。如果第二行是 False说明当前安装的是 CPU 版本或者驱动与 PyTorch 的 CUDA 版本不匹配。实际项目里经常遇到一种情况安装时用的是某个 CUDA 版本运行大模型时报出类似you need pytorch with cu130 or higher to use optimized cuda operations的警告。这条警告说明 PyTorch 内置的 CUDA 版本低于驱动或算子要求优先方案是升级 PyTorch 到与驱动匹配的版本而不是强行在旧版本上继续训练。2.2 依赖清单与版本选择除 PyTorch 之外还需要图像处理、可视化和进度显示相关的库。下面是一个最小依赖清单。依赖库用途常见版本torch模型构建与训练2.0 及以上torchvision数据集工具、图像变换与 torch 版本对应numpy数组与数据预处理1.24 及以上Pillow读取图片9.0 及以上matplotlib查看分割结果3.5 及以上tqdm训练进度条4.60 及以上版本选择上有一条原则torch 和 torchvision 的版本必须配套。如果原始安装命令没有明确指定版本落地前要先用pip show torch torchvision确认本机版本再决定是否升级。pip install numpy Pillow matplotlib tqdm2.3 项目目录结构建议按下面的结构组织代码模型、数据、训练脚本和预测脚本分开便于后续调试和扩展。unet-segmentation/ ├── data/ │ ├── images/ # 原始图片 │ └── masks/ # 分割掩码黑底白目标 ├── checkpoints/ # 模型权重保存目录 ├── src/ │ ├── model.py # U-Net 模型定义 │ ├── dataset.py # Dataset 与数据变换 │ ├── train.py # 训练脚本 │ ├── predict.py # 推理脚本 │ └── utils.py # 指标计算、可视化工具 ├── requirements.txt └── README.md这种结构在学习阶段显得“小题大做”但一旦项目开始增长比如需要新增验证集、增加多类别分割、接入 TensorBoard 或接入 API 服务拆分文件的优势会立刻体现出来。2.4 数据集准备为了快速跑通推荐先使用公开数据集。Oxford Pets 数据集同时提供原始图片和三值掩码掩码包含背景、宠物、轮廓三类适合用来练习二分类分割。下载后把图片和掩码分别放入data/images和data/masks文件名一一对应掩码统一处理成单通道灰度图。如果你有自己的业务数据例如广告牌检测区域分割需要注意三点。第一图片和掩码的文件名必须一一对应掩码中目标区域统一为白色255背景为黑色0。第二训练集、验证集、测试集要在原始文件层面划分不要在同一个 DataLoader 里随机切分避免同一张图出现在训练和验证集合里。第三掩码必须使用最近邻插值缩放不能用双线性插值否则会因为灰度插值产生介于 0 和 255 之间的伪边界。# 掩码缩放必须使用最近邻插值 mask_transform transforms.Compose([ transforms.Resize((256, 256), interpolationtransforms.InterpolationMode.NEAREST), transforms.ToTensor() ])3. 用 PyTorch 从零搭建 U-Net 模型3.1 基础卷积模块DoubleConvU-Net 的基本单元是连续两次 3x3 卷积每次卷积后接 BatchNorm 和 ReLU。为什么用两次而非一次因为连续的两次 3x3 卷积拥有和一次 5x5 卷积相同的感受野但参数量更少同时非线性表达能力更强。import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(out_channels) self.relu1 nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(out_channels) self.relu2 nn.ReLU(inplaceTrue) def forward(self, x): x self.relu1(self.bn1(self.conv1(x))) x self.relu2(self.bn2(self.conv2(x))) return x这里所有卷积都设置padding1目的是保证卷积不改变特征图尺寸。如果漏掉 padding特征图每经过一层就会变小后续拼接时会因为尺寸不一致而报错。3.2 下采样 Down 与上采样 Up下采样模块由最大池化和 DoubleConv 组成最大池化把分辨率减半DoubleConv 增加通道数。class Down(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.pool nn.MaxPool2d(2) self.double_conv DoubleConv(in_channels, out_channels) def forward(self, x): return self.double_conv(self.pool(x))上采样模块是理解跳跃连接的关键。转置卷积先把特征图放大两倍然后从编码器对应层取出同尺寸的高分辨率特征在通道维度拼接再通过 DoubleConv 融合。class Up(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size2, stride2) self.double_conv DoubleConv(in_channels, out_channels) def forward(self, x1, x2): x1 self.up(x1) diff_y x2.size(2) - x1.size(2) diff_x x2.size(3) - x1.size(3) x1 F.pad(x1, [diff_x // 2, diff_x - diff_x // 2, diff_y // 2, diff_y - diff_y // 2]) x torch.cat([x2, x1], dim1) return self.double_conv(x)在forward(self, x1, x2)中x1是上一解码层上采样后的特征x2是编码器的跳跃连接特征。F.pad用来把x1补齐到和x2一致这样即使用奇数尺寸的输入也不会因为尺寸差一两个像素而报错。当输入是 256x256 这类偶数尺寸时diff_x和diff_y通常为 0F.pad不会产生实际填充。3.3 完整 U-Net 模型定义组合上面的模块得到完整的 U-Net。输入是三通道 RGB 图片输出是单通道分割结果。class UNet(nn.Module): def __init__(self, in_channels3, out_channels1, features(64, 128, 256, 512)): super().__init__() self.inc DoubleConv(in_channels, features[0]) self.down1 Down(features[0], features[1]) self.down2 Down(features[1], features[2]) self.down3 Down(features[2], features[3]) self.pool nn.MaxPool2d(2) self.bottleneck DoubleConv(features[3], features[3] * 2) self.up1 Up(features[3] * 2, features[3]) self.up2 Up(features[3], features[2]) self.up3 Up(features[2], features[1]) self.up4 Up(features[1], features[0]) self.outc nn.Conv2d(features[0], out_channels, kernel_size1) def forward(self, x): x1 self.inc(x) x2 self.down1(x1) x3 self.down2(x2) x4 self.down3(x3) x5 self.bottleneck(self.pool(x4)) x self.up1(x5, x4) x self.up2(x, x3) x self.up3(x, x2) x self.up4(x, x1) out self.outc(x) return out编码器部分产生四个跳跃连接特征x1分辨率 256、通道 64x2分辨率 128、通道 128x3分辨率 64、通道 256x4分辨率 32、通道 512。瓶颈层x5分辨率 16、通道 1024。解码器逐层上采样最后用 1x1 卷积把通道数压缩到out_channels。最后输出的 1x1 卷积很关键它不做信息融合只负责通道压缩。分割网络通常用它把 64 通道特征映射回类别通道数避免使用全连接层带来的尺寸限制。3.4 用随机张量验证输入输出维度写完模型先别急着训练用随机张量验证一次前向传播这是检查模型书写错误最快的方式。model UNet(in_channels3, out_channels1) x torch.randn(2, 3, 256, 256) out model(x) print(out.shape) # torch.Size([2, 1, 256, 256])如果输出尺寸和输入尺寸一致说明网络的整体结构正确。如果报维度错误优先查看报错信息里提示的是哪一行再看该层的输入输出通道数是否对上。维度问题在 U-Net 里最常见的原因是Up模块中in_channels计算错误拼接后的通道总数不是in_channels而是两侧通道之和。4. 数据加载、损失函数与训练循环4.1 Dataset 与 DataLoader 实现自定义 Dataset 负责读取图片、对齐掩码、执行变换并返回张量。下面是一个二分类分割的通用实现。import os from PIL import Image from torch.utils.data import Dataset from torchvision import transforms class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size256): self.image_dir image_dir self.mask_dir mask_dir self.image_size image_size self.image_names [ f for f in os.listdir(image_dir) if f.lower().endswith((.jpg, .jpeg, .png)) ] self.image_transform transforms.Compose([ transforms.Resize((image_size, 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, image_size), interpolationtransforms.InterpolationMode.NEAREST), transforms.ToTensor() ]) def __len__(self): return len(self.image_names) def __getitem__(self, idx): image_name self.image_names[idx] image Image.open(os.path.join(self.image_dir, image_name)).convert(RGB) mask_name image_name.rsplit(., 1)[0] .png mask Image.open(os.path.join(self.mask_dir, mask_name)).convert(L) image self.image_transform(image) mask self.mask_transform(mask) mask (mask 0.5).float() return image, mask掩码处理有一个容易忽略的细节原始掩码中目标区域是 255经过ToTensor()后变成 1.0背景保持 0.0。但有的数据集掩码不是严格的 0 和 255或者缩小时产生了中间灰度值所以用mask 0.5做一次阈值化统一成 0 和 1。这样训练时输入给损失函数的就是明确的二值标签。DataLoader 参数中batch_size影响显存占用和训练稳定性。num_workers在 Linux 环境下可以设置大于 0 来加速数据读取Windows 环境下设置为 0 更稳妥。from torch.utils.data import DataLoader train_dataset SegmentationDataset(data/images, data/masks, image_size256) train_loader DataLoader(train_dataset, batch_size8, shuffleTrue, num_workers2)4.2 损失函数为什么选 BCE Dice二分类分割最基础的是二元交叉熵损失BCEWithLogitsLoss。它内部先做 sigmoid再计算交叉熵数值稳定性比“先手动 sigmoid 再算 BCE”更好。交叉熵对每个像素独立计算损失在目标区域很小的数据上模型容易偏向把所有像素预测为背景因为背景占绝大多数。Dice 损失从区域重叠角度计算更关注前景区域的覆盖率两者组合能避免类别不平衡问题。class DiceLoss(nn.Module): def __init__(self, smooth1e-6): super().__init__() self.smooth smooth def forward(self, pred, target): pred torch.sigmoid(pred) pred_flat pred.view(pred.size(0), -1) target_flat target.view(target.size(0), -1) intersection (pred_flat * target_flat).sum(dim1) dice (2.0 * intersection self.smooth) / ( pred_flat.sum(dim1) target_flat.sum(dim1) self.smooth ) return 1.0 - dice.mean() bce nn.BCEWithLogitsLoss() dice DiceLoss() def combined_loss(pred, target): return bce(pred, target) dice(pred, target)DiceLoss返回1.0 - dice因为优化器总是做梯度下降要最小化损失Dice 分数越高越好所以用 1 减去它。如果要做多类别分割把out_channels改为类别数量损失函数换成nn.CrossEntropyLoss()模型输出不再经过 sigmoid而是配合torch.argmax获取类别索引。4.3 评估指标Dice 与 IoU训练过程中不能只看 loss还要看真实分割指标。二分类分割最常用的是 Dice 系数和 IoU交并比。def dice_score(pred, target, smooth1e-6): pred (pred 0).float() intersection (pred * target).sum() return (2.0 * intersection smooth) / ( pred.sum() target.sum() smooth ) def iou_score(pred, target, smooth1e-6): pred (pred 0).float() intersection (pred * target).sum() union pred.sum() target.sum() - intersection return (intersection smooth) / (union smooth)计算指标时必须先对预测结果做阈值化。模型输出是连续概率不处理就计算指标结果会和真实阈值下的指标有明显偏差。smooth参数防止除零当某一批次完全没有前景时分母为 0 会导致 NaN。4.4 训练主循环与模型保存下面是一个最小训练循环包含损失计算、反向传播、参数更新和模型保存。import argparse parser argparse.ArgumentParser() parser.add_argument(--epochs, typeint, default60) parser.add_argument(--batch-size, typeint, default8) parser.add_argument(--lr, typefloat, default1e-3) parser.add_argument(--image-size, typeint, default256) args parser.parse_args() device torch.device(cuda if torch.cuda.is_available() else cpu) model UNet(in_channels3, out_channels1).to(device) optimizer torch.optim.Adam(model.parameters(), lrargs.lr) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.1) for epoch in range(args.epochs): model.train() total_loss 0.0 for images, masks in train_loader: images images.to(device) masks masks.to(device) optimizer.zero_grad() outputs model(images) loss combined_loss(outputs, masks) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) avg_loss total_loss / len(train_dataset) scheduler.step() print(fEpoch {epoch 1}/{args.epochs}, loss: {avg_loss:.4f}) torch.save(model.state_dict(), fcheckpoints/unet_epoch_{epoch 1}.pth)训练里的几个参数需要根据实际环境调整参数常见值调小调大说明batch_size4 到 16训练更稳定但速度慢梯度更新更准但显存占用高显存不足时优先调小lr1e-4 到 1e-3收敛慢容易来回震荡不收敛观察 loss 曲线决定epochs50 到 100可能欠拟合可能过拟合配合早停策略image_size256 或 512显存占用低细节保留更好分辨率越高速度越慢不要只保存最后的模型建议每个 epoch 保存一次或者按验证指标保存最优模型。这样如果第 80 个 epoch 开始过拟合可以回退到第 60 个 epoch 的权重。5. 推理、可视化与结果验证5.1 加载模型进行推理训练完成后预测脚本要做四件事加载模型权重、读取图片、前向传播、输出掩码。model UNet(in_channels3, out_channels1) model.load_state_dict(torch.load(checkpoints/unet_best.pth, map_locationcpu)) model.eval() image Image.open(data/test_image.jpg).convert(RGB) 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) with torch.no_grad(): output model(input_tensor) prob torch.sigmoid(output) pred (prob 0.5).float() print(pred.shape) # torch.Size([1, 1, 256, 256])进行推理前一定要调用model.eval()。它会关闭 BatchNorm 的统计更新和 Dropout让模型使用训练阶段累计的均值和方差结果更稳定。反向传播不需要的中间结果也会被释放节省显存。5.2 后处理阈值、连通域与掩码叠加二分类预测是一个概率图阈值的选择会直接影响效果。默认使用 0.5但如果目标区域很小或者边界模糊可以统计验证集上不同阈值下的 IoU选择最优阈值。import numpy as np import matplotlib.pyplot as plt pred_np pred.squeeze().cpu().numpy().astype(np.uint8) # shape: [256, 256] pred_np pred_np * 255 image_np image.resize((256, 256)) image_np np.array(image_np) # 原图与掩码叠加显示 plt.figure(figsize(10, 5)) plt.subplot(1, 2, 1) plt.imshow(image_np) plt.title(Original) plt.subplot(1, 2, 2) plt.imshow(pred_np, cmapgray) plt.title(Prediction) plt.show()如果预测结果中有很多细小噪点可以用连通域分析去掉面积过小的区域。OpenCV 的cv2.connectedComponentsWithStats可以统计每个连通区域的面积只保留面积大于阈值的区域这在实际工程中非常常用。5.3 模型保存与恢复的两种方式PyTorch 有两种常见保存方式。第一种只保存状态字典占用空间小推荐在训练复现和模型部署时使用。torch.save(model.state_dict(), unet.pth) model.load_state_dict(torch.load(unet.pth, map_locationcpu))第二种保存完整模型包含结构信息和权重加载时不需要重新定义模型类但受代码结构变动影响较大。torch.save(model, unet_full.pth) model torch.load(unet_full.pth)实际项目中推荐第一种。原因是state_dict只保存权重模型类升级时更容易做权重兼容处理完整模型对象会把类定义和版本信息一起固化代码改动后很容易加载失败。6. 常见错误与排错链路6.1 问题速查表U-Net 训练和推理过程中错误集中在几个固定环节。下面是实际项目中最常遇到的问题。问题现象常见原因检查方式处理建议前向传播时报维度错误跳跃连接两侧特征尺寸不一致或通道数写错打印各层张量 shape检查卷积 padding 和 Up 模块拼接后的通道总数CUDA out of memorybatch_size 过大、输入分辨率过高查看nvidia-smi显存占用调小 batch_size、降低分辨率、使用梯度累积loss 不下降学习率不合适、掩码预处理错误、标签全为零先对少量样本过拟合测试调小学习率检查 mask 是否被正确二值化先跑 5 张图预测结果全黑或全白阈值设置错误、模型未收敛、sigmoid 后未处理打印预测概率的 min/max/mean调整阈值检查 loss 曲线确认输出经过 sigmoid训练集指标高但验证集差很多过拟合、数据划分不一致对比训练和验证 loss 曲线增加数据增强、使用早停、降低模型容量出现cuXXX版本警告PyTorch 内置 CUDA 版本与算子要求不匹配查看 torch.version.cuda 和驱动版本升级到匹配的 PyTorch 版本训练速度很慢但 CPU 占用不高数据读取成为瓶颈检查num_workers设置在 Linux 上增加num_workersWindows 上优先排查数据读取逻辑6.2 典型排查路径遇到训练异常不建议直接改代码乱试。按下述顺序排查通常能快速定位问题。第一确认输入数据。打印一批训练数据的 shape、数值范围和 mask 的像素分布。一个常见的低级错误是数据增强里对图片做了归一化却忘记对掩码也施加同样处理导致 mask 变成小数值二值化后全部变成背景。第二确认模型前向传播。用随机张量跑一次模型观察输出 shape 是否与输入一致再比较输出和标签的 shape 是否匹配损失函数的输入要求。第三确认损失函数。单独计算一个已知的人工样本验证 loss 是否符合预期。例如输入全是 0 的概率交叉熵应该接近一个较大的正值Dice 损失应该接近 1。第四确认训练循环。观察第一个 batch 的 loss 数值如果第一个 batch 就能看到 loss 下降说明整体链路正常。如果出现 NaN优先检查学习率是否过大、标签中是否有异常值、损失函数中是否有除零。第五确认验证阶段。推理时必须使用model.eval()和torch.no_grad()否则测试阶段结果会不稳定显存也会被没必要的中间变量占用。注意不要只验证程序能启动还要验证输入、输出、异常分支和日志是否符合预期。训练一个 epoch 后保存一张可视化预测图对比原图和掩码比只看 loss 数字直观得多。7. 从学习环境到生产环境的最佳实践7.1 学习环境与生产环境的差异本地把模型训练出来和真正把它部署成服务中间还隔着不少工程问题。学习环境下代码可以全部写在.py脚本里训练后保存权重即可。生产环境至少还需要考虑下面几项。配置外置图片路径、模型路径、batch_size、阈值不要硬编码在代码里使用 YAML 或环境变量管理。日志和监控训练时记录每个 epoch 的 loss、Dice、IoU推理时记录请求耗时和异常次数。模型版本管理每次训练都保存模型文件、训练参数、数据版本和指标便于回滚。推理优化把模型转换成 TorchScript、ONNX 或 TensorRT 格式减少部署时的 Python 依赖。输入校验推理服务必须处理非图片文件、损坏图片、超分辨率输入等情况不能假设客户端永远发送合法数据。7.2 训练前检查清单每次开始新数据训练之前按这份清单过一遍能省下大量排错时间。图片和掩码目录是否存在文件名是否一一对应。掩码是否为单通道 0/1 或 0/255没有中间灰度值。训练集和验证集是否按文件划分没有重叠。模型输入通道数与图片通道数一致输出通道数等于类别数。是否用小学习率、少量数据做过一次拟合测试。是否用随机张量验证过模型输出 shape。训练脚本是否支持断点恢复是否定期保存 checkpoints。是否记录了数据版本和训练参数便于复现。7.3 扩展方向U-Net 作为基线模型验证出可行结果之后可以根据业务需求选择延伸方向。最直接的是把编码器替换成预训练的 ResNet 或 EfficientNet这类骨干网络在 ImageNet 上预训练过在小数据集上往往能更快收敛这就是经典的 U-Net 变体 ResNet 编码器方案。需要更精确的边界信息时可以考虑 Attention U-Net在跳跃连接中加入注意力门控让模型自动抑制无关背景区域。如果目标物体的尺度差异很大比如同时存在很大的建筑和很小的广告牌可以考虑 DeepLabV3 或 PSPNet 这类引入空洞卷积和多尺度池化结构的模型。如果最终要上线边缘设备可能还要做一些工程改造把模型转成 ONNX使用半精度推理调整输入分辨率以平衡速度和精度。不过这些都是在 U-Net 跑通之后的事第一步永远是先把最小可运行的 U-Net 项目落地跑通一遍完整的数据流和训练流再讨论如何优化。如果在复现过程中遇到维度不匹配优先检查卷积 padding 和 Up 模块的拼接逻辑遇到模型不收敛优先用少量数据做拟合测试确认链路本身没有错误。把这个流程跑熟之后再切换更复杂的网络结构会顺畅得多。
返回列表