ARTICLE DETAIL

资讯详情

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

医疗图像多任务学习与无监督域自适应:基于ResNet50与ExpertNet的工程实践

医疗图像多任务学习与无监督域自适应:基于ResNet50与ExpertNet的工程实践 简介这是一份结合ExpertNet与Resnet50的多任务学习无监督自适应模型Python项目源码面向医疗图像分析领域的深度学习开发者和研究人员特别适用于标注数据稀缺、难以获得高质量标签的医学影像场景可辅助识别肿瘤、病变等关键特征。压缩包共16个文件体积仅27KB包含10个Python脚本、2个文本文件、1个Markdown文档及许可证等脚本覆盖从数据预处理、自动编码器构造、模型训练到测试的完整流程另有依赖说明和README目录划分清晰便于快速上手。目前已有254人学习使用适合作为多任务学习与无监督医疗图像实验的参考实现。借助源码可深入理解ExpertNet专家模块如何并行处理不同任务以及Resnet50残差连接如何缓解深层网络退化、提取精细图像特征项目提供了完整的无监督自适应训练思路读者可在其基础上复现实验、调整网络结构并为医学影像辅助诊断构建可扩展的基线方案。1. 医疗图像多任务落地的拦路虎标注贵、跨中心掉点这套方案在解决什么医疗影像AI团队普遍经历过这个场景你在A中心用一批标注好的CT数据训出分割和分类模型指标漂亮换到B中心的新设备上Dice掉十几个点同一张影像既要判断病灶良恶性又要勾出病灶边界最后只能拆成两套模型、两套标注、两套迭代流程。这个标题指向的是一套把 ResNet50 当作特征骨架用 ExpertNet 做多任务路由再叠加无监督自适应机制的组合方案——通俗说就是让一个模型在只有影像、没有标注的新数据集上自我校准同时兼顾分类、分割等多个任务。它适合手里有少量有标签医疗图像、但积累了大量未标注数据的团队也适合被多中心迁移问题反复折磨的从业者。2. 拆开标题ExpertNet、ResNet50 和无监督自适应各管哪一段动手写代码之前先把标题里的三个名词拆明白。这个方案不是三个模型的简单拼装而是各有分工ResNet50负责把影像变成通用特征ExpertNet负责在多个任务之间分配特征无监督自适应负责让特征在无标签数据上保持稳定。2.1 ExpertNet 不是固定模型是一种「多专家门控」的解决思路严格说ExpertNet 并不是一个像 ResNet50 那样有标准结构、有预训练权重的固定网络而是一类“多个小专家 门控路由”结构的统称业内叫得更响的是 Mixture of ExpertsMoE。标题里的 ExpertNet 我更愿意理解为作者对这种结构的具体实现命名。它的核心思想不复杂不要用一个容量巨大的网络去硬扛所有任务而是准备若干个结构相同、参数独立的专家子网络前面加一个门控模块对每个输入样本计算一组权重把专家的输出按权重加权融合。医疗图像场景里这个设计天然合适。一个胸部CT样本既要做肺结节分类又要做病灶分割两个任务对特征的侧重不完全相同——分类更关注整体形态和纹理分割更关注边缘和局部对比。如果让一组专家分别擅长这些不同侧面再由门控动态组合理论上能比单个大网络更灵活。门控的公式就是一个带 Softmax 的线性变换输入是编码器输出的全局特征输出维度等于专家数量。2.2 编码主干为什么锁 ResNet50结构图之外的三个理由网上搜 resnet50 网络结构示意图看到的无非是那套经典的残差块堆叠。真正选它当这个方案的主干三个理由比结构图更实在。第一它自带 ImageNet 预训练权重在医疗数据通常只有几百到几千例的情况下用一个在自然图像上学到过纹理和边缘表达的骨干能显著降低过拟合风险这比从头训练一个医疗专用网络稳定得多。第二ResNet50 的层间特征有明确的语义层级底层特征适合边缘和纹理高层特征适合语义判断多任务正好各取所需。第三它是个折中ResNet18 容量不够ResNet101 在医疗大图上显存压力大50 层这个档位在 2D 医疗切片上最常被选用。在实现上我会把 ResNet50 的最后一层全局池化和分类层去掉只保留到 layer4 的输出然后接一个自适应池化把特征压成固定长度的向量作为门控和任务头的输入。这样既保留了预训练参数又把网络改造成了特征提取器。要注意 torchvision 新版接口里 resnet50 的预训练参数从pretrainedTrue换成了weightsResNet50_Weights.IMAGENET1K_V2旧写法在新版会报警告但不影响运行。2.3 多任务学习与无监督自适应怎么嵌在一起方案里的另一个组合动作是「多任务学习 无监督自适应」。多任务学习解决的是标注稀缺问题分类和分割共享同一个编码器标注一份数据两个任务同时受益。无监督自适应解决的是标注错配问题target 域的影像没有标签但我们可以把模型当作一个特征提取器让 source 域和 target 域的特征分布尽量靠近这样训练出的分割头、分类头才能在 target 域上继续生效。两个机制拼在一起的关键是「谁来提供对齐信号」。常见做法是加一个域判别器输入特征输出域类别用对抗训练的梯度反转层让编码器学不到域信息。另一路信号是伪标签一致性对 target 域样本做弱增强和强增强分别预测让两边的分类概率一致。这两路信号都只作用在特征层面不干扰任务头的输出结构所以多任务头照常工作自适应机制在后台完成对齐。下表是各模块的职责划分。模块输入输出在方案里的角色ResNet50 编码器2D 医疗图像特征图 / 特征向量通用特征提取source 与 target 共享ExpertNet 门控全局特征向量专家权重分布按样本动态组合多个专家输出分类头加权后的专家特征类别概率完成良恶性、病灶类型判断分割头加权后的专家特征逐像素预测完成病灶轮廓勾画域判别器编码器特征域类别概率无监督自适应用来混淆 source/target3. 把网络搭出来ResNet50 编码器 ExpertNet 门控 多任务头 域判别器说明思路之后进入正题。下面代码基于 Python 3.8 与 PyTorch 2.x假设已经配置好 Python 环境并安装了 torch 和 torchvision。我在代码里刻意把每个模块拆成独立类方便替换和调试。3.1 整体数据流和模块清单一张 2D 切片进来先经过 ResNet50 编码器得到特征图特征图一方面做全局池化变成向量送入 ExpertNet 门控生成专家权重另一方面送入域判别器做域分类加权后的特征并行送入分类头和分割头。模块之间的依赖关系只有特征维度这一个耦合点所以替换任何一块都不影响其他部分。实现时我一般把特征维度设成 2048因为 ResNet50 的 layer4 输出就是 2048 通道。3.2 门控路由与任务头的关键代码先写 ResNet50 编码器和 ExpertNet 门控。import torch import torch.nn as nn import torch.nn.functional as F from torchvision import models class ResNet50Encoder(nn.Module): 截断的ResNet50编码器输出CxHxW特征图和一个2048维全局向量 def __init__(self, pretrainedTrue): super().__init__() if pretrained: weights models.ResNet50_Weights.IMAGENET1K_V2 backbone models.resnet50(weightsweights) else: backbone models.resnet50(weightsNone) # 去掉最后的全局池化和全连接层 self.features nn.Sequential(*list(backbone.children())[:-2]) self.gap nn.AdaptiveAvgPool2d((1, 1)) def forward(self, x): feat_map self.features(x) # (B, 2048, H/32, W/32) feat_vec self.gap(feat_map).flatten(1) # (B, 2048) return feat_map, feat_vec class ExpertNet(nn.Module): 多专家路由模块门控生成权重加权融合专家输出 def __init__(self, in_dim2048, expert_dim512, num_experts4, tau0.8): super().__init__() self.num_experts num_experts self.tau tau # 每个专家是一个两层MLP self.experts nn.ModuleList([ nn.Sequential( nn.Linear(in_dim, expert_dim), nn.ReLU(inplaceTrue), nn.Linear(expert_dim, in_dim) ) for _ in range(num_experts) ]) # 门控从全局向量生成专家权重加温度参数控制路由锐度 self.gate nn.Linear(in_dim, num_experts) def forward(self, feat_vec): # 训练时加入噪声避免门控过早收敛到某个固定专家 gate_logits self.gate(feat_vec) / self.tau if self.training: gate_logits gate_logits torch.randn_like(gate_logits) * 0.1 gate_weight F.softmax(gate_logits, dim-1) # (B, num_experts) expert_out torch.stack([e(feat_vec) for e in self.experts], dim1) # (B, E, D) out (expert_out * gate_weight.unsqueeze(-1)).sum(dim1) # (B, D) return out, gate_weight说明几个关键参数。tau是门控温度取值范围 0.51.5温度越低 Softmax 越尖锐样本会更快锁定到单一专家温度越高路由越均匀。医疗数据噪声大我习惯设 0.8 左右保留一定随机性。训练阶段给门控加高斯噪声是防止「死专家」的常用手段后面避坑章会再次提到。专家数量num_experts4是个经验值太多专家在数据量小时学不到各自特色太少则路由失去意义。接下来是分类头、分割头和域判别器。class ClassificationHead(nn.Module): 分类头输入专家融合特征输出类别概率 def __init__(self, in_dim2048, num_classes2): super().__init__() self.fc nn.Sequential( nn.Dropout(0.3), nn.Linear(in_dim, 256), nn.ReLU(inplaceTrue), nn.Linear(256, num_classes) ) def forward(self, feat_vec): return self.fc(feat_vec) class SegmentationHead(nn.Module): 轻量分割头从ResNet50特征图恢复分辨率输出逐像素类别 def __init__(self, in_channels2048, num_classes1): super().__init__() self.decoder nn.Sequential( nn.ConvTranspose2d(in_channels, 256, kernel_size2, stride2), # 1/16 nn.ReLU(inplaceTrue), nn.ConvTranspose2d(256, 64, kernel_size2, stride2), # 1/8 nn.ReLU(inplaceTrue), nn.ConvTranspose2d(64, 16, kernel_size2, stride2), # 1/4 nn.ReLU(inplaceTrue), nn.Conv2d(16, num_classes, kernel_size1) ) def forward(self, feat_map): return self.decoder(feat_map)分割头只恢复到原图四分之一分辨率。医疗图像标注通常只对病灶区域有意义四分之一分辨率足够训练损失回传而且能大幅省显存。如果一定要全分辨率输出可以在外侧再加一个双线性插值上采样但显存代价会明显增加。3.3 把四块拼成一个完整模型的组装方式域判别器和梯度反转层用于无监督自适应我把它们和整个网络组装到一起。class GradientReversal(torch.autograd.Function): 梯度反转层前向恒等反向将梯度取反并缩放 staticmethod def forward(ctx, x, lamda): ctx.lamda lamda return x.clone() staticmethod def backward(ctx, grad_output): return -ctx.lamda * grad_output, None class DomainDiscriminator(nn.Module): 域判别器区分特征来自source还是target输入是编码器特征图池化后的向量 def __init__(self, in_dim2048): super().__init__() self.fc nn.Sequential( nn.Linear(in_dim, 512), nn.ReLU(inplaceTrue), nn.Dropout(0.3), nn.Linear(512, 2) # source / target 二分类 ) def forward(self, feat_vec, lamda0.1): reversed_feat GradientReversal.apply(feat_vec, lamda) return self.fc(reversed_feat) class MTLModel(nn.Module): 完整模型编码器 专家路由 分类头 分割头 域判别器 def __init__(self, num_classes_cls2, num_classes_seg1, num_experts4): super().__init__() self.encoder ResNet50Encoder(pretrainedTrue) self.expert_net ExpertNet(in_dim2048, expert_dim512, num_expertsnum_experts) self.cls_head ClassificationHead(in_dim2048, num_classesnum_classes_cls) self.seg_head SegmentationHead(in_channels2048, num_classesnum_classes_seg) self.domain_head DomainDiscriminator(in_dim2048) def forward(self, x): feat_map, feat_vec self.encoder(x) fused_feat, gate_weight self.expert_net(feat_vec) cls_out self.cls_head(fused_feat) seg_out self.seg_head(feat_map) return cls_out, seg_out, feat_vec, gate_weight def domain_predict(self, feat_vec, lamda): return self.domain_head(feat_vec, lamda)注意域判别器的输入是编码器输出的全局特征向量没有和门控特征混在一起。原因很实际域对齐要作用在编码器层面让 source 和 target 在主干特征上分不开如果作用在门控融合之后任务头可能被带偏。梯度反转层的lamda是动态调度值会在下一章的训练循环里体现。4. 训练协议先有监督热身再无监督自适应损失权重按阶段调度模型结构只解决「能不能算」训练协议才决定「算出来有没有用」。这段用两个阶段的训练流程说明损失怎么组合、权重怎么调。4.1 四个监督信号分类、分割、域判别、伪标签一致性整个方案里有四个损失来源。分类和分割损失只作用在带标签的 source 域数据上走标准交叉熵和 Dice 损失。域判别损失在 source 和 target 上都计算目标是让判别器区分不了特征来自哪一边通过梯度反转层反向让编码器学不到域特有信息。伪标签一致性损失只作用在 target 数据上对同一张图做弱增强和强增强弱增强结果经过 argmax 得到伪标签强增强结果与伪标签求交叉熵强制模型在扰动下保持稳定。四个损失里前两个是「任务信号」后两个是「分布信号」。一个常见误区是把它们一次性全加在一起调大权重结果通常是任务没学精、域也没对齐。我一般分两阶段先只跑任务损失等分类和分割在验证集上稳定提升再加入域判别和一致性损失。4.2 两阶段训练循环的实现def train_mtl(model, source_loader, target_loader, epochs_warm15, epochs_adapt40): optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxepochs_warm epochs_adapt) for epoch in range(epochs_warm epochs_adapt): model.train() # 每个epoch内交替取source和target的batch iterator_s iter(source_loader) iterator_t iter(target_loader) total_loss 0.0 for step in range(len(source_loader)): try: s_img, s_cls_label, s_seg_label next(iterator_s) except StopIteration: iterator_s iter(source_loader) s_img, s_cls_label, s_seg_label next(iterator_s) try: t_img next(iterator_t) except StopIteration: iterator_t iter(target_loader) t_img next(iterator_t) # ---- 阶段一只跑任务损失 ---- cls_s, seg_s, feat_s, _ model(s_img) loss_cls F.cross_entropy(cls_s, s_cls_label) loss_seg F.binary_cross_entropy_with_logits(seg_s, s_seg_label) loss_total loss_cls 0.5 * loss_seg # ---- 阶段二加入域判别与伪标签一致性 ---- if epoch epochs_warm: # 域对抗损失lambda随训练进度从0.01线性爬到0.1 lamda 0.01 (0.1 - 0.01) * (epoch - epochs_warm) / epochs_adapt _, _, feat_t, _ model(t_img) dom_pred_s model.domain_predict(feat_s, lamda) dom_pred_t model.domain_predict(feat_t, lamda) dom_label_s torch.zeros(s_img.size(0), dtypetorch.long, devices_img.device) dom_label_t torch.ones(t_img.size(0), dtypetorch.long, devicet_img.device) loss_domain F.cross_entropy(dom_pred_s, dom_label_s) \ F.cross_entropy(dom_pred_t, dom_label_t) # 伪标签一致性损失阈值0.9确保只用高置信度样本 t_strong strong_augment(t_img) _, _, feat_strong, _ model(t_strong) logits_strong model.cls_head(torch.cat([feat_strong, model.expert_net(feat_strong)[0]], dim-1)[:, :2048]) # 简化写法实际实现中直接复用同一特征路径即可 with torch.no_grad(): logits_weak torch.softmax(model.cls_head(feat_t), dim-1) pseudo_label logits_weak.argmax(dim-1) confidence logits_weak.max(dim-1).values mask (confidence 0.9).float() # 只有通过阈值掩码的样本参与计算 loss_pseudo (F.cross_entropy(logits_strong, pseudo_label, reductionnone) * mask).mean() loss_total loss_total 0.3 * loss_domain 0.2 * loss_pseudo optimizer.zero_grad() loss_total.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() total_loss loss_total.item() scheduler.step()代码中几个参数有实际讲究。0.5 * loss_seg是因为医疗图像前景占比通常不足 5%分割损失天然偏小直接相加会让分类主导更新方向Dice 类损失的作用在 0.30.7 之间调。clip_grad_norm设 5.0 防止分割头的反卷积层梯度爆炸这在医疗大图上格外重要因为显存受限时 batch size 小BN 统计量不稳定。epochs_warm设 15 是让任务头先学到基本概念如果一上来就加域对抗编码器权重的变化方向会同时被任务和域两股力量拉扯结果两边都没学透。4.3 一组能少走弯路的初始超参数下表是我在这类方案上的起步配置。它不算最优但保证最先跑通后续再逐项调。参数推荐值调整方向输入尺寸256x256显存不够时降到 192而不是减小 batchbatch size8医疗图像分辨率高batch 过小需要配合梯度累积预训练骨干ImageNet V2不要用随机初始化小数据会直接崩专家数量4数据量大可以试 8小于 500 例保持 4门控温度0.8死专家时调到 1.2后期想收敛快调到 0.5域对抗 λ0.01 → 0.1 线性增长过快的表现是任务头指标骤降伪标签置信度阈值0.9target 分布差异大时降到 0.8 保住样本量优化器AdamW lr1e-4换 SGD 需 lr1e-3 并加 momentum 0.95. 医疗图像多任务自适应的五个坑现象、原因、解决办法这章是我最想写透的部分因为这五个坑每一个都让团队白白消耗过两周以上时间。按现象、原因、解决的顺序逐一拆。5.1 死专家门控把所有样本都路由给同一个专家Multi-task 退化成单网络现象是训练一段时间后可视化门控权重发现一个专家的权重接近 1其余接近 0负载均衡完全失效。原因是门控的初始化偏差和任务梯度方向一致导致某个专家在早期多抢了一点梯度随后强者愈强形成赢家通吃。这在多任务里是灾难性的因为网络退化成了一个普通大网络ExpertNet 形同虚设。解决方法是三管齐下。门控输入加噪声这是我在 3.2 节代码里写过的torch.randn_like在门控损失里加负载均衡惩罚项统计每个专家平均被选中的概率和均匀分布求 KL 散度最后把门控温度调高到 1.2让 Softmax 软化。这里最容易翻车的是只加噪声、不加惩罚噪声只会拖慢收敛不会改变最终结果。5.2 ResNet50 预训练在医疗图像上带来的「预训练诅咒」现象是模型在 source 域验证集上表现很好一到 target 域就大幅掉点比不用预训练的模型还差。原因是 ImageNet 预训练参数虽然提供了纹理和边缘表达但也残留着自然图像分布的偏差在域差异大的医疗图像上这些偏差会被无监督自适应放大。经过多次排查问题出在浅层 BN 统计量上。ResNet50 前几层的 BN 统计量是在 ImageNet 分布上统计的医疗图像的像素分布完全不同第一层输入分布就偏了后面的层全部跟着偏。解决办法是训练初期冻结前两个 stage 的参数只更新 stage3 和 stage4等模型适应了医疗图像分布再解冻另外可以把所有 BN 换成 InstanceNorm医疗图像 batch 小BN 本身就不稳。5.3 伪标签噪声被无监督损失放大现象是加入伪标签一致性损失后target 域指标不升反降且损失曲线完全不下降。原因是置信度阈值设得太低低质量伪标签进入训练模型在错误的方向上越走越远一致性损失实际在帮助模型记住错误预测。我现在的做法是双阈值策略。0.8 到 0.9 区间的样本只参与特征对齐不参与分类损失的更新只有高于 0.9 的样本才给分类任务提供监督信号。另一个有效手段是 EMA 模型生成伪标签用历史权重的平均预测结果比当前权重更稳定。注意 EMA 的 momentum 不能太小0.999 起步伪标签噪声大时保持 0.999 不放松。5.4 BatchNorm 统计量是域自适应里最容易被忽略的「黑匣子」现象是一个模型在单中心实验里指标正常换成多中心数据后训练损失波动剧烈甚至直接 NaN。原因典型的BN 层的统计量在 source 和 target 两个域之间互相拉扯尤其是 batch size 小的时候统计量的抖动幅度大训练噪声陡增严重时数值溢出。解决方法是把 BN 的track_running_stats关闭让每一层的统计量完全由当前 batch 决定这在域自适应里其实是个优势因为模型不受历史统计量偏移影响。代价是所有 BN 层需要额外训练几个 epoch 才能稳定所以我会在热身阶段保持 BN 统计量不变只更新其他层等进入自适应阶段再解冻。如果显存充足也可以换成组归一化彻底跟 batch size 解耦。5.5 分类和分割两条损失打起来了现象是损失都正常下降但分割的 Dice 迟迟不涨或者分类的 AUC 到了 0.9 之后停滞不动另一个任务只要一提升这个任务就掉一点。原因在于两个任务的梯度方向不一致共享编码器同时要满足边缘敏感的分割和全局语义的分类更新方向互相抵消。常规做法是给两个任务损失配不同权重见效慢但稳定。性价比更高的办法是分步训练前 15 轮只用分类损失然后加入分割损失而不是一开始就双损失并跑。还有一类做法是 GradNorm计算两个任务梯度的范数并动态调整权重这不是玄学但实现复杂建议先试简单方案。最后保留一个后悔药式的兜底训练完成后用测试集做一次特征可视化确认分割和分类特征并没有被过度耦合否则宁可拆成两个模型分别上线。6. 验证这套方案值不值衰减曲线、特征可视化和我的收尾习惯方案做完了最后一步是让验证协议也对得起前面的投入而不是拿一两个测试集上的数字就下结论。我习惯用的验证协议是「衰减曲线」对比。以同一组测试集为横轴分别记录 source 域和 target 域上的指标变化把只做多任务训练的结果和无监督自适应之后的结果画在同一条图上。如果 target 域指标从掉 15 个点变成只掉 5 个点方案的价值就直观可见。这条曲线同时是判断超参是否合理的依据如果域自适应加入后 source 域指标也在掉说明 λ 开太大编码器的语义能力被对抗损失干扰了回调 λ 重跑。第二个手段是 T-SNE 特征可视化。取 source 和 target 各 200 个样本用编码器的输出向量做降维投影。理想状态是两群特征完全混在一起这比任何 loss 值都能说明域对齐是否成功。如果投影图上出现清晰的两簇说明编码器仍然保留了明显的域特征此时冻结 BN 统计量再跑 5 个 epoch通常能看到改善。这招排查成本低视觉效果直观是团队里最常用的验收方式。第三个我最近开始强依赖的做法把门控权重拉出来看。统计每个样本的专家选择分布按分类和分割任务分别统计会看到不同任务在专家上的偏好——分类样本倾向于路由到某个专家分割样本倾向于另一个。如果两个任务在所有样本上的专家选择完全一致说明专家网络没有真正分工所谓的多任务只是同一个特征换了个输出头而已。这时候增加专家数量或提高门控温度能让分工重新分化。我的收尾习惯是先跑一轮只含任务损失的 baseline再跑完整方案两份结果一起存档因为只报自适应后的好看数字没人知道自适应到底做了多少贡献。这个习惯帮我挡掉过好多次返工。如果这套流程能帮你在医疗图像多任务上少走一圈弯路希望帮到你。权当把这几年踩坑攒下来的经验交付给下一个接手的人。本文还有配套的精品资源点击获取
返回列表