拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch从零实现UNet:数据流、跳跃连接与小样本分割实战

简介本资源是一份面向深度学习初学者与图像分割实践者的PyTorch U-Net实战项目聚焦医学及通用场景下的自定义数据集训练全流程。内容涵盖U-Net网络结构实现、数据预处理含掩码生成、训练/测试/评估完整代码链以及配套说明文档与示例图像助用户快速掌握端到端图像分割建模能力。压缩包共27个文件含7个核心Python脚本如net.py、train.py、data.py、5张示例图像与结果图、5份Markdown说明文档含中英文README、5个XML标注文件另有.gitignore、LICENSE等工程配置文件整体仅602KB轻量易部署。目前已有760人学习下载目录结构清晰分层data/、utils/、evaluation/等所有模块解耦良好支持替换数据路径后一键训练附带make_mask_data.py等实用工具脚本显著降低从零复现门槛。1. 用 PyTorch 从零搭 UNet 不是调包而是掌控数据流、损失计算和梯度回传的完整链路很多人以为“PyTorch 搭建自己的 UNet”就是pip install unet后改几行model UNet(num_classes2)—— 实际上UNet 没有官方 PyTorch 实现所有pytorch-UNet.zip类项目本质是开发者手写编码的结构模板。它解决的不是“能不能跑”而是“如何让桥墩裂缝、皮肤病变、POI 地理栅格等小规模自定义数据集在显存有限、标注稀疏、类别不平衡的现实条件下稳定收敛”。你不需要复现原始论文全部细节但必须亲手定义下采样路径中的卷积核尺寸与 padding 策略、跳跃连接的通道对齐方式、以及nn.Upsample与ConvTranspose2d在上采样阶段的数值稳定性差异。本文面向已能写Dataset.__getitem__但常卡在loss.backward()报nan或CUDA out of memory的中级使用者不讲张量维度推导公式只给可粘贴验证的代码块、必调参数表和三类典型数据集桥墩病害、ACNE04 皮肤镜、POI 栅格的预处理适配逻辑。2. UNet 结构实现从 Encoder-Decoder 对称性到跳跃连接的通道对齐UNet 的核心不在“U”形外观而在特征金字塔内信息流动的确定性。原始论文中 encoder 每层输出通道数为[64, 128, 256, 512, 1024]decoder 对应为[512, 256, 128, 64]但直接照搬会导致跳跃连接时torch.cat([x_upsampled, x_skip], dim1)因通道数不匹配而报错。必须根据输入图像尺寸和显存约束动态调整而非硬编码。2.1 基础 UNet 模块ConvBlock 与 Down/Up 操作的封装逻辑ConvBlock是 UNet 的原子单元需同时支持 BN ReLU 卷积顺序且保证padding1时输出尺寸不变即H_out H_in。若使用Conv2d(3, 64, 3, padding1)而未设biasFalseBN 层会与 bias 冲突导致训练不稳定import torch import torch.nn as nn class ConvBlock(nn.Module): def __init__(self, in_ch, out_ch, kernel_size3, padding1): super().__init__() self.conv1 nn.Conv2d(in_ch, out_ch, kernel_size, paddingpadding, biasFalse) self.bn1 nn.BatchNorm2d(out_ch) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_ch, out_ch, kernel_size, paddingpadding, biasFalse) self.bn2 nn.BatchNorm2d(out_ch) def forward(self, x): x self.relu(self.bn1(self.conv1(x))) x self.relu(self.bn2(self.conv2(x))) return x注意inplaceTrue可节省显存但若后续需对x做梯度检查如torch.autograd.grad应设为FalsebiasFalse是强制要求因 BN 已含可学习偏置项。2.2 Encoder 路径下采样必须用 MaxPool2d 而非 stride 卷积UNet 原始设计强调无损下采样——即 pooling 不引入额外可学习参数避免梯度在压缩阶段被污染。虽然Conv2d(in_ch, out_ch, 3, stride2)更省显存但实测在桥墩病害数据集图像尺寸 512×512裂缝宽度仅 3–5 像素上会导致细小目标丢失。必须用MaxPool2d(2)并手动补零以保持尺寸整除class UNetEncoder(nn.Module): def __init__(self, in_ch3, features[64, 128, 256, 512]): super().__init__() self.enc_blocks nn.ModuleList() self.poolings nn.ModuleList() prev_ch in_ch for feat in features: self.enc_blocks.append(ConvBlock(prev_ch, feat)) self.poolings.append(nn.MaxPool2d(2)) prev_ch feat def forward(self, x): skip_connections [] for block, pool in zip(self.enc_blocks, self.poolings): x block(x) skip_connections.append(x) # 保存 H×W×C 特征图供 decoder 拼接 x pool(x) return x, skip_connections2.2.1 关键参数表不同输入尺寸对应的 features 配置建议输入图像尺寸推荐features列表显存占用单卡 RTX 3090适用场景256×256[32, 64, 128, 256]~2.1 GBPOI 栅格、ACNE04 皮肤镜小图512×512[64, 128, 256, 512]~4.7 GB桥墩病害高清图、Kitti 车道线分割1024×1024[64, 128, 256, 512, 1024]~11.3 GB高精度遥感影像、病理切片提示若显存不足优先削减features[-1]即 bottleneck 层通道数而非减少层数——因为浅层特征对定位至关重要bottleneck 层过大会导致信息瓶颈。2.3 Decoder 路径上采样必须用ConvTranspose2d并校验输出尺寸nn.Upsample(scale_factor2)仅插值不学习ConvTranspose2d可反向传播梯度但易出现棋盘效应checkerboard artifacts。解决方案是先Upsample插值再接Conv2d校正既规避棋盘纹又保留可学习性class UNetDecoder(nn.Module): def __init__(self, features[64, 128, 256, 512], num_classes2): super().__init__() self.up_convs nn.ModuleList() self.dec_blocks nn.ModuleList() self.features features[::-1] # [512, 256, 128, 64] # 上采样模块每个 up_conv 将通道减半并将 H/W ×2 for i in range(len(features)-1): up_conv nn.Sequential( nn.Upsample(scale_factor2, modebilinear, align_cornersTrue), nn.Conv2d(self.features[i], self.features[i1], 1) # 1×1 卷积降维 ) self.up_convs.append(up_conv) # 拼接后通道数翻倍需用 ConvBlock 重新融合 concat_ch self.features[i1] * 2 self.dec_blocks.append(ConvBlock(concat_ch, self.features[i1])) # 最终分类头 self.final_conv nn.Conv2d(self.features[-1], num_classes, 1) def forward(self, x, skip_connections): skip_connections skip_connections[::-1] # 逆序[enc4, enc3, enc2, enc1] for i, (up_conv, dec_block, skip) in enumerate(zip(self.up_convs, self.dec_blocks, skip_connections)): x up_conv(x) # 关键skip 与 x 尺寸必须严格一致否则 cat 失败 if x.shape ! skip.shape: # 强制双线性插值对齐应对 padding 导致的奇偶尺寸偏差 x torch.nn.functional.interpolate(x, sizeskip.shape[2:], modebilinear, align_cornersTrue) x torch.cat([x, skip], dim1) x dec_block(x) return self.final_conv(x)2.3.1 尺寸对齐调试技巧打印每层 shape 并定位 mismatch在forward中插入调试语句print(fBefore up_conv {i}: {x.shape}) x up_conv(x) print(fAfter up_conv {i}: {x.shape}, skip shape: {skip.shape})常见 mismatch 场景skip.shape [B, C, 128, 128],x.shape [B, C, 129, 129]→ 因padding1且H_in为奇数H_out H_in 2*padding - kernel_size 1 129解决方案在Dataset.__getitem__中强制transforms.Resize((512, 512))或用torch.nn.functional.pad补零至偶数尺寸。3. 自定义数据集构建针对桥墩病害、ACNE04、POI 栅格的三类预处理范式UNet 训练失败 70% 源于数据加载器返回的 tensor 不满足H % 16 0 and W % 16 0因 4 层下采样总 stride2⁴16。不能依赖transforms.Resize简单拉伸必须按数据特性定制 pipeline。3.1 桥墩病害数据集裂缝像素占比 0.5%需 foreground-aware 采样桥墩图像中裂缝区域常不足整图 0.3%直接随机裁剪 256×256 会导致 batch 内多数样本无正样本。解决方案是先统计 mask 中非零像素坐标再以这些坐标为中心做带偏移的裁剪import numpy as np from torchvision import transforms class BridgeDefectDataset(torch.utils.data.Dataset): def __init__(self, img_paths, mask_paths, crop_size256): self.img_paths img_paths self.mask_paths mask_paths self.crop_size crop_size self.transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) def __getitem__(self, idx): img Image.open(self.img_paths[idx]).convert(RGB) mask Image.open(self.mask_paths[idx]).convert(L) # 二值 mask # 获取所有裂缝像素坐标 mask_arr np.array(mask) coords np.argwhere(mask_arr 0) if len(coords) 0: # 随机选一个裂缝点作为 crop 中心 center coords[np.random.randint(len(coords))] h, w img.height, img.width # 计算 crop 区域确保不越界 top max(0, center[0] - self.crop_size//2) left max(0, center[1] - self.crop_size//2) bottom min(h, top self.crop_size) right min(w, left self.crop_size) img img.crop((left, top, right, bottom)) mask mask.crop((left, top, right, bottom)) else: # 无裂缝则随机 crop i, j, h, w transforms.RandomCrop.get_params(img, (self.crop_size, self.crop_size)) img transforms.functional.crop(img, i, j, h, w) mask transforms.functional.crop(mask, i, j, h, w) return self.transform(img), torch.tensor(np.array(mask) // 255, dtypetorch.long)提示mask // 255将 0–255 灰度转为 0–1 整型标签这是nn.CrossEntropyLoss的输入要求若用nn.BCEWithLogitsLoss则需mask.float() / 255.0。3.2 ACNE04 数据集多类别皮肤病变需 class-balanced batch 构建ACNE04 含 4 类粉刺、丘疹、脓疱、结节各类别样本数极不均衡粉刺占 62%结节仅 8%。WeightedRandomSampler可缓解但需按类别统计权重from torch.utils.data import WeightedRandomSampler # 统计每类样本数 class_counts [1240, 892, 653, 215] # 示例粉刺/丘疹/脓疱/结节 class_weights 1. / torch.tensor(class_counts, dtypetorch.float) samples_weight torch.zeros(len(dataset)) for idx, (_, label) in enumerate(dataset): samples_weight[idx] class_weights[label] sampler WeightedRandomSampler(samples_weight, num_sampleslen(dataset), replacementTrue) train_loader DataLoader(dataset, batch_size8, samplersampler, num_workers4)3.3 POI 栅格数据集地理坐标转像素需 affine transform 保持空间一致性POI 数据集常以(lat, lon)坐标存储需转为图像像素索引。关键约束resize 后的地理比例尺必须恒定否则 UNet 学习的尺度先验失效。推荐用rasterio读取 GeoTIFF提取仿射变换矩阵import rasterio from rasterio.transform import from_bounds def geo_to_pixel(transform, lon, lat): 将地理坐标转为像素坐标 col, row ~transform * (lon, lat) # 逆变换 return int(row), int(col) # row 对应 ycol 对应 x # 加载时获取 transform with rasterio.open(poi_raster.tif) as src: transform src.transform # Affine(a0.0001, b0, c116.0, d0, e-0.0001, f39.0) img src.read(1) # 读取第 1 波段 mask np.zeros_like(img) # 将 POI 坐标打点到 mask for lon, lat in poi_coords: y, x geo_to_pixel(transform, lon, lat) if 0 y img.shape[0] and 0 x img.shape[1]: mask[y, x] 14. 训练循环与损失函数Dice Loss Focal Loss 混合策略及梯度裁剪实战UNet 在小目标分割中易受背景主导单一CrossEntropyLoss会导致 dice score 0.4。必须组合 Dice Loss关注重叠率与 Focal Loss抑制易分样本。4.1 Dice Loss 实现避免分母为 0 的平滑项设置def dice_loss(pred, target, smooth1e-6): pred: [B, C, H, W], logits target: [B, H, W], long tensor pred_soft torch.softmax(pred, dim1) # 转概率 target_onehot torch.nn.functional.one_hot(target, num_classespred.shape[1]) target_onehot target_onehot.permute(0, 3, 1, 2).float() # [B, C, H, W] intersection (pred_soft * target_onehot).sum(dim(2,3)) # [B, C] union pred_soft.sum(dim(2,3)) target_onehot.sum(dim(2,3)) # [B, C] dice (2. * intersection smooth) / (union smooth) return 1 - dice.mean() # mean over batch and classes4.2 Focal Loss 实现gamma2.0 对桥墩裂缝提升显著class FocalLoss(nn.Module): def __init__(self, alpha1, gamma2, reductionmean): super().__init__() self.alpha alpha self.gamma gamma self.reduction reduction def forward(self, inputs, targets): ce_loss F.cross_entropy(inputs, targets, reductionnone) pt torch.exp(-ce_loss) focal_weight (1 - pt) ** self.gamma loss focal_weight * ce_loss if self.reduction mean: return loss.mean() return loss focal_loss FocalLoss(alpha1, gamma2)4.3 混合损失与训练主循环梯度裁剪阈值设为 1.0model UNet(in_ch3, num_classes2, features[64,128,256,512]) optimizer torch.optim.Adam(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, factor0.5, patience5) for epoch in range(100): model.train() total_loss 0 for batch_idx, (data, target) in enumerate(train_loader): data, target data.cuda(), target.cuda() optimizer.zero_grad() output model(data) # [B, 2, H, W] loss_dice dice_loss(output, target) loss_focal focal_loss(output, target) loss 0.7 * loss_dice 0.3 * loss_focal # 权重可调 loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 关键防止 nan optimizer.step() total_loss loss.item() # 验证 val_loss validate(model, val_loader) scheduler.step(val_loss) print(fEpoch {epoch}, Train Loss: {total_loss/len(train_loader):.4f}, Val Loss: {val_loss:.4f})注意clip_grad_norm_阈值设为1.0而非默认5.0因 UNet 梯度易在跳跃连接处爆炸若 loss 曲线在 epoch 3–5 后突然跳变立即检查是否漏了.cuda()或targetdtype 是否为long。5. 模型验证与推理优化用 sliding window 处理大图及 ONNX 部署关键参数UNet 输入尺寸固定但实际桥墩检测需处理 4000×3000 像素图像。直接 resize 会模糊裂缝细节必须用滑动窗口sliding window分块预测再拼接。5.1 Sliding Window 推理overlap1/3 且 padding 保证边缘完整性def sliding_window_inference(model, image, tile_size(512, 512), overlap0.33): image: [C, H, W] tensor c, h, w image.shape tile_h, tile_w tile_size step_h int(tile_h * (1 - overlap)) step_w int(tile_w * (1 - overlap)) # 输出初始化 prob_map torch.zeros((2, h, w), deviceimage.device) count_map torch.zeros((h, w), deviceimage.device) for y in range(0, h - tile_h 1, step_h): for x in range(0, w - tile_w 1, step_w): tile image[:, y:ytile_h, x:xtile_w].unsqueeze(0) # [1, C, H, W] with torch.no_grad(): pred torch.softmax(model(tile), dim1)[0] # [2, H, W] prob_map[:, y:ytile_h, x:xtile_w] pred count_map[y:ytile_h, x:xtile_w] 1 # 归一化 prob_map / count_map.unsqueeze(0) return torch.argmax(prob_map, dim0) # [H, W] # 使用示例 large_img torch.randn(3, 4000, 3000).cuda() pred_mask sliding_window_inference(model, large_img)5.2 ONNX 导出必须指定 dynamic_axes 以支持任意尺寸输入dummy_input torch.randn(1, 3, 512, 512).cuda() torch.onnx.export( model, dummy_input, unet_bridge.onnx, input_names[input], output_names[output], dynamic_axes{ input: {2: height, 3: width}, output: {2: height, 3: width} }, opset_version12 )5.2.1 ONNX Runtime 验证脚本检查输出 shape 是否与输入一致import onnxruntime as ort import numpy as np ort_session ort.InferenceSession(unet_bridge.onnx) inputs np.random.randn(1, 3, 512, 512).astype(np.float32) outputs ort_session.run(None, {input: inputs}) print(fONNX output shape: {outputs[0].shape}) # 应为 (1, 2, 512, 512)5.3 桥墩病害检测落地技巧后处理用 morphological closing 去除噪声裂缝预测 mask 常含孤立噪点OpenCV 的闭运算closing比简单阈值更鲁棒import cv2 def postprocess_mask(mask_np, kernel_size3): kernel np.ones((kernel_size, kernel_size), np.uint8) # 先膨胀再腐蚀填充细小空洞 closed cv2.morphologyEx(mask_np, cv2.MORPH_CLOSE, kernel) # 连通域分析只保留面积 50 像素的区域 num_labels, labels cv2.connectedComponents(closed) sizes [np.sum(labels i) for i in range(1, num_labels)] valid_labels [i1 for i, s in enumerate(sizes) if s 50] filtered np.isin(labels, valid_labels).astype(np.uint8) * 255 return filtered # 使用 mask_np pred_mask.cpu().numpy() * 255 # 转 uint8 clean_mask postprocess_mask(mask_np)提示cv2.MORPH_CLOSE的kernel_size需根据裂缝平均宽度设定——桥墩裂缝通常 3–8 像素宽故kernel_size3若处理 ACNE04 的粉刺直径 10–20 像素应设为7。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门