资讯动态

DINOv2+原型网络:少样本医学图像分割的实战指南

发布时间:2026/9/13 1:40:30 来源:尧图企业网站定制
简介面向医学图像分割研究者与深度学习开发者这份资源提供了基于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, n6) # 取倒数第二层提取最后一层CLS token的注意力权重 attentions model.get_last_selfattention(input_tensor) # shape: [1, heads, tokens1, tokens1] attention_map attentions[0, :, 0, 1:].mean(dim0).reshape(37, 37).numpy() plt.imshow(attention_map, cmapjet) plt.axis(off) plt.savefig(attn_map.png, dpi150, bbox_inchestight)这段代码里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_channels384, num_classes2): super().__init__() # 输入是DINOv2输出的patch token序列 self.conv1 nn.Conv2d(in_channels, 256, kernel_size3, padding1) self.norm1 nn.GroupNorm(8, 256) self.conv2 nn.Conv2d(256, 128, kernel_size3, padding1) self.norm2 nn.GroupNorm(8, 128) # 两倍上采样到74x74 self.upsample1 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.conv3 nn.Conv2d(128, 64, kernel_size3, padding1) # 从74x74上采样到148x148再插值到512x512交给loss处理 self.upsample2 nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.seg_head nn.Conv2d(64, num_classes, kernel_size1) 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到8BatchNorm统计量不稳定。训练时优化器选择也很关键。冻结阶段用AdamW、学习率3e-4weight 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, devicesupport_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, dim0) # [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, dim1) p_norm F.normalize(prototypes, dim1) # [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_class1): 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训练时用组合lossDice 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 python3.10 -y conda activate medical_dino # 安装核心依赖 pip install torch2.1.2 torchvision0.16.2 --index-url https://download.pytorch.org/whl/cu118 pip install einops timm monai1.3.0 opencv-python # 从本地权重加载DINOv2 python -c import torch from torchvision.models import vit_b_14 model vit_b_14(weightsNone) state_dict torch.load(dinov2_vits14_pretrain.pth, map_locationcpu) # 去掉头部的分类层只保留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_onlyTrue), EnsureChannelFirstd(keys[image, label]), ScaleIntensityRanged(keys[image], a_min-200, a_max400, b_min0.0, b_max1.0, clipTrue), Resized(keys[image, label], spatial_size(512, 512), mode(bilinear, nearest)), RandFlipd(keys[image, label], spatial_axis0, prob0.5), RandRotate90d(keys[image, label], prob0.5, max_k3), RandSpatialCropd(keys[image, label], roi_size(448, 448), random_sizeFalse), ] dataset Dataset(datafiles, transformtransforms) dataloader DataLoader(dataset, batch_size4, shuffleTrue, num_workers4)CT图像的窗宽窗位设置直接影响DINOv2的输入分布。肺部CT建议窗位-600、窗宽1500区间约在-1350到150代码里ScaleIntensityRanged的a_min-200, a_max400对应软组织窗如果做骨结构分割就要把窗位拉到400。Resized用双线性插值做图像、最近邻插值做标签前者对齐像素值有平滑作用后者防止标签产生训练集里不存在的插值灰度导致Dice计算失真。4.3 训练脚本和loss实现少样本训练当前的batch_size建议取4或6一个episode的support和query共享同一个batch。如果batch过大query图像的特征会被平均值稀释标签噪声的直接贡献也更高。class CombinedLoss(nn.Module): def __init__(self, dice_weight0.6, focal_weight0.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, dim1) targets_one_hot F.one_hot(targets.long(), num_classeslogits.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(), reductionnone) 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_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max200, eta_min1e-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_classes2) 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_norm1.0) optimizer.step() scheduler.step()冻结阶段把backbone的lr设成0但AdamW仍然会维护一阶和二阶动量这是很多人容易弄错的地方即便lr为0momentum状态也在更新解冻后会立即产生跳变。完全冻结应该用model.parameters()排除backbone之外的方式或者给backbone参数设置requires_gradFalse。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_locationdevice) 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次取平均指标。n12的 случайных切分会产生极大的方差一次验证的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_iter10): # 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如果模型分割结果仍然大致有合理形状说明模型过拟合了空间先验而不是真正在辨认结构。再一个验证办法是把支持集里某一张图的标注换成一个完全不相关的物体形状看模型预测是否跟着剧烈改变。如果不变说明原型计算根本不依赖标注只是学到了某个静态背景偏置。这两项测试通过后再谈论部署不迟。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价