深度可分离UNet:轻量化医学图像分割的PyTorch实现与训练指南

发布时间:2026/9/23 12:33:46
深度可分离UNet:轻量化医学图像分割的PyTorch实现与训练指南 简介这是一份面向医学图像分割场景的轻量级UNet实现采用深度可分离卷积替代标准卷积在保持分割精度的同时显著减少参数量适合部署于资源受限的医疗设备。代码提供标准卷积与深度可分离卷积两种模式可通过参数灵活切换并支持常见256×256输入与多类别分割任务。资源共10个文件以Python源码为主4个py脚本、3个pyc缓存另含项目说明docx、requirements.txt及README.md压缩包总大小仅28KB体量精简但覆盖完整训练链路。配套的SegmentationDataset类具备自动标签映射、图像-掩膜配对与one-hot编码能力数据增强与ImageNet标准化同步作用于掩膜训练模块支持Dice系数评估、两种损失函数、断点续训与早停机制并可实时绘制中文双语训练曲线。目前已有70人学习下载适合具备一定深度学习基础、希望在医学图像分析方向快速搭建高效分割模型的开发者参考。1. 深度可分离UNet把UNet的参数量打下来医学图像分割才跑得动一台只有4GB显存的老GPU一张512×512的CT片基础UNet一个batch塞两三个样本就到顶了。这正是很多医学图像分割课题被卡住的起点。深度可分离UNet的思路很直接把UNet里的普通卷积块换成深度可分离卷积块让参数量和计算量都矮一大截同时把Dice的损失控制在很小的范围里。它适合三类人正在做UNet模型改进的研究生、手里有标注数据但显存不够的医工交叉从业者以及想把轻量级分割模型跑到边缘设备上的工程师。不要以为要重写整个网络才能跑一个UNet网络改动比想象中小得多。2. 深度可分离卷积在UNet里的正确打开方式替换粒度、参数账与两个变体很多人在UNet上做轻量化第一反应是把编码器换成现成的分类主干换完才发现跳跃连接的语义对不上越改越乱。深度可分离UNet走的是另一条路网络骨架完全不动只把构成编码器和解码器的卷积块换掉。这个改动工程量小、边界清楚出问题也容易回退是我在项目里最常用的一档方案。2.1 深度可分离卷积与普通卷积的参数账普通卷积在处理输入特征图时每个输出通道都要对输入的所有通道做一次3×3卷积并求和参数量是C_in×C_out×3×3。深度可分离卷积把这一件事拆成两步第一步叫Depthwise卷积输入有多少通道就分成多少组每组只对单个通道做3×3卷积参数量是C_in×3×3第二步叫Pointwise卷积用1×1卷积把C_in个通道线性组合成C_out个通道参数量是C_in×C_out。总参数量从C_in×C_out×9降成C_in×9C_in×C_out。以3×3卷积、输出通道是输入两倍的情况来算深度可分离卷积的参数量只有普通卷积的八分之一到九分之一。我把三个典型层的账算出来放在表里输入通道→输出通道普通卷积参数量深度可分离参数量压缩比16→324608656约7.0倍64→128737288768约8.4倍256→5121179648133376约8.8倍从表里能读出一个关键规律压缩比跟通道数成正比。在16→32这种浅层压缩比只有7倍而在256→512的深层接近理论极限8.8倍。这里还没算FLOPs的差异但结论方向一致。所以深度可分离UNet的替换策略重心应该放在网络深部的下采样层而不是第一层。BN的参数按通道数线性增长在两种方案里差异很小不影响上面的结论。2.2 两个容易被混为一谈的变体“深度可分离UNet”在论文和开源代码里至少指两种结构。第一种是把UNet的每个普通卷积块替换成DWPW的轻量块整个网络保持对称的U形改动最小中小数据集上最容易复现前面那张表算的就是这个方案。第二种是MobileNetV2式的倒残差结构先1×1升维再3×3深度可分离卷积降维中间层的通道数会比输入输出大好几倍。这种结构单看参数量未必比普通UNet少但它每一层的FLOPs都更低而且有ImageNet预训练权重可以用精度往往更好。在项目里怎么选我的判断标准是数据量。几百张到一两千张的医学图像用第一种数据量到万张级或者要跟别的模型做系统性的精度对比再用第二种。第二种的代价是要重新设计跳跃连接不能直接把UNet编码器的feature map拿过来用因为倒残差块的通道分布和普通卷积完全不同强行拼接会让解码器前几层学到一堆冗余特征。2.3 替换粒度浅层保留、深层替换还是全换第一个档位是全部替换。它最省显存适合GPU显存小于6GB但必须跑512×512输入的场景代价是浅层边缘特征的连续性会变差。原因不复杂DW卷积对每个通道独立操作浅层只有三四个通道每个通道被单独卷完再做1×1混合底层信息的组合方式不如普通卷积丰富。对比实验里能直接看到全换之后第一个编码器块的feature map明显更碎边缘断点变多。第二个档位是只替换编码器第三、第四个下采样块前两个块保留普通卷积。这个方案的Dice损失通常能控制在0.5个百分点以内参数量却能省下大几十个百分点。如果项目对精度敏感这是我默认推荐的做法。第三个档位是解码器保留普通卷积。上采样之后的卷积直接决定分割边缘的锐度深度可分离卷积在这里省下的参数不多但有可能让边界变模糊。因为上采样会产生插值噪声DW卷积没有跨通道的信息融合能力对这种结构噪声的抑制不如普通卷积。这三个档位不是互斥的可以在代码里作为配置项来回切换。我一般锚定“浅层保留深层替换解码器保留”这个基线再根据显存余量决定是否往全换方向调。每次改动只动一个变量Dice掉了也容易定位是哪一层出的问题——轻量化改造最怕一次改太多最后翻车了都不知道该回退哪一步。3. 用PyTorch实现深度可分离UNet模型定义、数据集加载与训练脚本这一章给出一套能在单卡上跑起来的最小实现。我用PyTorch写模型用公开的“原图同名mask”分割格式来举例这是ISIC皮肤病变、DRIVE眼底血管这类公开数据集通用的组织方式也是unet训练自己的数据集时最常见的起步格式。同样的代码改一下路径和通道数就能换到另一批数据上。3.1 定义深度可分离卷积块与UNet主体先定义两个基础块普通双卷积块留作浅层和对比实验深度可分离卷积块作为替换单元。两块代码放在同一个文件里。import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): 普通UNet使用的双卷积块保留在浅层和跳跃连接附近 def __init__(self, in_channels, out_channels): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.conv(x) class DSConvBlock(nn.Module): 深度可分离卷积块DW 3x3 PW 1x1后面接BN和ReLU def __init__(self, in_channels, out_channels): super().__init__() self.dw nn.Conv2d(in_channels, in_channels, kernel_size3, padding1, groupsin_channels) self.pw nn.Conv2d(in_channels, out_channels, kernel_size1) self.bn nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) def forward(self, x): x self.dw(x) x self.pw(x) x self.bn(x) return self.relu(x)groupsin_channels是DSConvBlock的关键它让卷积在通道维度上完全分组输入有多少通道就并行卷多少个独立通道这就是Depthwise的含义。后面的1×1卷积再把通道重新组合保证信息能跨通道流动。BN放在PW之后而不是DW之后因为PW的输出才是这个块的真正输出通道数BN在真实输出空间上做归一化才有效。ReLU的inplaceTrue能省一点显存对轻量级模型有帮助。下面是UNet主体。跳跃连接用torch.cat拼接而不是相加保留更多空间信息上采样用双线性插值而不是转置卷积转置卷积在医学图像上容易出现棋盘格伪影。class DSUNet(nn.Module): def __init__(self, in_ch1, base_ch32, n_cls1): super().__init__() chs [base_ch * (2 ** i) for i in range(4)] # 32, 64, 128, 256 # use_ds 为 False 的层保留普通卷积避免浅层过度压缩 self.enc1 self._stage(in_ch, chs[0], use_dsFalse) self.enc2 self._stage(chs[0], chs[1], use_dsFalse) self.enc3 self._stage(chs[1], chs[2], use_dsTrue) self.enc4 self._stage(chs[2], chs[3], use_dsTrue) self.bottleneck self._stage(chs[3], chs[3] * 2, use_dsTrue) up_ch chs[3] * 2 self.dec4 self._stage(up_ch chs[3], chs[3], use_dsFalse) self.dec3 self._stage(chs[3] chs[2], chs[2], use_dsFalse) self.dec2 self._stage(chs[2] chs[1], chs[1], use_dsFalse) self.dec1 self._stage(chs[1] chs[0], chs[0], use_dsFalse) self.out nn.Conv2d(chs[0], n_cls, kernel_size1) def _stage(self, in_ch, out_ch, use_ds): block DSConvBlock if use_ds else DoubleConv return nn.Sequential(block(in_ch, out_ch), block(out_ch, out_ch)) def forward(self, x): e1 self.enc1(x) e2 self.enc2(F.max_pool2d(e1, 2)) e3 self.enc3(F.max_pool2d(e2, 2)) e4 self.enc4(F.max_pool2d(e3, 2)) b self.bottleneck(F.max_pool2d(e4, 2)) d4 self.dec4(torch.cat([F.interpolate( b, scale_factor2, modebilinear, align_cornersFalse), e4], dim1)) d3 self.dec3(torch.cat([F.interpolate( d4, scale_factor2, modebilinear, align_cornersFalse), e3], dim1)) d2 self.dec2(torch.cat([F.interpolate( d3, scale_factor2, modebilinear, align_cornersFalse), e2], dim1)) d1 self.dec1(torch.cat([F.interpolate( d2, scale_factor2, modebilinear, align_cornersFalse), e1], dim1)) return self.out(d1)use_ds开关直接对应上一章的替换策略前两个编码器块保留普通卷积第三个下采样块之后全部使用深度可分离卷积解码器统一保留普通卷积。base_ch32是医学图像小数据的常用起点数据量大或者图像分辨率高时可以调到48或64。align_cornersFalse让上采样时像素中心对齐到输入网格避免特征错位。3.2 数据集加载原图加同名mask的通用做法如果你要unet训练自己的数据集最常见的问题出在数据加载这一步而不是模型结构。通用格式是一张原图对应一张同名mask原图和mask分别在两个目录。下面的数据集类按这种格式读取灰度图进来RGBA的mask也会被转成单通道。import os import glob import cv2 import numpy as np import torch from torch.utils.data import Dataset class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, size(256, 256), trainTrue): self.img_paths sorted(glob.glob(os.path.join(img_dir, *.png))) self.mask_dir mask_dir self.size size self.train train def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img cv2.imread(self.img_paths[idx], cv2.IMREAD_GRAYSCALE) name os.path.basename(self.img_paths[idx]) mask cv2.imread(os.path.join(self.mask_dir, name), cv2.IMREAD_GRAYSCALE) img cv2.resize(img, self.size) mask cv2.resize(mask, self.size, interpolationcv2.INTER_NEAREST) if self.train: if np.random.rand() 0.5: img img[:, ::-1] mask mask[:, ::-1] angle np.random.uniform(-30, 30) M cv2.getRotationMatrix2D((self.size[0] // 2, self.size[1] // 2), angle, 1.0) img cv2.warpAffine(img, M, self.size) mask cv2.warpAffine(mask, M, self.size, flagscv2.INTER_NEAREST) img torch.from_numpy(img.astype(np.float32) / 255.0).unsqueeze(0) mask torch.from_numpy((mask 127).astype(np.float32)).unsqueeze(0) return img, maskmask的resize必须用INTER_NEAREST不能用线性插值否则mask边缘会出现0.3这样的小数Dice计算和损失函数全部对不上。灰度图像直接除以255归一化不要套ImageNet的mean/std医学图像CT、超声、病理切片的灰度分布和自然图像差得很远套错反而把对比度压低了。旋转角度±30°对皮肤病变、细胞团这类没有方向先验的目标是安全的如果是肝脏、肾脏这类有明确解剖朝向的器官建议缩到±10°。3.3 训练脚本混合损失、AdamW与Dice指标损失函数用Dice Loss和BCE以1:1混合。Dice Loss对类别不平衡更友好BCE提供平滑梯度两者配合在大多数分割任务里比单用任何一个都稳。模型最后一层是线性输出BCE用binary_cross_entropy_with_logits这个函数内部先算sigmoid再做交叉熵数值上比手动sigmoidBCE稳定。def dice_coef(pred, target, smooth1.0): pred pred.view(pred.size(0), -1) target target.view(target.size(0), -1) intersection (pred * target).sum(dim1) return (2.0 * intersection smooth) / (pred.sum(dim1) target.sum(dim1) smooth) def train_one_epoch(model, loader, optimizer, scheduler): model.train() total_loss, total_dice 0.0, 0.0 for img, mask in loader: img, mask img.cuda(), mask.cuda() pred model(img) bce nn.functional.binary_cross_entropy_with_logits(pred, mask) dice 1 - dice_coef(torch.sigmoid(pred), mask).mean() loss 0.5 * bce 0.5 * dice optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() total_dice dice_coef(torch.sigmoid(pred), mask).mean().item() scheduler.step() return total_loss / len(loader), total_dice / len(loader)Dice计算时要对sigmoid之后的概率做因为目标mask是0/1概率乘以mask相当于只统计前景区域的贡献。AdamW比Adam更适合这类小数据训练weight_decay设1e-4到1e-5之间太大损失曲线会变得毛糙。学习率1e-3起步配合CosineAnnealingLR配一个Early Stopping验证集Dice连续15个epoch不升就保存当前最佳权重。这里要注意一点如果开了混合精度训练torch.cuda.ampDice和损失都要在FP32下计算AMP只加速卷积和矩阵乘评价指标用FP16会损失精度。注意判断一个轻量化改造是否值得不要只看参数量。训练时同时观察GPU利用率和每轮耗时这两个指标比FLOPs更接近真实体验。4. 深度可分离UNet训练避坑Dice卡住、小目标丢失和GPU变慢的解法这一章讲的坑大多不是模型结构写错而是配置和数据层面的问题。我把训练深度可分离UNet时遇到过的、以及帮同行排查过的问题整理成五条按现象、原因、解决三步写。4.1 现象Loss规律下降Dice却卡在0.5附近不动训练曲线显示loss一直在降Dice就是上不去。如果损失函数里BCE权重过大而目标只占图像面积的几个百分点模型很快学到“输出全背景”这个局部最优——此时BCE已经很低因为背景占了绝大多数像素但Dice几乎为零。解决方法是把混合损失改成Dice为主比如0.7 Dice加0.3 BCE。另一个容易被忽略的原因是初始学习率过大1e-2起步会让权重在最优解附近反复横跳Dice表现为前期剧烈波动、后期卡死。遇到这种情况先不要动结构把学习率降到1e-3重跑一次。调学习率本来就有玄学成分但1e-2在分割任务里绝大多数时候就是偏大先排除这个再调别的。4.2 现象小目标器官在预测mask上直接消失在胰脏、小淋巴结这类目标上模型输出的mask经常是全黑或只有零星几个点。UNet下采样四次后特征图分辨率只有输入的1/16小目标在最低分辨率上只剩几个像素。深度可分离卷积对通道独立处理通道间没有信息交换低分辨率下特征更容易被池化吞掉。解决手段一般有两个。第一是深度监督在解码器每一层输出上都接一个1×1卷积把预测结果上采样回原尺寸后分别计算损失这样浅层也能收到梯度信号。第二是保留最浅层跳跃连接把它和最后一层解码器的输出做一次拼接相当于把底层边缘信息直接送到输出端。两个手段可以同时用对小目标的提升很直接代价是训练时间增加百分之十几但省下来的调参时间远比这个多。4.3 现象显存降下来了训练反而变慢GPU利用率不到40%这是深度可分离卷积最容易被忽略的代价。DW卷积在GPU上等于把一个大矩阵乘拆成几十个小矩阵乘kernel启动和访存的开销占比变大FLOPs降了但墙钟时间可能不降反升batch size越小越明显。解决方式是在GPU上优先用组卷积代替极端的深度可分离卷积比如groups8或groups16。参数压缩效果依然可观但并行效率高得多。如果坚持用DW卷积至少把batch size调大让GPU吞吐覆盖kernel启动开销。batch size为1时深度可分离UNet在消费级显卡上会有明显的“跑不满”感觉这属于深度可分离卷积的固有特性不是你的代码写错了。4.4 现象训练验证都很漂亮一换设备数据就崩医学图像分割的跨设备泛化问题普遍存在。不同扫描仪、不同层厚的图像灰度分布差异很大统一用全局mean/std归一化等于把这种设备差异直接送进网络。深度可分离卷积削弱了通道间的信息混合能力模型会更依赖单通道特征而单通道特征对设备差异最敏感。做法是先把图像按目标器官区域做自适应裁剪再做分位点归一化。CT图先做窗宽窗位调整到目标器官的CT值范围再归一化比任何模型层面的trick都有效。如果换设备的差异实在难以消除推理时用TTA水平翻转和90度旋转各预测一次取平均通常能把Dice拉回几个点代价是推理时间翻几倍。4.5 现象batch size只有2或4时训练早期就发散医学图像标注成本高很多项目batch size调不到16以上。普通UNet在小batch下勉强能训深度可分离UNet的BN层会放大这种不稳定每个通道单独归一化统计量噪声本来就大batch再小训练早期很容易直接发散去。改用GroupNorm是单卡场景最稳的方案把通道分成8组做归一化不依赖batch统计量参数和计算开销几乎没有变化。BatchNorm在小batch下的抖动是统计学问题不是超参数调不好的问题硬调学习率只会把另一个正常配置搞坏。如果项目必须用BN那就减小输入分辨率或者加batch size到8以上没有别的省事路径。注意每改一个配置建议同时记录Dice均值、训练每轮耗时和显存占用。没有这三个基线数据后续任何优化都等于闭着眼睛调。5. 让深度可分离UNet真正交付交叉验证、特征可视化和一个复用后处理一个分割模型能不能作为结论写进论文或者交给临床科室用我看的不是训练集Dice而是三件事方差是否足够小、模型到底在看图像的什么位置、预测mask是否需要人工大量修图。5.1 验证Dice稳定性5折交叉验证与特征图可视化交叉验证是最先要做的。用5折交叉验证取代单次划分医学图像标注主观性强单次划分的运气成分太大。5折的Dice均值和标准差能说明模型是否稳定均值高而标准差大的模型往往只在某几例数据上表现好这样的模型不适合上线。深度可分离UNet省下来的显存可以用来跑更大的batch刚好让交叉验证的耗时缩短一些。特征可视化用来验证模型是否学到器官结构。把编码器最后一层的feature map按通道求均值缩放到原图大小后叠加到原图上。高响应区域集中在器官边界和内部纹理说明模型学到的是形状特征高响应集中在前景目标之外说明模型学到的是扫描仪伪影。这个检查十分钟就能做完能避免大量无效调参。不要等模型训完才去打开黑匣子训练第二天就做一次可视化早点发现学歪了还能及时调整。5.2 交付前的最后一步一个可复用的mask后处理后处理直接决定交付时的观感。我长期在分割项目里用下面这段代码每次预测完自动滤掉零散假阳并填补边界空洞import cv2 import numpy as np def refine_mask(prob, min_area50): mask (prob 0.5).astype(np.uint8) n, labels, stats, _ cv2.connectedComponentsWithStats(mask, connectivity8) keep [i for i in range(1, n) if stats[i, cv2.CC_STAT_AREA] min_area] if not keep: return np.zeros_like(mask) out np.isin(labels, keep).astype(np.uint8) kernel np.ones((3, 3), np.uint8) return cv2.morphologyEx(out, cv2.MORPH_CLOSE, kernel)connectedComponentsWithStats先找出所有连通域面积小于min_area的直接丢弃这一步能滤掉大部分噪声假阳。MORPH_CLOSE用3×3核做闭运算对mask边缘的毛刺和内部空洞做一次修补。min_area按目标在图像中的像素数来设256×256图像里的皮肤病灶通常大于50像素噪声点往往只有几个像素这个值可以先统计训练集mask的分布再定。我自己跑分割项目的习惯是把交叉验证、特征可视化和后处理写成一个固定的评估脚本每次训练一结束先跑这个脚本而不是先看训练loss。很多模型改进能不能被接受基本在脚本跑完的前十分钟就能判断。这套流程在深度可分离UNet上适用换到别的分割网络同样适用希望帮到你。本文还有配套的精品资源点击获取