ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

深度可分离UNet:轻量医学图像分割与显存优化实战

深度可分离UNet:轻量医学图像分割与显存优化实战 简介这套代码实现了一个基于深度可分离卷积的轻量级UNet模型专为高效医学图像分割设计在保持分割精度的同时大幅压缩参数量适合部署在算力有限的医疗设备上。压缩包共10个文件包含4个Python脚本、3个pyc缓存、1份docx项目说明书、1个requirements.txt及1个Markdown说明整体仅28KB目录按models、utils、train.py等模块划分便于直接复用。配套四份文档分别覆盖模型、数据、训练与主控模型支持标准卷积与深度可分离卷积灵活切换通道数最高1024数据处理实现自动标签映射、图像-掩膜配对和动态one-hot编码训练采用Dice系数与BCE/CrossEntropy双损失并支持断点续训、早停主控模块可实时绘制双语训练曲线、自动检测GPU/CPU。已有70人学习下载适合医学影像算法研究与轻量化分割应用开发者可基于源码快速改造成多类别分割工具。1. 深度可分离UNet轻量版医学图像分割真的能省显存吗医学图像分割一直有个绕不开的现实矛盾UNet 这类编码器-解码器结构精度稳但参数量和显存占用都不小真要把它部署到床边超声、嵌入式 GPU 这类资源受限的设备上常常第一步就被卡住。这套「深度可分离UNet」把 UNet 里的普通卷积整体换成深度可分离卷积参数量直接降一个量级同时保留 use_separable 开关一行参数就能切回标准卷积做对比实验。资源把模型、数据流水线、训练评估脚本配齐了适合想改进 UNet 结构、或需要把分割模型压到小显存设备上的从业者从读代码到跑通训练再到换自己的数据集是一条完整的链路。2. 网络结构改造普通卷积和深度可分离卷积的参数量差在哪2.1 先算一笔账普通卷积和深度可分离卷积差在哪深度可分离卷积不是新东西MobileNet 靠它把分类网络压到能在手机端跑这套 UNet 沿用同样思路把标准卷积拆成 Depthwise 逐通道卷积和 Pointwise 1×1 卷积两步。标准 3×3 卷积把所有输入通道同时卷参数量是 3×3×in_ch×out_ch。拆开后Depthwise 层每个通道单独卷参数只有 3×3×in_chPointwise 层负责跨通道融合参数是 in_ch×out_ch。合计 9×in_ch in_ch×out_ch。卷积方式参数量公式inoutCC64 时相对标准卷积标准 3×3 卷积9×C²36864100%深度可分离卷积9×C C²4672约 12.7%C64 时省得最明显通道越深比例越接近 1/9 加一个小尾巴。这套 UNet 的 encoder 从 64 通道一路翻倍到 1024深层省下的参数量非常可观。要注意两点第一这个账只算了卷积层BatchNorm 和激活不在内整体结论不变但数字会略低第二当通道数小时比如输入层 3→64Depthwise 部分省不了多少但 Pointwise 已经把大头压下来了。还有一个经常被忽略的点深度可分离卷积里 Depthwise 用的仍然是 3×3 卷积感受野和标准卷积相同所以替换之后 UNet 的编码-解码结构完全不用动跳跃连接、池化位置全都保持不变。这就是为什么这类改进能直接套进现成 UNet 代码里也是它作为 UNet 模型改进落点最方便的原因。对这层结构心里有数之后读 models/unet.py 时就不会被代码行数带偏重点只在一个 conv_block 函数。2.2 use_separable 开关两种卷积模式怎么落到代码上模型把卷积封装成函数读代码第一个要看的就是这个块def conv_block(in_ch, out_ch, use_separableTrue): if use_separable: # Depthwisegroupsin_ch每个输入通道独立做 3x3 卷积 depthwise nn.Conv2d(in_ch, in_ch, kernel_size3, padding1, groupsin_ch) # Pointwise1x1 卷积把各通道结果融合到 out_ch pointwise nn.Conv2d(in_ch, out_ch, kernel_size1) return nn.Sequential( depthwise, pointwise, nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) else: return nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) )关键在 groupsin_ch它让 Conv2d 每个输入通道单独走一遍 3×3 卷积互不共享权重这就是 Depthwise 的精髓。后面的 1×1 Pointwise 把通道重新融合相当于一次跨通道线性组合。两个卷积 padding 都是 1保持分辨率不变所以这个块能直接替换 UNet 里任何标准卷积块输入输出形状完全一致不需要动周围的池化和跳跃连接。两个细节值得注意。第一bias 处理标准分支 biasFalse因为后面紧跟 BatchNormBN 自带可学习的偏移留着 bias 既浪费参数又引入冗余这是 PyTorch 里卷积BN 的标准写法深度可分离分支为了行为一致也这么处理。第二BN 放在 Pointwise 之后而不是 Depthwise 之后这是 MobileNet 验证过的顺序比每层各放一个 BN 更省参数也更稳。切到深度可分离模式后梯度流动方式和标准卷积不同我一般会把学习率从 1e-3 降到 8e-4 再跑对比实验否则前期 loss 波动会明显偏大。2.3 模型实例化通道翻倍到 1024 与跳跃连接结构这套 UNet 的 encoder 每下采样一次通道翻倍base_channels64 起64→128→256→512→1024bottleneck 处最高 1024。对照 unet 网络结构图来读代码最省力down 路径每层是一个 conv_block 加一次 MaxPool分辨率减半通道翻倍到 bottleneck 之后进 decoder每层先上采样再把 encoder 对应层的输出按通道 concat 进来最后接一个 conv_block。concat 是 UNet 的招牌操作decoder 的输入通道数是「上采样后的通道数 encoder 同层通道数」所以跳跃连接处的通道匹配是模型能直接实例化的前提。from models.unet import UNet # 二分类背景 病灶输出 2 通道走 softmax model UNet(in_channels3, num_classes2, use_separableTrue, base_channels64) # 纯二值分割输出 1 通道走 sigmoid BCEWithLogitsLoss model_binary UNet(in_channels3, num_classes1, use_separableTrue, base_channels64)实例化只要把 use_separable 改成 False就完整回到标准 UNet。我建议第一次跑先把两种模式的参数量都打出来对比眼见为实def count_params(model): return sum(p.numel() for p in model.parameters()) print(separable:, count_params(UNet(3, 2, True))) print(standard :, count_params(UNet(3, 2, False)))输入 256×256 的 RGB 图像输出和输入同尺寸通道数等于 num_classes。base_channels 可以整体缩放显存吃紧时把 64 改成 32参数量直接掉一个量级代价是分割精度可能下降几个点适合先跑通流程再决定要不要加回去。提示use_separableFalse 就是标准 UNet正好用来做「改进前后」的对照实验不用维护两套代码。3. 数据流水线SegmentationDataset 的标签映射与 one-hot 陷阱3.1 自动标签映射处理非连续标注的第一道坎医学分割数据集最常见的一个坑是标签不连续。画好的掩膜里可能是 0、5、10 这类值而不是 0、1、2直接拿去 one-hot 会让 np.eye(num_classes)[mask] 数组越界。这套代码的 SegmentationDataset 用一张 label_map 字典解决class SegmentationDataset(Dataset): def __init__(self, image_paths, mask_paths, num_classes, label_mapNone, augmentFalse): self.image_paths image_paths self.mask_paths mask_paths self.num_classes num_classes # 示例{0:0, 5:1, 10:2}把原始标签映射到 0~num_classes-1 self.label_map label_map if label_map is not None else {i: i for i in range(num_classes)} self.augment augment def _map_labels(self, mask): mapped np.zeros_like(mask, dtypenp.int64) for src, dst in self.label_map.items(): mapped[mask src] dst return mapped映射逻辑是逐标签遍历赋值先初始化为全 0 再按字典覆盖。注意这里不要直接在原数组上改mapped 是新开的一块内存因为后面增强、one-hot 都可能复用 mask 的原始值。label_map 没覆盖到的标签会被静默归到背景类 0这既是方便也是隐患——如果掩膜里有 255 这种边界值会被悄悄当成背景等训练 loss 异常时再回头查就费劲了。拿到数据第一件事我一般会把所有掩膜的 np.unique 结果都打印出来看一眼到底有哪些标签值再写 label_map。不要想当然认为标签一定从 0 开始连续医学标注软件导出的数据经常不是这样。3.2 智能配对与动态 one-hot输出形状总对不上的根源SegmentationDataset 的配对机制保证同一索引取出的图像和掩膜来自同一个样本getitem里按 idx 同时取 image 和 mask再统一做 resize 和类型转换。这里有个隐藏约定图像 resize 用双线性插值掩膜 resize 必须用最近邻插值。掩膜是离散类别双线性插值会在类别边界上生成 1.5、2.7 这种非整数标签后续 np.eye 直接越界报错。one-hot 转换是这部分重点也最容易翻车def __getitem__(self, idx): image self._load_image(self.image_paths[idx]) # (H, W, 3)float32 mask self._load_mask(self.mask_paths[idx]) # (H, W)int64 mask self._map_labels(mask) if self.augment: image, mask self._sync_augment(image, mask) # 动态 one-hot(H, W) - (H, W, C) - (C, H, W) one_hot np.eye(self.num_classes)[mask].transpose(2, 0, 1).astype(np.float32) image image.transpose(2, 0, 1).astype(np.float32) return torch.from_numpy(image), torch.from_numpy(one_hot)np.eye(num_classes)[mask] 是 NumPy 的花式索引把每个像素值替换成对应的单位向量一步到位生成 one-hot比循环快得多。之后 transpose 把通道维挪到最前面符合 PyTorch 的 (C, H, W) 约定。注意 one-hot 之后 mask 转成 float32而原始 mask 是 int64这是为了后面 Dice 计算时能和 softmax 输出直接相乘PyTorch 里两个张量类型不一致会直接报错提前转 float32 省一个麻烦。3.3 数据增强同步与 get_data_loaders训练和验证必须同一套预处理医学图像分割的数据增强必须图像掩膜同步做。随机水平翻转时 image 左右翻mask 也必须左右翻否则模型学到的对应关系是错位的。固定随机种子让两个变换用同一个随机状态是常见做法def _sync_augment(self, image, mask): seed random.randint(0, 2**32) torch.manual_seed(seed) image transforms.RandomHorizontalFlip(p0.5)(torch.from_numpy(image)) torch.manual_seed(seed) mask transforms.RandomHorizontalFlip(p0.5)(torch.from_numpy(mask)) return image.numpy(), mask.numpy()两个变换用同一个种子重置随机源翻不翻就完全一致。这段代码演示的是同步增强的核心思路换成随机旋转、随机裁剪同理。get_data_loaders 把训练集和验证集分开构建关键点在于验证集不启用增强但标准化参数必须和训练集一致否则验证时输入分布和训练时不同Dice 会莫名下降def get_data_loaders(train_imgs, train_masks, val_imgs, val_masks, num_classes, batch_size8, label_mapNone): train_ds SegmentationDataset(train_imgs, train_masks, num_classes, label_maplabel_map, augmentTrue) val_ds SegmentationDataset(val_imgs, val_masks, num_classes, label_maplabel_map, augmentFalse) train_loader DataLoader(train_ds, batch_sizebatch_size, shuffleTrue, num_workers4, drop_lastTrue) val_loader DataLoader(val_ds, batch_sizebatch_size, shuffleFalse, num_workers4, drop_lastFalse) return train_loader, val_loader标准化用的 ImageNet 均值方差对医学图像不一定最优但胜在通用图像归一化之后掩膜不做归一化保持原始标签值。Windows 下 num_workers 超过 0 时记得把训练脚本入口包在 ifname main 里否则多进程 DataLoader 会递归加载报错这是 PyTorch 在 Windows 上的经典翻车点。4. 训练闭环Dice 评估、双损失适配与断点续训的工程细节4.1 Dice 系数平滑因子不是玄学是除零保护医学图像分割的评估指标里 Dice 比 IoU 更常用因为它对前景占比小的区域更敏感。train_utils.py 里的实现关键在 smooth 参数def dice_coeff(pred, target, smooth1e-5): # pred 需先经过 sigmoid二分类或 softmax多分类 inter (pred * target).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target.sum(dim(2, 3)) dice (2 * inter smooth) / (union smooth) return dice.mean().item()smooth 的作用是防止除零如果某个 batch 里全是背景pred 和 target 的和都是 0分母为 0Dice 直接变 NaN。加一个小常数 1e-5同时加在分子分母上让全背景的 batch 返回接近 1 的数值而不是无穷大。另一个细节是 pred 必须先过激活函数。训练循环里算损失时用 logitsBCEWithLogitsLoss 内部自带 sigmoid但算 Dice 时必须手动把 logits 转成概率。如果拿 logits 直接算 Dice数值会在一个完全不同的尺度上得到的指标没有参考意义。多类别时我一般会算每个类别包括背景的 Dice 再取平均训练时看的这个数和最终提交指标口径一致避免线上线下对不上。4.2 损失函数自动适配二分类和多分类的 target 形状完全不同train_utils 里按 num_classes 自动选损失函数def build_criterion(num_classes): if num_classes 1: # BCEWithLogitsLoss 内部含 sigmoid数值更稳定 return nn.BCEWithLogitsLoss() else: return nn.CrossEntropyLoss()二分类走 BCEWithLogitsLoss模型输出 1 通道target 是 (B, 1, H, W) 的 float one-hot多分类走 CrossEntropyLoss模型输出 C 通道target 是 (B, H, W) 的 long 索引不带通道维。两者 target 的形状和类型都不一样这是新手最容易踩的类型坑。BCEWithLogitsLoss 优先于「sigmoid BCELoss」是因为它把 sigmoid 和交叉熵做了数值融合logits 绝对值很大时不会出现 log(0) 导致 NaN。CrossEntropyLoss 同理内部自带 log_softmax。如果之后想换带权重的损失比如给前景类别更高权重CrossEntropyLoss 直接传 weight 参数就行BCEWithLogitsLoss 要自己构造 pos_weight这是两个 API 的差异点。4.3 断点续训只存模型权重等于没存train_utils 里断点续训的 checkpoint 同时存模型、优化器和 epoch这一点非常关键。只加载模型权重继续训练优化器的动量状态丢了训练会有一段「重新热身」过程loss 不降反升。def save_checkpoint(state, filename): torch.save({ model: state[model].state_dict(), optimizer: state[optimizer].state_dict(), epoch: state[epoch], best_dice: state[best_dice] }, filename) def load_checkpoint(model, optimizer, filename): ckpt torch.load(filename, map_locationcpu) model.load_state_dict(ckpt[model]) optimizer.load_state_dict(ckpt[optimizer]) return ckpt[epoch] 1, ckpt[best_dice]load 之后返回 start_epoch epoch 1训练循环从断点接着跑。map_locationcpu 是先加载到 CPU 再转目标设备避免 GPU 显存里残留旧 tensor。恢复时同时恢复 best_dice早停机制才不会误判如果只恢复模型不恢复 best_dicebest_dice 重置为 0验证 Dice 很容易超过它导致模型被错误地覆盖保存。4.4 train.py 主控早停、设备检测与训练曲线落盘train.py 是主控模块所有训练参数都能通过命令行覆盖典型用法python train.py --data_root ./data --num_classes 2 \ --use_separable --loss bce --batch_size 8 \ --epochs 100 --lr 1e-3 --resume ./checkpoints/best.pth--use_separable 是 store_true 类型不加这个参数模型就跑标准卷积--resume 传路径进入断点续训。设备检测用 torch.cuda.is_available() 自动选 GPU 或 CPU不用手动改代码。训练循环里每个 epoch 结束输出平均 loss 和 Dice验证阶段先 model.eval() 再包 torch.no_grad()这对显存节省非常明显。早停逻辑是验证 Dice 连续 patience 个 epoch 不提升就终止同时始终保存验证 Dice 最高的模型best_dice 0 patience 10 for epoch in range(start_epoch, args.epochs): train_loss, train_dice run_train_epoch(...) val_dice run_eval_epoch(...) if val_dice best_dice: best_dice val_dice patience args.patience save_checkpoint({model: model, optimizer: optimizer, epoch: epoch, best_dice: best_dice}, best.pth) else: patience - 1 if patience 0: print(fearly stop at epoch {epoch}) break训练曲线用 matplotlib 同时画 loss 和 Dice横轴是 epoch双语标题方便中英文习惯不同的同事都能直接看。曲线图自动存 PNG训练完翻文件就能复盘不需要等整个训练结束才确认结果。5. 避坑指南跑通这套 UNet 训练脚本的四个典型问题5.1 验证 Dice 虚高保存的预测图却全黑现象训练时验证 Dice 能到 0.9但保存出来的预测掩膜全黑一张都看不到前景。原因最常见的是输出通道数和保存逻辑不匹配。二分类任务用了 num_classes1 的 BCE 模式模型输出只有 1 个通道保存时却用 argmax 取最大索引单通道张量 argmax 恒为 0全图自然全黑。解决保存预测图前先判断模型是单通道还是多通道if pred.shape[1] 1: mask (torch.sigmoid(pred) 0.5).squeeze(1) else: mask torch.argmax(pred, dim1)这个判断看起来简单但很容易被忽略我见过有人调了一晚上最后发现只是保存逻辑写错。5.2 mask 标签映射后出现越界或 loss 为 NaN现象训练刚开始 loss 就是 NaN或者 one-hot 转换时报 index 超出范围。原因掩膜里存在 label_map 没覆盖的标签值。比如标注文件里混入 255 这类边界值np.eye(num_classes)[mask] 下标直接越界或者映射后出现负数标签CrossEntropyLoss 不接受负标签。解决加载数据后先加一段全量检查for mp in mask_paths: m np.array(Image.open(mp)) uniq np.unique(m) if uniq.max() num_classes - 1 or uniq.min() 0: print(mp, uniq) # 定位问题样本宁可多跑这一段扫描也别在训练中段才发现数据脏。医学标注文件偶尔会有杂色边或画笔残留这类问题在数据检查阶段解决成本最低错一个标签往往要白跑几十个 epoch。5.3 断点续训后 loss 反弹早停被误触发现象resume 之后第一个 epoch 的 loss 明显高于上次保存时的值然后连续几个 epoch 不涨patience 耗尽提前退出。原因大多数情况是断点只存了模型权重优化器的动量状态和调度器状态全部丢失训练重新「预热」。另外 DataLoader 默认 shuffle新 epoch 数据顺序变化会带来波动但幅度远小于优化器状态丢失的影响。解决checkpoint 里按 train_utils 的写法把 optimizer.state_dict() 一并存进去加载时恢复如果用了 lr scheduler把它的 state_dict 也存上。恢复之后先跑一个 epoch 观察 loss 是否回到存档水平再判断早停是否继续不要第一个 epoch 掉一点就手动停掉。5.4 CUDA out of memory显存到底消耗在哪现象batch_size8 都跑不起来提示 CUDA out of memory。原因这套代码在 256×256 输入下显存主要花在 encoder 中间层和 decoder 的跳跃连接上。通道数开到 1024 之后半分辨率层的特征图是 128×128 乘通道数batch 一大很容易爆。验证阶段如果忘了包 torch.no_grad()显存会额外多一份计算图。解决先用 use_separableTrue 跑显存通常能降一截再把 batch_size 从 8 降到 4 或 2确认能跑之后再逐步往上加。验证阶段务必 model.eval() torch.no_grad()这是最基本也最有效的一手。如果还是不够考虑梯度累积或混合精度不建议暴力把输入尺寸压到 128 以下分割任务对分辨率很敏感图片一缩细节就丢了。6. 换到自己的数据集三个必改点加一次快速冒烟验证拿到这套代码最常干的事是迁移到自己的医学数据集上。三个必改点和一次冒烟验证照做基本不会翻车。第一个改点是 label_map 和 num_classes。自己的掩膜是什么标签体系打印 np.unique 之后对着写映射字典num_classes 必须是映射后最大的目标类别数加一不要拿原始标签的最大值去填。第二个改点是 get_data_loaders 的路径参数确保 train_imgs 和 train_masks 一一配对、长度一致不一致时先对齐再进 Dataset配对错位是最隐蔽的错误训练时不报错但指标一直上不去。第三个改点是损失函数的选择纯前后景二分类用 bce多组织分割用 ce框架已经按 num_classes 自动选了只要保证 num_classes 传对就行。冒烟验证我一般这么跑只挑 2~3 张图epochs 设 5batch_size 设 1开着 use_separable 直接训练。如果 5 个 epoch 内训练 loss 明显下降、Dice 从 0 升到 0.5 以上说明模型、数据、损失三者对得上可以放心开全量。如果 loss 纹丝不动多半是数据配对或标签映射出了问题这时候逐一排查比直接调参效率高得多。验证保存的预测图时用前面 5.1 里那个单通道多通道分支判断避免又一次出全黑图。注意换数据集后先检查掩膜的 dtype 和最大值这两个字段最容易在数据导来导去的过程中被改变。我自己踩过最狠的一次是把一个含 5 类组织的数据集拿过来忘了改 label_map训练了两天发现第三类组织从来没被预测出来过查代码才发现原始标签里有两个类别映射到了同一个目标类。从那以后我每次换数据集都强制走一遍三查np.unique 查标签length 对齐查配对5 个 epoch 冒烟查收敛。这个习惯救过我不止一次希望帮到你。本文还有配套的精品资源点击获取
返回列表