MedicalNet代码改造:从3D分割到3D分类的迁移学习实战
去年有个项目让我印象很深手里的3D分割代码跑得很顺基于MedicalNet的3D ResNet预训练权重做体素分割结果新任务突然变成了3D分类——判断某个区域的影像属不属于高风险病灶。我第一反应是去找一个现成的3D分类预训练模型翻了一圈发现MedicalNet官方提供的还是那批在医学影像大数据上预训练好的3D ResNet权重。这反而提醒了我与其另起炉灶不如把已经跑通的分割代码改造成分类代码既能复用数据预处理、训练管线也能把预训练权重里学到的三维解剖特征直接迁移过来。这篇文章就是这次改造过程的完整记录。我会从MedicalNet的分割实现讲起逐个拆解从分割到分类需要动的代码位置包括网络结构替换、预训练权重加载、数据管线调整、训练策略设置以及我在实践中踩过的几个坑。如果你手里也有一套分割代码或者正打算用MedicalNet做3D医学影像分类这篇文章可以直接当改造手册用。1. 动手改造前先看懂MedicalNet这套迁移学习的真实边界1.1 分割与分类共享的东西比想象中多很多人在分割和分类之间划了一道很硬的界线分割是逐体素的分类是图像级的听起来像是两个完全不同的任务。但从特征提取的角度看它们的底层逻辑是一样的。分割需要理解每个体素属于哪个解剖结构或病灶区域分类需要判断整个图像或者图像中的某个区域整体属于哪一类。前者要求模型保住空间细节后者要求模型把空间信息聚合成一个全局判断。这个差异主要体现在网络的“后半段”前半段对纹理、边界、空间结构的特征提取能力是完全可以共用的。MedicalNet提供的预训练权重本质上就是前半段——一个在大规模医学影像数据上学过的3D ResNet骨架。官方把这套模型用在分割任务上预训练时用的是多模态医学影像数据通过大量分割标注让模型学会了识别三维医学图像里的基本结构。这些结构特征迁移到分类任务时几乎无损因为不管是CT还是MRI不管是肿瘤、器官还是血管分类任务同样需要模型先“看见”边界和纹理才能做出正确判断。这也是为什么从分割改到分类不需要推倒重来。把分割解码器丢掉换上分类头骨架部分直接加载预训练权重整套代码的改动量其实远比你想象中少。1.2 MedicalNet没有提供的东西需要自己补上这里要特别说清楚一个容易误解的地方MedicalNet官方开源的内容是预训练权重和配套的3D ResNet基础网络它没有给你一套“开箱即用的分类器”。我在项目里下载下来的resnet_50.pth加载进去之后发现模型结构里自带一个conv_seg分割头。也就是说官方给的是“适合做分割迁移的模型”而不是“已经帮你接好分类输出的模型”。需要自己补的东西主要有这几件分类头把ResNet输出的特征图聚合成一个固定维度的向量再接全连接层输出类别数。池化层分割任务一般不需要全局池化但分类任务通常需要一个AdaptiveAvgPool3d把不同尺寸的特征图压成1x1x1。损失函数DiceLoss那一套要换成适合分类的交叉熵或者带权重的交叉熵。预处理约定预训练权重的输入尺寸、归一化方式、裁剪策略都需要和预训练条件保持一致否则迁移效果会打折扣。这些内容都不是MedicalNet自带的而是“从分割到分类”改造时需要自己动手的部分。我在实际操作中的经验是先想清楚哪些可以复用哪些必须重写再动代码这样能少走很多弯路。1.3 什么情况下不该用这套方案虽然MedicalNet的迁移学习在大多数3D医学影像任务里都有效但它不是万能的。我总结了几种不太适合硬套的场景任务本身是2D切片分类一张CT或者MRI按切片输入。这种情况用2D预训练模型更合适强行上3D模型只会增加计算量而且3D预训练的特征在单张切片上不一定比2D ImageNet特征更好。数据量非常庞大比如你已经有几百万个带有标注的3D样本。这种情况从头训练可能比迁移学习更干净没必要承担预训练分布带来的偏差。输入对象非常小比如几个体素就能判断的局部任务。预训练骨架的卷基层感受野远大于输入尺寸特征提取效率会很低换一个小模型反而更合理。除此之外MedicalNet这套方案通常都能稳定带来收益尤其是医学影像数据量只有几百到几千例的典型场景。2. 从分割代码到分类代码核心改动集中在这四处2.1 网络结构替换把解码器换成分类头MedicalNet官方仓库里的ResNet3DMedNet在基础ResNet后面接了一个conv_seg作为分割输出层。如果沿用这个类来搭建分类器最简单的做法是实例化模型之后把分割头丢掉换成一个由全局池化和全连接层组成的分类头。实际改造代码大致长这样import torch import torch.nn as nn from medicalnet import resnet50 # 假设已经从MedicalNet仓库导入 def build_classifier(checkpoint_path, num_classes2): # 实例化3D ResNet这里的num_seg_classes传1只是为了满足初始化 model resnet50( sample_input_D64, sample_input_H64, sample_input_W64, num_seg_classes1 ) # 去掉原来的分割头 model.conv_seg nn.Identity() # 替换成分类头 model.avgpool nn.AdaptiveAvgPool3d((1, 1, 1)) model.fc nn.Sequential( nn.Dropout(0.3), nn.Linear(2048, 512), nn.ReLU(inplaceTrue), nn.Dropout(0.3), nn.Linear(512, num_classes) ) # 加载预训练权重时跳过分类头 state_dict torch.load(checkpoint_path, map_locationcpu) new_state_dict {k.replace(module., ): v for k, v in state_dict.items()} missing_keys, unexpected_keys model.load_state_dict(new_state_dict, strictFalse) if missing_keys: print(未加载的层, missing_keys) if unexpected_keys: print(多余的层, unexpected_keys) return model这里有几个关键决策为什么用AdaptiveAvgPool3d((1, 1, 1))而不是直接x.view(x.size(0), -1)因为自适应池化可以接受任意输入尺寸模型在推理阶段就不用被训练时的尺寸绑死。而且平均池化能把全局空间信息聚合起来对分类任务来说比直接展平更稳。为什么分类头里加了Dropout3D医学影像数据量通常不大ResNet50的参数量又很大不加Dropout很容易过拟合。两个Dropout层分别放在两个全连接之间实际测试下来对缓解过拟合很有帮助。为什么加载权重时用strictFalse因为预训练权重里有分割头的参数而当前模型已经没有这些层了strictTrue会直接报错。但注意strictFalse也意味着你写错层名时不会报错所以加载完一定要打印missing_keys和unexpected_keys人工核对一下。2.2 前向传播的输出语义变化分割模型的前向输出是[B, C, D, H, W]其中C是分割类别数每个体素位置对应一个概率分布。而分类模型的前向输出是[B, num_classes]每个样本对应一个类别得分。改动之后原来在分割代码里对输出做argmax(dim1)得到逐体素标签的地方全部要改成argmax(dim1)得到逐样本标签。这一处改动看起来简单实际很容易忽略的是后续的验证可视化逻辑。分割项目里通常会有把输出mask叠加到原图上的脚本这些在分类项目里不再适用需要换成绘制ROC曲线、混淆矩阵这类分类评估工具。### 2.3 损失函数和评估指标需要一次性换掉 分割任务最常用的损失是DiceLoss和交叉熵的组合评估指标是Dice和IoU。分类任务则要换成更适合分类的损失和指标。我在改造时直接列了一个对照表 | 环节 | 分割任务常用 | 分类任务常用 | |----------|--------------------------------------|----------------------------------------| | 损失函数 | DiceLoss、FocalLoss、CEDice组合 | CrossEntropyLoss、带权重CE | | 输出层 | Softmax或者Sigmoid逐体素输出 | Softmax多分类、Sigmoid二分类 | | 评估指标 | Dice、IoU、Hausdorff距离 | Accuracy、Precision、Recall、F1、AUC | | 可视化 | 分割mask叠加原图 | ROC曲线、混淆矩阵、Grad-CAM | 损失函数的具体改动可以这样理解分割任务里你会写loss dice_loss(output, label)分类任务里则是loss criterion(output, label)label从形状[B, D, H, W]变成[B]。交叉熵的target不能再是one-hot的张量而是要直接传一个包含类别索引的一维张量。 对于二分类场景我建议用带权重的交叉熵或者直接改用BCEWithLogitsLoss它内部会把logits通过sigmoid再计算交叉熵数值上更稳定而且天然支持类别权重。 ### 2.4 训练循环里的维度对齐问题 分割代码的dataloader返回的通常是(image, label_mask)image的形状是[B, 1, D, H, W]label_mask的形状是[B, D, H, W]。改成分类之后label_mask要变成[B]每个元素是一个整数类别索引。 训练循环本身改动不大但有几个细节值得注意 - 如果使用混合精度训练scaler.scale(loss).backward()的写法不变但分类头最后输出的logits规模可能偏大导致fp16下精度下降。可以调整分类头最后一层的初始化或者用torch.nn.init.xavier_normal_对fc层做初始化。 - 验证阶段分割模型通常要在一个病例的多个patch上分别预测再做拼接分类模型则要看你是按patch分类还是按subject分类。如果按subject分类多个patch的输出要做平均或投票这部分逻辑需要额外写。 - 如果之前分割代码里有torch.argmax(outputs, dim1)的操作注意现在dim1仍然适用但outputs已经是三维结果是一维tensor不再需要做任何解码。 ## 3. 预训练权重的加载细节决定了迁移学习的地基稳不稳 ### 3.1 去掉权重文件里的module前缀 MedicalNet官方给出的权重文件是用DataParallel训练的所以所有键名都带有module.前缀。直接加载到普通模型里会报找不到键的错误需要先过滤掉这个前缀。 我通常用下面这段代码做通用加载 python import torch def load_medicalnet_weight(model, checkpoint_path): state_dict torch.load(checkpoint_path, map_locationcpu) if module. in list(state_dict.keys())[0]: state_dict {k.replace(module., ): v for k, v in state_dict.items()} missing_keys, unexpected_keys model.load_state_dict(state_dict, strictFalse) return missing_keys, unexpected_keys加载完成后一定要看打印信息。预期结果是missing_keys里包含分类头相关层unexpected_keys里包含原分割头相关层。如果missing_keys里有某个resnet层的名字说明你的模型层名和预训练权重对不上这多半是resnet50的构造参数和MedicalNet不一致导致的比如zero_init_residual或者groups这些参数不同。3.2 加载权重后先做一次冒烟测试很多人在加载权重之后直接开训结果发现loss不降然后开始怀疑梯度、优化器、学习率排查半天才发现是权重根本没有正确加载。我建议在开始训练前先手动构造一个随机输入跑一次前向和反向确认形状和梯度流都正常。model build_classifier(resnet_50.pth, num_classes2) x torch.randn(2, 1, 64, 64, 32) y model(x) print(y.shape) # 期望输出 torch.Size([2, 2]) loss torch.nn.functional.cross_entropy(y, torch.tensor([0, 1])) loss.backward() print(冒烟测试通过)这一步能快速判断网络结构是否改对尤其是分类头输出的类别数是不是和目标一致。如果y.shape不对后面的训练脚本再怎么写都是错的。3.3 输入尺寸、归一化参数要和预训练条件对齐MedicalNet的预训练是在相对较小的patch上做的常见的是64x64x32或64x64x64。这个尺寸的由来不是拍脑袋定的而是医学影像数据在裁掉无关背景之后主要解剖结构刚好能落在这个空间范围内。如果任务的数据分辨率更高比如原始CT是512x512x300不建议直接把整张图塞进模型。显存不够只是原因之一更关键的是预训练骨架在64x64x32尺度上学习的特征模式放到512x512x300上不一定适用。一个合理的做法是先用中心裁剪或重采样把数据统一到预训练尺寸附近。归一化方面MedicalNet的实践通常采用z-score归一化每个样本自身减均值再除以标准差而不是用固定的全局均值和标准差。因为不同扫描设备、不同医院的影像灰度范围差异很大全局统计量很容易被个别样本带偏。我在代码里用的就是每次读取一个样本后对非背景区域计算均值和标准差再归一化。4. 数据准备与训练策略三份可以直接参考的配置4.1 NIfTI读取与预处理流程医学影像最常见的格式是NIfTI后缀为.nii或.nii.gz。用nibabel读取之后需要做重采样、裁剪、归一化这三步。下面是我在分类任务里常用的预处理流程import nibabel as nib import numpy as np def load_and_preprocess(nii_path, target_size(64, 64, 32)): img nib.load(nii_path).get_fdata() # 1. 裁剪非零区域减少背景干扰 nonzero np.nonzero(img) if len(nonzero[0]) 0: z_min, z_max nonzero[0].min(), nonzero[0].max() 1 y_min, y_max nonzero[1].min(), nonzero[1].max() 1 x_min, x_max nonzero[2].min(), nonzero[2].max() 1 img img[z_min:z_max, y_min:y_max, x_min:x_max] # 2. 缩放到目标尺寸 from scipy.ndimage import zoom zoom_factors [target_size[i] / img.shape[i] for i in range(3)] img zoom(img, zoom_factors, order3) # 3. z-score归一化 img (img - np.mean(img)) / (np.std(img) 1e-8) # 增加通道维度 img np.expand_dims(img, axis0) return img.astype(np.float32)这个写法有个地方要说明裁剪非零区域这一步很关键。原始CT数据里除了人体区域外全是空气灰度值接近0。如果不做裁剪直接用整个扫描体积做归一化那些空气区域会把均值往0拉导致器官和病灶的灰度对比被削弱。裁剪之后只对有效区域做归一化特征分布会稳定很多。4.2 冻结策略和分层学习率迁移学习的经典做法是先冻结骨干只训练分类头等loss降到一定程度再解冻整个网络做微调。但MedicalNet和ImageNet预训练模型有一点不同3D医学影像数据量通常很小整个ResNet50一旦解冻很容易过拟合。我更推荐的做法是分层学习率而不是完全冻结。具体来说optimizer torch.optim.AdamW([ {params: [p for n, p in model.named_parameters() if fc not in n], lr: 1e-4}, {params: [p for n, p in model.named_parameters() if fc in n], lr: 1e-3} ], weight_decay1e-4)骨干部分用较小的1e-4分类头用较大的1e-3。这样既能让预训练特征在后向传播中缓慢调整又能让随机初始化的分类头更快收敛。这里有个容易被忽视的坑如果你真的想完全冻结骨干不能只设置param.requires_grad False因为PyTorch中BatchNorm层的running_mean和running_var在模型处于train()模式时仍然会更新。即使把整个骨干的requires_grad都关了BN的统计量还是会被训练数据带偏。解决方法是手动把包含BN的层设置成eval()模式或者干脆不要完全冻结用分层学习率的方式让骨干参与微调反而更省心。4.3 显存控制、数据增强和epoch安排3D模型对显存的消耗是2D模型的几十倍。64x64x32的输入batch size为8ResNet50在12G显存上已经站得很勉强。如果显存不够通常有几个选择减小batch size到2或4同时把梯度累积步数增大模拟大batch的效果。把输入尺寸降为48x48x32或64x64x24但要注意预训练骨架可能对长宽比有一定依赖不要降太多。使用混合精度训练torch.cuda.amp在A100或者3090这类卡上能节省大量显存。数据增强方面3D医学影像分类常用的增强包括随机翻转、小角度旋转、随机缩放和随机裁剪。注意不要用太强的几何增强因为医学影像的解剖结构有固定的空间方向左右翻转可以前后翻转就要看任务是否允许。我在项目中用到的增强组合是随机水平翻转概率0.3随机小角度旋转范围±10度随机缩放范围0.9~1.1随机裁剪回原尺寸epoch的设置在3D医学影像任务里通常不需要太多。我一般先跑30个epoch看趋势如果验证AUC还在上升就继续。大量实验下来MedicalNet预训练模型通常在前10个epoch内就能看到明显的收敛迹象如果20个epoch后验证指标还在原地踏步大概率是数据管线或者标签出了问题而不是训练不够充分。5. 从分割改到分类后我踩过的五种坑5.1 加载预训练权重后loss不降问题出在BatchNorm这是我第一次改造时遇到的最诡异的坑。加载了MedicalNet预训练权重分类头正常初始化跑第一个epoch时loss确实在降但第二个epoch开始就原地不动了。排查了很久最后发现是前面提到的BN统计量问题。因为我没有完全冻结骨干只是把骨干的学习率设得很小但BN层依然像从头训练那样更新running_mean和running_var很快就破坏掉了预训练权重里的特征分布。解决办法有两种。第一种是前面说的在训练代码里把骨干部分切换到eval()模式确保BN完全不更新第二种是干脆解冻骨干让它以更小的学习率微调让BN跟着一起适应新数据。第二种在数据不是特别少的情况下更稳定我最终选了第二种。5.2 类别极度不平衡时冻结骨干会放大偏置还有一个项目是二分类任务正负样本比例接近1:10。我一开始用默认的CrossEntropyLoss前20个epoch验证集loss一直在降但AUC只有0.6左右。后来检查预测结果发现模型几乎把所有样本都预测成了多数类。原因不难理解随机初始化的分类头一开始输出概率接近均匀分布在类别严重不平衡的情况下交叉熵的梯度会不断把预测推向多数类。如果此时骨干被冻结分类头只能靠自己调整几乎没有能力从数据中学到少数类的特征。解决办法是给交叉熵加类别权重。PyTorch里可以直接传一个weight参数class_weights torch.tensor([1.0, 10.0]).cuda() criterion nn.CrossEntropyLoss(weightclass_weights)实际测试中加上类别权重之后AUC很快从0.6涨到了0.82。如果类别权重效果不明显还可以尝试FocalLoss它能降低易分类样本的梯度贡献让模型更关注少数类。5.3 验证集AUC很高实际使用时误判依然多训练和验证指标都很漂亮AUC达到0.95但放到真实临床数据上测试时误判率明显偏高。后来发现原因是训练和推理时的数据粒度不一致。训练时我把每个case裁成了多个patch每个patch单独算一个标签。但真实使用场景里一个case只有一个标签多个patch的预测结果需要合并。我之前在推理阶段直接对每个patch取argmax然后统计次数但这种方式没有考虑patch的重要程度。一个case的病灶区域可能只占很小一部分大量正常patch的预测结果会把病灶patch的投票稀释掉。正确的做法是对patch级的预测概率做平均再做最终决策。比如每个patch输出一个[0, 1]之间的概率把所有概率取平均超过阈值就判定为阳性。这样既保留了patch的局部信息又不会让多数正常patch淹没少数病灶patch。5.4 显存不够降低分辨率还是滑窗裁剪我遇到过显存刚好不够的情况64x64x32的输入batch size设成8直接OOM。当时心里想着换个更大的batch来提升稳定性但其实3D模型对batch size的敏感程度没有2D那么高。两个可选方案我都试过降低输入分辨率到48x48x24显存占用下降明显但精度也掉了不少尤其是边界模糊的病灶。保持64x64x32把batch size降到2开启混合精度显存占用能降一半左右精度几乎没有影响。最终选择了后者。这里也推荐一个很实用的技巧如果显存还是不够可以把最后一个stage通常是layer4的输出先做一次全局池化再算loss这样梯度回传时不需要保留layer4所有中间特征图的梯度节省的显存非常可观。具体做法是把模型输出改成两个分支一个用于训练一个用于推理。5.5 分类头初始化不当导致训练初期异常波动分类头的全连接层如果用默认初始化在某些情况下会让最后一层logits的绝对值变得非常大。Softmax在输入值很大时梯度会接近饱和训练初期loss下降极为缓慢甚至出现NaN。我后来在分类头里对最后一个全连接层单独做了初始化nn.init.xavier_normal_(model.fc[-1].weight, gain1.0) nn.init.constant_(model.fc[-1].bias, 0.0)这样能保证最后一层输出在初始阶段处于一个合理的数值范围前几个epoch的loss曲线会更平滑。如果是多分类任务还可以把最后一层的bias初始化为log(1 / num_classes)让初始输出接近均匀分布配合类别不均衡问题会更稳。6. 改造是否成功的评判标准用两个小实验自查6.1 先过拟合一个小batch改造完成之后先别急着跑全量数据。我每次改网络结构都会做这个实验取8个样本反复训练同一个小batch观察loss能不能降到接近0。这个实验的价值在于快速验证整个链路是否畅通。如果过拟合小batch都做不到说明模型结构或者数据管线里一定有bug此时跑全量数据只会浪费时间和显卡。我在一次改造中就遇到过这个情况模型输出的类别数和标签不匹配loss恒等于固定值小batch实验一下子就暴露了问题。6.2 预训练权重和随机权重各跑20个epoch如果小batch实验通过下一步可以做对比实验一组加载MedicalNet预训练权重另一组用同样的网络结构但随机初始化其余训练配置保持一致各跑20个epoch。预期结果是预训练组的验证AUC明显高于随机初始化组且预训练组的loss下降速度更快。如果两组曲线几乎重合说明预训练权重没有真正参与训练或者你的输入数据分布和预训练数据差异太大导致迁移学习没有任何优势。这个实验能帮你快速判断“迁移学习是否有效”。# 伪代码示意 train_with_pretrained train(model_pretrained, train_loader, val_loader) train_with_random train(model_random_init, train_loader, val_loader) print(f预训练组 AUC: {train_with_pretrained[best_auc]:.4f}) print(f随机初始化组 AUC: {train_with_random[best_auc]:.4f})如果预训练组的AUC优势不明显我的经验是先去检查预处理pipeline尤其是归一化和裁剪方式。很多时候不是预训练权重没用而是你的输入影像和预训练数据差异太大导致预训练特征根本派不上用场。6.3 用激活图检查模型关注区域训练完成后可以用简单的Grad-CAM或注意力可视化的方式看一下模型在分类时到底关注图像的哪些区域。对医学影像分类来说模型应该关注病灶区域而不是背景或正常组织。这个步骤不需要特别复杂的工具最简单的方式是取出最后一层特征图做空间平均后可视化。如果模型在正常组织上也有很高的激活值说明它可能学到了不该学的伪相关特征比如扫描设备和病灶位置之间的相关性。这时候需要增加数据清洗或者用更强的数据增强来打破这种伪相关。如果让我再走一遍从分割到分类的改造过程我会先搭一条最小闭环一个小模型、一个batch、一个分类头跑通之后再考虑加预训练权重和复杂结构。很多从分割改到分类的任务问题往往不是模型本身而是数据pipeline没有跟着任务类型一起改过来。希望这些记录能帮你少踩几个坑。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →