~附:安装依赖库及工程源码)
别再死磕枯燥理论了AI 时代拿实战作品说话才是硬道理原创不易哈希望可以帮到还有些许学习劲儿的同学们【进阶版还在创作中耗费精力中……】跳转到专栏目录你学习更有方向和思路……入门实践工程十四U-Net 语义分割在合成形状数据上训练~附:安装依赖库及工程源码简介实现经典 U-Net 编码-解码 跳跃连接架构在合成形状数据集随机圆形/矩形 噪声背景掩码为对应形状像素上从零训练理解「像素级分类」语义分割全流程。工程详细介绍核心思想U-Net 用对称的编码-解码结构与跳跃连接解决分割中细节丢失问题——编码器逐步提取语义分辨率降、通道升、解码器逐步恢复分辨率跳跃连接把编码器高分辨率特征拼到解码器对应层从而保留边缘细节实现像素级分类。实现方法数据numpy 程序化生成随机圆形/矩形 噪声背景图像 掩码成对无需联网。模型U-Net3 次下采样到瓶颈 8×8×1283 次上采样 跳跃拼接输出 1 通道 logits。流程二分类 BCE 损失训练 8 轮 → 用 IoU 评估 → 取 6 张测试图可视化「输入 / 真实掩码 / 预测掩码」对照。输出测试 IoU 0.90 segmentation_result.png unet.pth。目录结构14_unet_segmentation/ ├── main.py # 数据生成 U-Net 训练 IoU 评估 出图 ├── requirements.txt ├── unet.pth # 训练后的权重 └── segmentation_result.png正确安装pipinstalltorch numpy matplotlib python main.py数据由 numpy 程序化生成无需联网下载安装后即可运行。运行方式pipinstall-rrequirements.txt python main.py说明数据由 numpy 程序化生成图像 掩码无需联网、无需外部数据集。U-Net 从零训练编码器下采样 3 次到瓶颈解码器上采样 跳跃连接逐步还原。CPU 8 个 Epoch 约 2-5 分钟。预期结果测试集 IoU 通常 0.90生成segmentation_result.png输入 / 真实掩码 / 预测掩码 三行对照扩展方向换为真实医学数据视网膜血管 / 细胞分割微调多类别语义分割输出多通道Dice 损失替换为 DeepLabV3 / Mask2Former 等现代分割模型部署为实时分割视频流工程源码main.py 入门实践工程十四U-Net 语义分割在合成形状数据上训练 实现经典 U-Net 编码-解码 跳跃连接架构在合成形状数据集 随机圆形/矩形 噪声背景掩码为对应形状像素上从零训练 理解「像素级分类」语义分割全流程并可输出预测掩码可视化。 全程无外部数据、无需联网纯 numpy torch 生成数据。CPU 可跑。 运行 python main.py importosimportmatplotlib matplotlib.use(Agg)importmatplotlib.pyplotaspltimportnumpyasnpimporttorchimporttorch.nnasnnimporttorch.nn.functionalasFfromtorch.utils.dataimportDataset,DataLoader BASE_DIRos.path.dirname(os.path.abspath(__file__))DEVICEtorch.device(cudaiftorch.cuda.is_available()elsecpu)IMG_SIZE64BATCH_SIZE32EPOCHS8# ---------------- 合成数据生成 ----------------defmake_sample(sizeIMG_SIZE):生成一张图背景 1~3 个随机形状返回 image[1,H,W], mask[H,W]imgnp.random.normal(0.2,0.05,(size,size)).astype(np.float32)masknp.zeros((size,size),dtypenp.float32)n_shapesnp.random.randint(1,4)for_inrange(n_shapes):kindnp.random.choice([circle,rect])cx,cynp.random.randint(10,size-10,2)rnp.random.randint(6,14)ifkindcircle:yy,xxnp.ogrid[:size,:size]region(yy-cy)**2(xx-cx)**2r*relse:# rectx0,x1sorted(np.random.randint(0,size,2))y0,y1sorted(np.random.randint(0,size,2))regionnp.zeros((size,size),dtypebool)region[y0:y1,x0:x1]Truemask[region]1.0img[region]np.random.uniform(0.6,0.95)returnimg[None],mask# [1,H,W], [H,W]classShapeDataset(Dataset):def__init__(self,n512,sizeIMG_SIZE):self.data[make_sample(size)for_inrange(n)]def__len__(self):returnlen(self.data)def__getitem__(self,idx):img,maskself.data[idx]returntorch.from_numpy(img),torch.from_numpy(mask)# ---------------- U-Net ----------------classDoubleConv(nn.Module):def__init__(self,in_ch,out_ch):super().__init__()self.blocknn.Sequential(nn.Conv2d(in_ch,out_ch,3,padding1),nn.ReLU(inplaceTrue),nn.Conv2d(out_ch,out_ch,3,padding1),nn.ReLU(inplaceTrue),)defforward(self,x):returnself.block(x)classUNet(nn.Module):def__init__(self):super().__init__()# 编码器self.enc1DoubleConv(1,16)self.enc2DoubleConv(16,32)self.enc3DoubleConv(32,64)self.poolnn.MaxPool2d(2)# 瓶颈self.bottomDoubleConv(64,128)# 解码器self.up3nn.ConvTranspose2d(128,64,2,stride2)self.dec3DoubleConv(128,64)self.up2nn.ConvTranspose2d(64,32,2,stride2)self.dec2DoubleConv(64,32)self.up1nn.ConvTranspose2d(32,16,2,stride2)self.dec1DoubleConv(32,16)self.outnn.Conv2d(16,1,1)defforward(self,x):e1self.enc1(x)# [16,64,64]e2self.enc2(self.pool(e1))# [32,32,32]e3self.enc3(self.pool(e2))# [64,16,16]bself.bottom(self.pool(e3))# [128,8,8]d3self.dec3(torch.cat([self.up3(b),e3],dim1))# [64,16,16]d2self.dec2(torch.cat([self.up2(d3),e2],dim1))# [32,32,32]d1self.dec1(torch.cat([self.up1(d2),e1],dim1))# [16,64,64]returnself.out(d1)# [1,64,64]# ---------------- 训练 评估 ----------------defiou_score(pred,target,threshold0.5):pred(predthreshold).float()inter(pred*target).sum(dim(1,2,3))union(predtarget).clamp(0,1).sum(dim(1,2,3))return(inter/(union1e-6)).mean().item()defmain():torch.manual_seed(42)print(f设备:{DEVICE})print(生成合成形状数据集 ...)train_setShapeDataset(n512)test_setShapeDataset(n128)train_loaderDataLoader(train_set,batch_sizeBATCH_SIZE,shuffleTrue)test_loaderDataLoader(test_set,batch_sizeBATCH_SIZE)print(f训练集:{len(train_set)}测试集:{len(test_set)})modelUNet().to(DEVICE)optimizertorch.optim.Adam(model.parameters(),lr1e-3)forepochinrange(1,EPOCHS1):model.train()total_loss,total_iou,n0.0,0.0,0forimg,maskintrain_loader:img,maskimg.to(DEVICE),mask.to(DEVICE).unsqueeze(1)optimizer.zero_grad()outmodel(img)lossF.binary_cross_entropy_with_logits(out,mask)loss.backward()optimizer.step()total_lossloss.item()*img.size(0)total_iouiou_score(out,mask)*img.size(0)nimg.size(0)print(fEpoch{epoch}/{EPOCHS}损失{total_loss/n:.4f}训练IoU{total_iou/n:.4f})# 测试集评估model.eval()test_iou0.0withtorch.no_grad():forimg,maskintest_loader:img,maskimg.to(DEVICE),mask.to(DEVICE).unsqueeze(1)outtorch.sigmoid(model(img))test_iouiou_score(out,mask)*img.size(0)test_iou/len(test_set)print(f\n测试集 IoU:{test_iou:.4f})# 可视化输入 / 真实掩码 / 预测掩码withtorch.no_grad():img,masknext(iter(test_loader))imgimg[:6].to(DEVICE)predtorch.sigmoid(model(img)).cpu()fig,axesplt.subplots(3,6,figsize(16,8))forcinrange(6):axes[0,c].imshow(img[c,0].cpu(),cmapgray);axes[0,c].axis(off)axes[1,c].imshow(mask[c].cpu(),cmapgray);axes[1,c].axis(off)axes[2,c].imshow(pred[c,0],cmapgray);axes[2,c].axis(off)forr,tinenumerate([输入图像,真实掩码,预测掩码]):axes[r,0].set_title(t,fontsize12)plt.suptitle(U-Net 语义分割效果,fontsize14)plt.tight_layout()fig_pathos.path.join(BASE_DIR,segmentation_result.png)plt.savefig(fig_path,dpi120)print(f分割效果对照图已保存到:{fig_path})pthos.path.join(BASE_DIR,unet.pth)torch.save(model.state_dict(),pth)print(f模型权重已保存到:{pth})if__name____main__:main()