简介:面向医学图像分割研究者与深度学习开发者,这份资源提供了基于DINOv2自监督的少样本分割算法实现,能够在标注数据稀缺的医学场景下,利用自监督特征与少量标注完成精确分割,同时兼顾模型的泛化能力。项目围绕完整训练流程展开,源码涉及数据加载、数据增强、模型骨干、注意力模块、损失函数、评估指标等关键环节,并配有Shell脚本便于一键训练与验证;Jupyter Notebook则用于数据探索、结果可视化与实验对比,方便研究者结合自己的数据集进行二次开发。压缩包共27个文件,其中23个Python源码、2个Shell脚本、1个Markdown说明和1个Notebook,整体仅约86KB,体积轻量、结构清晰,适配主流深度学习环境。目前已有160人学习下载,对少样本医学图像分割、自监督预训练与模型微调感兴趣的开发者,可直接基于这份源码进行算法复现、调试与改进,显著缩短从论文到工程落地的距离。
1. 少样本医学图像分割为什么需要DINOv2这类自监督模型
医学图像分割在临床场景中的痛点从来不是模型结构不够深,而是标注样本太少。一个3D肝脏CT数据集里能拿到几十例带精细标注的病例已经算得上“数据充足”,更多时候只有五例、十例,甚至只能拿到几张切片。传统做法把U-Net反复调参,在十例样本上做到过拟合是常态,换一个采集协议立刻失效。近几年自监督预训练的思路恰好补上这个短板:先用海量无标注数据把视觉特征学出来,再在下游只靠极少标注样本做微调。DINOv2在自然图像上证明了self-distillation可以学到对分割任务有效的密集特征,而且特征本身具备一定的类别区分度,把它迁移到医学图像场景就成了一个值得尝试的路径。
但直接照搬DINOv2到医学图像上通常会碰壁。自然图像和CT、MRI的成像分布差距太大,预训练权重里的底层特征虽然通用,高层的语义却和医学结构对不上。所以实际项目里很少直接拿原版ViT backbone做端到端微调,而是把DINOv2当作特征提取器或初始化权重,再挂一个轻量分割head,或者用少样本学习中常见的原型网络思路。这篇文就围绕这类方案里最常见的一条技术路线展开:用DINOv2自监督预训练权重提取特征,配合原型对齐和轻量解码器,在十例级别的医学图像数据上把分割模型跑起来,同时给出可复现的命令、参数和排错方向。
2. DINOv2自监督原理与医学图像分割的适配点
2.1 自监督预训练为什么比ImageNet监督预训练更适合少样本医学任务
医学图像分割的少样本困境本质上是分布偏移和标注成本的双重问题。ImageNet监督预训练虽然让模型学到了大量物体轮廓和纹理特征,但自然图像里不存在“肺结节在CT上呈现为磨玻璃影”这种密度映射关系。更关键的是,监督预训练强迫模型把特征压缩到1000类分类边界上,特征表达过多服务于类别判别,而医学分割需要的是连续、稠密、对边界敏感的空间特征。
DINOv2的做法是把这个问题换一个解法。它用self-distillation的方式训练ViT,让teacher分支和student分支对同一张图的不同裁剪视图输出一致的特征。训练过程中没有任何类别标签参与,模型被迫从像素级和区域级的一致性中学习空间结构。这个性质对医学图像非常关键:CT值范围、MRI的加权序列差异、超声的噪声纹理,这些底层成像特性不需要语义标签就能被自监督捕捉到。更重要的是,DINOv2的特征图天然保留了位置信息和局部纹理对比度,而这两者恰好是分割任务最依赖的线索。
实际使用时还有一个容易被忽略的点:DINOv2的patch size是14,以CT切片512x512输入为例,输出的特征图是37x37左右。这个空间分辨率直接决定了分割head的设计思路,而不是像U-Net那样从底层就开始逐步恢复分辨率。很多项目在这个地方翻车,拿到特征图直接上采样回原图大小,边界自然糊成一团。
2.2 DINOv2的注意力图可视化与医学图像语义发现
用DINOv2做医学图像分割之前,值得先花半小时看看它的自注意力图到底学到了什么。这一步既是可行性验证,也能帮助确定后续特征提取用哪个层。
import torch from torchvision import transforms from PIL import Image import numpy as np import matplotlib.pyplot as plt # 加载DINOv2 small版本,输出patch token特征 model = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14') model.eval() # 读取一张灰度医学图像,模拟CT切片 img = Image.open('ct_slice.png').convert('L') transform = transforms.Compose([ transforms.Resize((518, 518)), transforms.ToTensor(), transforms.Normalize(mean=[0.5], std=[0.5]) ]) input_tensor = transform(img).unsqueeze(0) with torch.no_grad(): intermediate_output = model.get_intermediate_layers(input_tensor, n=6) # 取倒数第二层,提取最后一层CLS token的注意力权重 attentions = model.get_last_selfattention(input_tensor) # shape: [1, heads, tokens+1, tokens+1] attention_map = attentions[0, :, 0, 1:].mean(dim=0).reshape(37, 37).numpy() plt.imshow(attention_map, cmap='jet') plt.axis('off') plt.savefig('attn_map.png', dpi=150, bbox_inches='tight')这段代码里get_last_selfattention拿到的是最后一层所有head的注意力矩阵,取CLS token对其他token的注意力并做跨head平均,得到的就是模型当前最关注的区域分布。如果输入的是肺部CT,通常能在注意力图上看到高响应区域集中在解剖结构边界附近,比如胸膜线和血管束。如果注意力图完全是一片均匀噪声,说明预训练特征对这个成像域完全不敏感,建议直接放弃迁移,改用医学图像自监督权重或者做更大规模的领域内预训练。
需要注意这里使用get_intermediate_layers和get_last_selfattention时,两个方法都会前向一次模型。调试阶段无所谓,正式训练代码里应该只前向一次,把中间特征和注意力一次取出来,避免双倍显存开销。
2.3 冻结backbone还是微调backbone:少样本下的最优解
这是整个项目里最值得花时间做对比实验的问题。十例训练样本下,微调ViT的所有参数几乎必然导致灾难性过拟合,模型会把训练集噪声当成语义特征。常见做法是分阶段走:第一阶段冻结DINOv2 backbone,只训练分割head;第二阶段用极低学习率解冻最后两三个transformer block。这个策略相当于在避免破坏预训练特征的同时,让深层特征向医学语义做有限度的偏移。
import torch.nn as nn class VitSegHead(nn.Module): def __init__(self, in_channels=384, num_classes=2): super().__init__() # 输入是DINOv2输出的patch token序列 self.conv1 = nn.Conv2d(in_channels, 256, kernel_size=3, padding=1) self.norm1 = nn.GroupNorm(8, 256) self.conv2 = nn.Conv2d(256, 128, kernel_size=3, padding=1) self.norm2 = nn.GroupNorm(8, 128) # 两倍上采样到74x74 self.upsample1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) self.conv3 = nn.Conv2d(128, 64, kernel_size=3, padding=1) # 从74x74上采样到148x148,再插值到512x512交给loss处理 self.upsample2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False) self.seg_head = nn.Conv2d(64, num_classes, kernel_size=1) def forward(self, x): # x: [B, num_patches, C],需要reshape成2D特征图 B, N, C = x.shape H = W = int(N ** 0.5) x = x.permute(0, 2, 1).reshape(B, C, H, W) x = torch.relu(self.norm1(self.conv1(x))) x = torch.relu(self.norm2(self.conv2(x))) x = self.upsample1(x) x = torch.relu(self.conv3(x)) x = self.upsample2(x) return self.seg_head(x)这个head设计刻意跳过了复杂的FPN结构,原因是少样本下参数量就是最大的敌人。Conv1把384维压缩到256维,Conv3只保留64维,整个head参数量不到100万。num_patches对512输入是37x37,两层上采样到148x148,最后在loss计算时用双线性插值把logits对齐到原图尺寸。GroupNorm比BatchNorm更适合少样本场景,因为batch size通常只有4到8,BatchNorm统计量不稳定。
训练时优化器选择也很关键。冻结阶段用AdamW、学习率3e-4,weight decay设为0.05,这个配置对ViT类结构比较稳。解冻阶段学习率降到3e-5,并且只对解冻层生效,通过parameter groups实现。
3. 基于DINOv2特征的原型少样本分割方案设计与实现
3.1 原型网络为什么和DINOv2是天然搭档
少样本分割常见的解决方案分为两类:一类是端到端的元学习,比如MAML和Reptile,通过跨任务梯度更新学习一个易于微调的初始化;另一类是基于度量的原型方法,把支持集的标注区域嵌入到特征空间得到类别原型,查询集样本和原型做相似度比较完成分割。医学图像场景下,元学习的问题在于每个任务的数据量太少,一个episode里的support set可能只有两到三张图,梯度估计的方差大到训练完全不稳定。而且医学图像的类内差异非常大,同一个器官在不同病人体内的形状和灰度分布都不一致,元学习的跨任务泛化优势很难体现。
原型方法对DINOv2来说是顺理成章的组合。DINOv2的特征具备很强的语义一致性,同一器官在不同切片上的特征分布相对紧凑,类原型能够较稳定地刻画其中心。另外原型方法不需要为每个episode维护一套内部的梯度更新状态,推理时只需要一次前向计算和一次特征平均。
具体方案里使用原型对齐的思路:用支持集的粗标注或少量精细标注生成每个语义类别的原型向量,查询图像的每个像素特征和原型向量做余弦相似度,得出像素级的概率图。关键点在于避免计算整图每个像素和原型的点积后直接下采样。医学图像里的目标往往只占整图的很小比例,背景原型会对前景预测形成强偏置,需要用背景抑制策略修正。
3.2 少样本医学图像分割的完整训练流程
实践中最稳的协议是episode training和object-level augmentation的组合。每个episode从训练集里随机采样两个不同的病例,一个作为support set,一个作为query set,支持集提供掩码,查询集要求输出分割结果。对于十例级别的数据,所有病例都既当support又当query,通过随机裁剪和弹性变形增加episode之间的差异度。
import torch import torch.nn.functional as F def compute_prototypes(support_features, support_masks, num_classes): """ support_features: [B, C, H, W] support_masks: [B, H, W],像素值为0~num_classes-1 返回每个类别的原型向量,背景类单独用全局统计 """ prototypes = [] B, C, H, W = support_features.shape for cls in range(num_classes): mask = (support_masks == cls).float() if mask.sum() < 1: prototypes.append(torch.zeros(C, device=support_features.device)) continue # 特征按mask加权平均 mask = mask.unsqueeze(1) # [B, 1, H, W] masked_feat = support_features * mask proto = masked_feat.sum(dim=(0, 2, 3)) / mask.sum(dim=(0, 2, 3)) prototypes.append(proto) return torch.stack(prototypes, dim=0) # [num_classes, C] def prototype_segment(query_features, prototypes): """ query_features: [B, C, H, W] prototypes: [num_classes, C] 返回像素级logits [B, num_classes, H, W] """ B, C, H, W = query_features.shape # 展平特征 q = query_features.view(B, C, -1) # [B, C, H*W] # 计算余弦相似度 q_norm = F.normalize(q, dim=1) p_norm = F.normalize(prototypes, dim=1) # [num_classes, C] # 相似度矩阵 [B, num_classes, H*W] similarity = torch.einsum('n c, b c l -> b n l', p_norm, q_norm) logits = similarity.view(B, -1, H, W) / 0.07 return logitscompute_prototypes里需要特别关注背景类。直接在整张图上平均特征会让背景原型偏向高密度出现的组织类型,导致脏像素也被拉向背景。在实践中用一个空间衰减权重更可取:离标注前景区域越远,背景原型的统计权重越低。
def background_prototype_with_distance(support_features, support_masks, fg_class=1): fg_mask = (support_masks == fg_class).float() # 计算到前景区域的L2距离 dist_map = distance_transform(fg_mask) # 使用scipy.ndimage.distance_transform_edt # 背景权重随距离衰减 bg_weight = torch.sigmoid((dist_map - 20) / 10) masked_feat = support_features * bg_weight.unsqueeze(1) bg_proto = masked_feat.sum(dim=(0, 2, 3)) / bg_weight.sum() return bg_proto训练时用组合loss:Dice loss加带温度缩放的focal loss。Dice loss让预测和真值之间在区域重叠上直接对齐,focal loss对边界像素的困难样本施压。温度为0.07来自CLIP的经验取值,作用是在softmax前放大相似度差异,开太大训练初期梯度消失,太小则导致所有类别概率几乎一致。
3.3 数据增强策略:少样本时代的立身之本
少样本医学图像分割里,增强策略的重要性高于模型结构和loss调参。标准做法是引入随机弹性形变、高斯噪声、强度偏移和cutout。弹性形变的网格sigma取值5到8、平滑系数0.5到1.0之间效果较好;强度偏移按均值为0、标准差为原图灰度标准差10%水平抽样。这些增强只作用于query图像,不作用于support set,否则会破坏原型参考的保真度。Kaiming He团队在自监督对比学习里的经验同样适用:强增强会帮助模型学到更invariant的特征。
DataLoader实现上要注意,同时加载两张图会翻倍内存开销。常见做法是设support batch size为1,保证support path上的梯度不流通,以节约显存。
4. 从零跑通DINOv2少样本医学分割项目:代码与关键参数
4.1 环境搭建和权重准备
整个项目最耗时的一步其实是把DINOv2的权重文件下载下来并正确加载。torch.hub拉取权重需要访问外网,国内环境下经常中途断连。
# 创建conda环境,Python版本不要超过3.10 conda create -n medical_dino python=3.10 -y conda activate medical_dino # 安装核心依赖 pip install torch==2.1.2 torchvision==0.16.2 --index-url https://download.pytorch.org/whl/cu118 pip install einops timm monai==1.3.0 opencv-python # 从本地权重加载DINOv2 python -c " import torch from torchvision.models import vit_b_14 model = vit_b_14(weights=None) state_dict = torch.load('dinov2_vits14_pretrain.pth', map_location='cpu') # 去掉头部的分类层,只保留backbone权重 filtered = {k: v for k, v in state_dict.items() if k.startswith('backbone.')} model.load_state_dict({k.replace('backbone.', ''): v for k, v in filtered.items()}) "这份代码的load_state_dict方式因为在PyTorch官方权重接口里不直接支持facebook的权重格式,才用字符串替换做兼容。如果没有本地权重下载条件,也可以直接从HuggingFace上拉取。
4.2 基于MONAI的医学图像预处理流程
医学图像没有通用文件格式,NIfTI、DICOM切片、mha和PNG序列各有各的坑。用MONAI框架的transforms可以统一读写路径,把大部分格式转换细节屏蔽掉。
from monai.transforms import ( LoadImaged, ScaleIntensityRanged, EnsureChannelFirstd, RandSpatialCropd, RandRotate90d, RandFlipd, Resized ) from monai.data import DataLoader, Dataset import glob files = [{'image': p, 'label': p.replace('image', 'label')} for p in sorted(glob.glob('data/train/img/*.nii.gz'))] transforms = [ LoadImaged(keys=['image', 'label'], image_only=True), EnsureChannelFirstd(keys=['image', 'label']), ScaleIntensityRanged(keys=['image'], a_min=-200, a_max=400, b_min=0.0, b_max=1.0, clip=True), Resized(keys=['image', 'label'], spatial_size=(512, 512), mode=('bilinear', 'nearest')), RandFlipd(keys=['image', 'label'], spatial_axis=0, prob=0.5), RandRotate90d(keys=['image', 'label'], prob=0.5, max_k=3), RandSpatialCropd(keys=['image', 'label'], roi_size=(448, 448), random_size=False), ] dataset = Dataset(data=files, transform=transforms) dataloader = DataLoader(dataset, batch_size=4, shuffle=True, num_workers=4)CT图像的窗宽窗位设置直接影响DINOv2的输入分布。肺部CT建议窗位-600、窗宽1500,区间约在-1350到150,代码里ScaleIntensityRanged的a_min=-200, a_max=400对应软组织窗,如果做骨结构分割就要把窗位拉到400。Resized用双线性插值做图像、最近邻插值做标签,前者对齐像素值有平滑作用,后者防止标签产生训练集里不存在的插值灰度,导致Dice计算失真。
4.3 训练脚本和loss实现
少样本训练当前的batch_size建议取4或6,一个episode的support和query共享同一个batch。如果batch过大,query图像的特征会被平均值稀释,标签噪声的直接贡献也更高。
class CombinedLoss(nn.Module): def __init__(self, dice_weight=0.6, focal_weight=0.4): super().__init__() self.dice_weight = dice_weight self.focal_weight = focal_weight def forward(self, logits, targets): # logits: [B, C, H, W], targets: [B, H, W] in [0, C-1] probs = F.softmax(logits, dim=1) targets_one_hot = F.one_hot(targets.long(), num_classes=logits.shape[1]).permute(0, 3, 1, 2).float() # Dice loss intersection = (probs * targets_one_hot).sum(dim=(0, 2, 3)) union = probs.sum(dim=(0, 2, 3)) + targets_one_hot.sum(dim=(0, 2, 3)) dice_score = (2.0 * intersection + 1e-5) / (union + 1e-5) dice_loss = 1.0 - dice_score.mean() # Focal loss ce_loss = F.cross_entropy(logits, targets.long(), reduction='none') pt = torch.exp(-ce_loss) focal_weight_tensor = (1 - pt) ** 2 focal_loss = (focal_weight_tensor * ce_loss).mean() return self.dice_weight * dice_loss + self.focal_weight * focal_loss # 训练循环片段 optimizer = torch.optim.AdamW([ {'params': backbone.parameters(), 'lr': 0.0}, # 冻结阶段不更新 {'params': head.parameters(), 'lr': 3e-4} ], weight_decay=0.05) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=200, eta_min=1e-6) for epoch in range(200): for support_img, support_mask, query_img, query_mask in train_loader: support_img = support_img.cuda() support_mask = support_mask.cuda() query_img = query_img.cuda() with torch.no_grad(): support_feat = backbone(support_img) query_feat = backbone(query_img) prototypes = compute_prototypes(support_feat, support_mask, num_classes=2) logits = prototype_segment(query_feat, prototypes) loss = criterion(logits, query_mask) optimizer.zero_grad() loss.backward() # 梯度裁剪:抑制ViT深层可能产生的异常梯度 torch.nn.utils.clip_grad_norm_(head.parameters(), max_norm=1.0) optimizer.step() scheduler.step()冻结阶段把backbone的lr设成0,但AdamW仍然会维护一阶和二阶动量,这是很多人容易弄错的地方:即便lr为0,momentum状态也在更新,解冻后会立即产生跳变。完全冻结应该用model.parameters()排除backbone之外的方式,或者给backbone参数设置requires_grad=False。
5. 微调策略、时序集成与验证技巧
5.1 分阶段解冻和判别性学习率
当一个模型在少样本数据上经过初始的全冻结训练后,头部已经能够把DINOv2特征映射到当前任务语义空间。此时从最后一个block开始逐步解冻,每观察训练loss下降趋缓后再解开前面一层。学习率分配按照“离head越近越大”的原则,最后一个block用1e-4,再往前依次乘0.2。
判别性学习率也可以直接用正则化方式替代,比如对backbone参数用更大的weight decay,这样即使解冻,也不会让特征偏离预训练权重分布太远。
5.2 用Model Soup做时序集成提升Dice
少样本训练中模型权重在最优解附近会来回震荡,单次checkpoint通常不是最强泛化点。一个实用技巧是把训练日志里连续5个最低验证loss的checkpoint拿出来做权重平均,也就是Model Soup里的uniform soup。代码上用下面的方式加载多个权重并进行逐层平均:
import copy def model_soup(model_class, ckpt_paths, device): # 先加载第一个权重 model = model_class().to(device) state_dicts = [torch.load(p, map_location=device) for p in ckpt_paths] avg_state = copy.deepcopy(state_dicts[0]) for key in avg_state.keys(): for sd in state_dicts[1:]: avg_state[key] += sd[key] avg_state[key] /= len(state_dicts) model.load_state_dict(avg_state) return model做权重平均时要注意,模型的backbone和head要一次性平均,不能只对head层做。还有,最后一层的bias也参与平均,因为它直接对应分类偏移量。
5.3 验证:leave-one-out交叉验证而不是随机划分
十例训练数据分成train/val根本没有统计意义。更稳做法是leave-one-out,每次拿一例当测试,其余所有当训练,跑N次取平均指标。n=12的 случайных切分会产生极大的方差,一次验证的Dice从0.5跳到0.9完全有可能,而leave-one-out能把这种偏差抹平。
计算Dice时注意背景类一般不看,只看前景类的Dice,避免背景占99%像素导致Dice虚高。还可以额外报告边界Hausdorff距离,少样本模型常在边界上丢细长突起,Dice接近但边界差很远。
6. 边界样本、自定义分割目标和失败模式排查
6.1 遇到没有结构边界的模糊区域怎么办
医学图像里很多结构天然不具备清晰边界,比如肝脏和周围脂肪、脑灰质和白质。DINOv2的patch embedding会把这些区域的特征混合在一起,原型边界上的像素处于特征空间的中间地带,无论怎么调温度都不可能精确分割。这种情况下停止继续调参,直接加一个条件随机场后处理层,用像素邻域的一致性把零散误分区域修整掉。常见用monai里的CRF包装成final post-process。
import SimpleITK as sitk def crf_postprocess(image_path, logits_np, num_iter=10): # image_path: 原始灰度图,logits_np是模型输出的概率图 sitk_img = sitk.ReadImage(image_path, sitk.sitkFloat32) prob_sitk = sitk.GetImageFromArray(logits_np[1]) # 前景概率 # SimpleITK自带的CRF实现比较粗糙,实际项目可用pydensecrf result = sitk.BinaryThreshold(prob_sitk, 0.4, 1.0) return sitk.GetArrayFromImage(result)6.2 在DINOv2特征之上直接调参还是重训backbone
很多项目拿到DINOv2权重后,在少量医学数据上继续做领域自适应预训练,这个做法在样本接近100例时有用,但只有十例时高度容易跑偏。十例样本的统计噪声会直接盖过真正的领域特征,继续预训练只是在拟合噪声。如果真的要做,最稳的限制方式是用masked image modeling的objective,把输入图像随机遮盖60%的patch,重建原始灰度值。这个任务不会强迫模型把语义压缩到类别边界,损失方向也更保守。但建议项目初期跳过这一步,先验证原型方案在冻结特征上的表现。
6.3 结果好的时候,还该检查什么
少样本模型最隐蔽的坑是“语义信息是从图像本身还是从上下文捷径里出来”。做一个简单的消融验证:把测试图像的灰度值随机shuffle,如果模型分割结果仍然大致有合理形状,说明模型过拟合了空间先验而不是真正在辨认结构。再一个验证办法是把支持集里某一张图的标注换成一个完全不相关的物体形状,看模型预测是否跟着剧烈改变。如果不变,说明原型计算根本不依赖标注,只是学到了某个静态背景偏置。这两项测试通过后,再谈论部署不迟。
本文还有配套的精品资源,点击获取