简介:面向医学图像分割研究团队与算法工程师的自监督双路径网络实现项目,针对CT、MRI等影像中病灶区域自动分离问题,提供了一套无需大量标注即可完成特征学习的完整实践方案。压缩包共19个文件,含FuseNet本地与Colab两个ipynb版本、模型训练工具脚本utils.py和model_utils.py、样例输入及分割效果对比图(bmp/png)、requirements.txt及README说明,整体仅1.34MB,结构紧凑便于快速上手。项目采用自监督学习策略预训练双路径网络,一条路径提取全局上下文,另一条保留局部细节,有效缓解标注数据不足的痛点;代码内附图像预处理与后处理模块,支持替换自有数据集进行迁移或二次开发,运行notebook即可还原端到端分割流程。已有135人学习下载,适合具备一定深度学习基础、希望掌握医学图像分割算法的中高级开发者参考与复用。
1. 从“标注不够”到“自监督双路径”:这个项目在解决什么问题
拿到的这个项目标题里挂着医学图像分割、自监督、双路径网络和项目源码四个关键词,实际上对应的是医疗AI落地中最常见的两难:标注数据稀缺,但分割模型又对边界和上下文极度敏感。自监督学习让模型可以在无标注的CT、MRI、皮肤镜影像上先学一遍通用特征,再用很少的标注样本微调;而双路径网络则从结构上让模型既能看到“整个器官在哪”,又能看清“病灶边缘在哪”。这个标题里的项目源码正是把这两件事封装在一起,适合读研读博做医学影像方向、或者刚进医院AI团队做落地的工程师。下面我从原理到代码,再到踩坑,把这条技术路线完整拆一遍。
2. 自监督双路径网络:为什么两条腿比一条腿稳
2.1 医学图像分割的标注瓶颈,以及自监督为什么能切入
医学图像分割对标注质量的要求远高于自然图像。一张尺寸为512×512的CT切片,医生需要用多边形逐层勾画器官轮廓,平均耗时以小时计。三维数据则更夸张,一个肝脏的精细标注可能需要两三天。所以医疗影像项目里,标注样本从几十例到几百例是常态,完全不足以训练一个从头开始的深层UNet。
自监督学习正是为了缓解这个矛盾。它的核心思路是:不依赖人工标签,而是从影像自身设计一个代理任务。模型在完成这个代理任务的过程中,能学到解剖结构、纹理走向、边界连续性这类通用特征。之后再用带标注的数据做微调,只用原来十分之一的样本就能达到可用效果。在医学图像上,常见的代理任务包括对比学习、掩码重建、旋转预测、灰度变换预测。其中对比学习在自然图像上效果最突出,但医学图像因为器官形态和位置相对固定,选择正负样本时反而要格外小心,不然会学到“模态不变性”以外的噪声。
2.2 双路径网络的结构逻辑:上下文路径与细节路径
双路径网络不是新概念,它最早火起来是在实时语义分割领域,代表作是BiSeNet。它的设计动机非常直观:一个分割网络在编码阶段不断下采样,能获得大感受野,但代价是丢失空间细节;如果不做下采样,每个像素都保留,显存又扛不住。与其在一个网络里纠结,不如把两条路径拆开。
上下文路径使用带stride的卷积或池化快速降低分辨率,比如一路卷积到原图的1/8或1/16,感受野覆盖整个器官;细节路径则保持高分辨率(通常是原图的1/2或1/4),浅层特征直接保留边缘信息。两条路径在最后通过一个融合模块合并,让预测结果既能定位整个病灶区域,又能保持边界锐利。
在医学图像里,我一般会把上下文路径设计成类似ResNet的前几层,细节路径则用轻量的卷积分支。这样做的理由是:医学影像的背景相对统一,器官位置也有先验,重上下文路径太深容易过拟合,轻量一点反而泛化更好。最后融合时,要充分考虑两者分辨率的差异,常见做法是把上下文路径上采样后逐元素相加,或者用空间注意力做加权融合。
2.3 自监督预训练与双路径结合的三种训练策略
自监督不是只能放在预训练阶段。第一种策略是两段式:先在大量无标注医学影像上对整套双路径网络做代理任务预训练,然后冻结编码器的一部分,只微调分割头。第二种策略是端到端联合训练:在分割损失基础上加上一个辅助重建损失,让网络在训练时同时学语义和细节。第三种策略更适合半监督场景:先用无标注数据做对比学习得到预训练权重,再切回经典UNet结构做微调,双路径只在预训练阶段出现,相当于一个特征提取器。
这三种策略里,两段式最稳定,也最好复现。项目源码里如果只含一套训练流程,往往是第一种。至于选择哪条路径做冻结,我的经验是:冻结上下文路径,微调细节路径。因为上下文路径学的是语义类别,和分割任务强相关,预训练后已经很可靠;细节路径学的纹理边缘更依赖具体数据集,需要继续更新。这个细节在微调阶段非常影响上手速度。
3. 搭建双路径分割网络:从模型定义到自监督损失
3.1 定义双路径编码器:让两条路径各司其职
这里我用PyTorch写一个简化版本,目的是展现双路径的关键操作,而不是复刻完整ResNet。实际项目里可以替换成预训练好的backbone,但结构上保持两条路径独立。
import torch import torch.nn as nn import torch.nn.functional as F class ConvBlock(nn.Module): def __init__(self, in_ch, out_ch, stride=1): super().__init__() self.conv = nn.Conv2d(in_ch, out_ch, 3, stride, padding=1, bias=False) self.bn = nn.BatchNorm2d(out_ch) self.relu = nn.ReLU(inplace=True) def forward(self, x): return self.relu(self.bn(self.conv(x))) class DualPathEncoder(nn.Module): # 上下文路径:快速下采样,输出低分辨率高语义特征 # 细节路径:保持高分辨率,输出浅层纹理特征 def __init__(self, in_ch=3, base_ch=32): super().__init__() # 上下文路径,逐步下采样到1/8 self.context_stem = nn.Sequential( ConvBlock(in_ch, base_ch, stride=2), ConvBlock(base_ch, base_ch * 2, stride=2), ConvBlock(base_ch * 2, base_ch * 4, stride=2), ) # 细节路径,只做一次下采样,到1/2 self.detail_stem = nn.Sequential( ConvBlock(in_ch, base_ch, stride=2), ConvBlock(base_ch, base_ch), ConvBlock(base_ch, base_ch), ) def forward(self, x): # x: (B, C, H, W) ctx = self.context_stem(x) # (B, 4*base_ch, H/8, W/8) det = self.detail_stem(x) # (B, base_ch, H/2, W/2) return ctx, det逻辑说明:模型输入是原始或预处理后的图像,ctx表示上下文路径输出,det表示细节路径输出。上下文路径用三次stride=2卷积把分辨率降到1/8,通道数从32涨到128,语义信息更抽象;细节路径保持在1/2分辨率,通道数不变,保留高分辨率细节。参数说明:base_ch控制网络宽度,医学图像通常用16或32,数据集大时可以升到64;in_ch在灰度CT或MRI上设为1,在彩色内镜或皮肤镜图像上设为3。
3.2 自监督预训练任务:对比损失与重建损失怎么选
如果手上有大量无标注原始影像,推荐先做掩码重建。医学图像的纹理和结构有强局部相关性,掩码重建能逼着路径学习可迁移的解剖先验。这里用一个简化的掩码重建任务来演示。
class MaskReconstructLoss(nn.Module): def __init__(self, mask_ratio=0.75): super().__init__() self.mask_ratio = mask_ratio def forward(self, encoder, x): # x: 原始医学图像 (B, C, H, W) B, C, H, W = x.shape # 生成随机掩码,将75%的像素块置零,模拟信息缺失 mask = torch.rand(B, 1, H, W, device=x.device) < self.mask_ratio x_masked = x * (~mask).float() # 双路径编码器提取特征 ctx, det = encoder(x_masked) # 一个简单的重建头:把高低分辨率特征上采样后拼接再卷积 ctx_up = F.interpolate(ctx, size=(H, W), mode='bilinear', align_corners=False) det_up = det # 已经是1/2分辨率,再上采样一次 det_up = F.interpolate(det_up, size=(H, W), mode='bilinear', align_corners=False) feat = torch.cat([ctx_up, det_up], dim=1) recon = self.recon_head(feat) # 这里recon_head需要预先定义 # 只计算被掩码区域的L2损失 loss = F.mse_loss(recon * mask, x * mask) return loss逻辑说明:MaskReconstructLoss先生成一个与输入同尺寸的随机掩码,比例高达75%,打乱大部分像素后送入双路径编码器。由于细节路径只下采样一次,重建时仍能保留局部纹理;上下文路径则提供全局结构信息。然后两路特征上采样回原尺寸拼接,送入重建头输出重建图像,最后只对掩码区域计算MSE损失。参数说明:mask_ratio是关键,医学图像上我建议调到0.75~0.85,过高会导致任务太难,过低则学不到全局语义。
3.3 分割头与联合训练:让两条路径的输出融合
预训练结束后,需要把模型切成“编码器 + 融合模块 + 分割头”。融合模块的目的是让上下文路径的语义信息和细节路径的边界信息互相补全。
class FusionModule(nn.Module): def __init__(self, ctx_ch, det_ch, out_ch=64): super().__init__() # 先把上下文路径上采样到细节路径分辨率 self.ctx_conv = nn.Conv2d(ctx_ch, out_ch, 1) self.det_conv = nn.Conv2d(det_ch, out_ch, 1) # 空间注意力,学习哪些位置该看细节路径 self.attn = nn.Sequential( nn.Conv2d(out_ch * 2, 1, 3, padding=1), nn.Sigmoid() ) def forward(self, ctx, det): # ctx: 1/8分辨率,det: 1/2分辨率 ctx = F.interpolate(ctx, size=det.shape[2:], mode='bilinear', align_corners=False) ctx = self.ctx_conv(ctx) det = self.det_conv(det) combined = torch.cat([ctx, det], dim=1) attn = self.attn(combined) # 注意力权重,0-1 out = ctx * (1 - attn) + det * attn return out逻辑说明:融合模块先通过双线性插值把上下文路径上采样到与细节路径相同分辨率,再用1×1卷积统一通道数。两组特征拼接后生成一个空间注意力图,注意力高表示更信任细节路径,低则更信任上下文路径。这种方式比直接相加更灵活。参数说明:out_ch决定后续分割头的通道数,通常设64或128;attn里用3×3卷积融合邻域信息,避免单一像素的误判。
4. 让模型真正跑起来:数据预处理、训练脚本与关键参数
4.1 从公开数据集到输入张量:预处理与增强
医学图像数据集的格式千差万别,有的是nii.gz三维文件,有的是png二维切片。我建议先把所有数据统一成一个接口,再喂给训练脚本。这里用二维切片为例,因为大多数双路径实现都是基于2D切片逐层处理的。
import os import cv2 import numpy as np import torch from torch.utils.data import Dataset class MedicalSliceDataset(Dataset): def __init__(self, image_dir, mask_dir, target_size=(256, 256), train=True): self.image_paths = sorted(os.listdir(image_dir)) self.mask_paths = sorted(os.listdir(mask_dir)) self.target_size = target_size self.train = train def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img = cv2.imread(os.path.join(self.image_dir, self.image_paths[idx]), cv2.IMREAD_GRAYSCALE) mask = cv2.imread(os.path.join(self.mask_dir, self.mask_paths[idx]), cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, self.target_size, interpolation=cv2.INTER_LINEAR) mask = cv2.resize(mask, self.target_size, interpolation=cv2.INTER_NEAREST) # 将uint8转成float,并归一化到[0,1] img = img.astype(np.float32) / 255.0 mask = (mask > 127).astype(np.float32) if self.train: # 数据增强:随机翻转和随机旋转,增强分割模型的平移不变性 if np.random.rand() > 0.5: img = cv2.flip(img, 1) mask = cv2.flip(mask, 1) angle = np.random.randint(-10, 10) if angle != 0: M = cv2.getRotationMatrix2D((self.target_size[0] // 2, self.target_size[1] // 2), angle, 1.0) img = cv2.warpAffine(img, M, self.target_size, flags=cv2.INTER_LINEAR) mask = cv2.warpAffine(mask, M, self.target_size, flags=cv2.INTER_NEAREST) # 转换为PyTorch张量:1个通道在首位 img = torch.from_numpy(img).unsqueeze(0) mask = torch.from_numpy(mask).unsqueeze(0) return img, mask逻辑说明:这个数据集类把图像和mask读成灰度图,统一resize到256×256。训练模式下加入随机翻转和旋转,配合增强能有效减少过拟合。参数说明:interpolation=cv2.INTER_LINEAR用于图像,cv2.INTER_NEAREST用于mask,这是最容易被忽视的细节——如果用线性插值去缩放mask,会在边界产生介于0和1之间的小数,导致训练时模型被迫学习“模糊边界”,最终预测结果也会糊成一片。
4.2 自监督预训练阶段的训练循环
预训练阶段不参与分割损失,只做掩码重建。这里给一个标准的训练循环框架,重点是学习率和优化器选择。
def train_pretrain(model, dataloader, loss_fn, epochs=200, lr=1e-4): optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) for epoch in range(epochs): model.train() total_loss = 0.0 for x in dataloader: # 注意:预训练阶段不需要标签,只需要图像 x = x.to(device) optimizer.zero_grad() loss = loss_fn(encoder=model, x=x) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() scheduler.step() if (epoch + 1) % 20 == 0: print(f"Epoch {epoch+1}/{epochs}, Loss: {total_loss/len(dataloader):.4f}")逻辑说明:整个预训练循环不需要任何mask标签,只需把原始图像传入loss_fn。使用AdamW是因为它在医学图像这种小数据量任务上通常比SGD稳。余弦退火学习率可以避免在预训练末期震荡。梯度裁剪clip_grad_norm_是防止重建任务在训练初期出现大梯度。参数说明:预训练epoch建议200~500,取决于数据量;lr=1e-4是通用起始点,如果损失不降,可以调到3e-4,但不要超过1e-3。
4.3 微调阶段与关键超参:冻结什么、学习率多少
微调阶段的分割头通常随机初始化,编码器则加载预训练权重。关键超参归纳如下表。
| 超参数 | 推荐值 | 说明 |
|---|---|---|
| 冻结层 | 冻结上下文路径前两层 | 保留通用语义特征 |
| 微调学习率 | 1e-4~3e-4 | 太高会破坏预训练特征 |
| 分割头学习率 | 1e-3 | 新初始化的层需要更快收敛 |
| 总损失 | 0.6 DICE + 0.4 Focal | 稀疏标签时更稳 |
| batch size | 8~16(256×256) | 显存不够则减少 |
| 训练epoch | 100~150 | 配合早停策略 |
微调循环和预训练的区别在于输入带mask标签,损失是分割损失。这时可以给编码器和分割头分别设置不同学习率,用两个优化器,或者直接用一个优化器但通过param_groups分组。我一般这么做:
def set_param_groups(model): encoder_params = [] head_params = [] for name, param in model.named_parameters(): if 'fusion' in name or 'head' in name: head_params.append(param) else: encoder_params.append(param) return encoder_params, head_params逻辑说明:这个分组函数会把融合模块和分割头的参数挑出来,其余都算编码器参数。然后可以在优化器里给两组参数分配不同的学习率。参数说明:是否需要分三组(上下文、细节、头)可以后续再调,但多分组会让训练脚本难维护,收益未必明显。先分两组足以解决“预训练特征被淹没”的问题。
5. 训练医学分割模型的5个血泪坑:现象、原因与排查
5.1 显存溢出:OOM发生在自监督阶段
现象:预训练或微调刚跑几步,PyTorch直接抛出CUDA out of memory,程序崩溃。
原因:双路径网络本身就比单路径多一份显存开销。很多医学图像输入分辨率是512×512甚至更高,细节路径又在高分辨率下计算,显存占用容易被瞬间打满。也有可能是torch.utils.checkpoint没有用,导致反向传播缓存了太多中间特征。
解决:先把batch size降到2或4,确认能跑通后再逐步增加。其次用混合精度训练,torch.cuda.amp可以在不降低精度的前提下大幅省显存。另外,细节路径的通道数可以减半,因为边缘信息不需要很深的通道。如果还不行,把输入尺寸从512降到384,通常分割精度损失很小。
5.2 自监督预训练损失不收敛,重建结果一直是模糊的
现象:预训练时损失卡在某个值不动,或者下降非常慢,重建出的图像只有轮廓,没有纹理。
原因:掩码比例过高,任务难度超过模型学习能力。另一个常见原因是没有做梯度裁剪,损失震荡触发“毁灭性梯度”,导致BN统计量不稳定。
解决:先把mask_ratio调到0.6验证一下能否收敛,如果收敛了再逐步增加。同时给模型加上torch.nn.utils.clip_grad_norm_,max_norm设为1.0。还要检查输入图像是否归一化到0-1区间,如果输入值是0-255,MSE损失在数值上会大很多,梯度也随之变大。
5.3 融合模块尺寸不匹配,报错信息来自interpolate
现象:前向传播时,F.interpolate(ctx, size=det.shape[2:])报错,提示输入和输出形状不一致,或者出现非预期的空间尺寸。
原因:上下文路径和细节路径的stride设置不一致时,上采样结果不一定严格对齐。比如上下文路径经过三次stride=2得到原图1/8,但如果输入尺寸是256,输出就是32×32;细节路径经过一次stride=2加两次stride=1得到128×128。用size=指定目标尺寸是稳妥的,但有时代码里会误用scale_factor=4,由于整除问题导致偏差。
解决:统一用size=det.shape[2:]而不是scale_factor。另外在定义路径时尽量让下采样倍数明确,比如上下文路径1/8、细节路径1/2,融合时上采样倍数正好是4。如果你在项目源码中看到类似“1/16与1/4融合”的组合,注意上采样倍数不同,但原理一样。
5.4 Dice损失在稀疏标签下训练震荡甚至梯度爆炸
现象:训练时损失出现NaN,或者Dice值一直在0附近跳,模型什么也学不到。
原因:医学图像中病灶区域可能只占整张图的5%甚至更低。Dice损失在前景极稀疏时,2 * intersection和union都很小,梯度数值不稳定。尤其当预测完全为背景时,分母可能趋近于0,导致NaN。
解决:在Dice损失中加一个平滑项eps=1e-5,保证分母不为零。更推荐把Dice和Focal Loss结合,Focal Loss可以压制易分背景样本的梯度,让模型关注少数前景像素。如果病灶区域太小,还可以在数据加载时做前景样本过采样,确保每个batch里至少有一张非空mask的样本。
5.5 预训练权重被微调阶段“洗掉”,效果反而不如随机初始化
现象:预训练后微调100个epoch,验证集Dice比从头训练还低,模型表现像被污染过。
原因:微调阶段使用了过大的学习率,编码器在第一步backward中就被破坏。另一个原因是在预训练和微调之间切换了输入分布:预训练时用0-1归一化,微调时又用了z-score归一化,导致编码器学到的特征完全不匹配。
解决:先冻结整个编码器,只训练融合模块和分割头,跑50个epoch后解冻上下文路径的后半段,用学习率1e-4继续微调。还有一个更省事的办法:微调阶段优化器不要携带预训练的momentum状态,重建一个新的AdamW,相当于给模型一个“后悔药”,避免旧优化器状态带着预训练任务的惯性。
6. 验证你的模型:Dice之外还该看哪些指标,以及一个可视化技巧
模型训练完,直接报一个Dice值往往不够。Dice对边界偏移不敏感,两个相同Dice的模型可能在边缘细节上天差地别。我在实际项目里通常再补两个指标:Hausdorff距离用来衡量最大边界误差,适合关注手术规划的场景;体积相关误差则用来评估器官体积估算,适合放疗剂量计算。
可视化方面,一个最实用的技巧是用Grad-CAM类激活图叠加在原始影像上。不用额外装库,用torch.autograd.grad就能实现。先让模型预测一次,拿到分割头的最后一个特征图,对目标类别(比如病灶)计算梯度,再把梯度做全局平均池化得到权重,最后对特征图加权求和并上采样。叠加原图后,你能直接看出模型是被真实病灶区域激活,还是被图像边缘的伪影欺骗。我第一次用这个技巧检查自己的分割模型时,发现模型居然被CT扫描床的边缘高响应激活了,就是因为训练数据里病灶总出现在图像中心区域,而模型偷懒学了位置。这个发现直接促使我改了数据增强里的随机裁剪策略。
自监督和双路径网络的组合还有更广的延展空间。预训练模型可以继续用半监督方式扩展到更多未标注数据;双路径中的细节路径也可以替换成Transformer分支,用来捕捉长距离依赖。但从工程上手角度,先把当前这套结构跑稳,再考虑这些进阶方向才是正路。我自己的习惯是每次实验都保留一个最简配置的备份,确保任何一次修改失败都能回滚。医学图像分割最怕的不是模型不work,而是你不知道它为什么不work——把可视化、指标、代码版本管理这三件事做好,比堆模型结构有意义得多。希望帮到你。
本文还有配套的精品资源,点击获取