简介:面向医学影像分析与深度学习入门者的肺炎诊断工具包,基于Transformer架构并结合ResNet34预训练权重,完成胸部X光图像的肺炎分类任务。模型经过400轮训练,批量大小32,学习率0.0001,并内置混淆矩阵评估模块,便于直观查看各类别诊断准确率与误判情况。整套系统基于PyTorch实现,代码结构清晰,适合作为医疗AI课题的参考基线或教学案例。
压缩包共13个文件,核心为9个Python脚本,涵盖模型定义、数据加载、训练、预测及混淆矩阵绘制等完整流程;另含json类别索引、md说明、txt配置说明及docx附赠文档,可辅助环境搭建与参数理解。资源包仅55KB,轻量易用。目前已有55人学习下载,适合希望快速上手Transformer医学影像分类,或需要一套可复现评估流程的研究者与开发者。
1. 用ResNet34+Transformer做胸片肺炎诊断:400轮训练背后真正值得关注的设计点
胸部X光肺炎诊断这两年已经从“有没有肺炎”升级到“要区分细菌性还是病毒性”,再进一步就是给病灶定位。纯CNN模型比如ResNet34,对局部磨玻璃影和间质性纹理很敏感,但胸片上的病灶常常是分散的、跨肺野的,CNN的局部感受野容易漏掉全局上下文;纯Transformer又需要海量数据和很长的训练时间,在几千张胸片这种小规模数据集上直接翻车是常事。所以这个标题的工程思路很聪明:用ResNet34预训练权重先把图像咬成稳定的局部特征,再接Transformer做全局交互,让模型同时拿到细节和位置关系。400轮、batch_size=32、lr=0.0001这套配置看似常规,但每个数字都要跟模型架构、数据量、预训练权重状态配合,否则要么200轮过拟合,要么跑完400轮验证集还在震荡。
这篇文章按“模型结构 → 数据与训练 → 混淆矩阵评估 → 排坑 → 验证”的顺序,把整个方案的复现路径讲透。适合手里有胸片数据集、想用PyTorch做可解释分类系统的算法工程师或医工交叉方向的研究生。下面所有代码都按PyTorch 1.13及以上版本书写,如果你还在用旧版本,注意TransformerEncoderLayer的激活函数参数名要改成activation="gelu"这种写法。
2. 模型主干设计:为什么是ResNet34预训练权重而不是纯Transformer
把“基于Transformer架构”和“ResNet34预训练权重”放在同一个标题里,很多人第一反应是“这俩怎么拼”。其实这在医学影像分类里是过去两年最稳的混合架构:CNN负责提取局部纹理,Transformer负责建模长距离依赖。下面对模型每一层做拆解,包括张量流、权重复用和冻结策略。
2.1 混合架构的张量尺寸变化:从512通道特征图到64个token序列
胸部X光原图分辨率通常很大,但医学影像数据集样本量有限,直接整图喂Transformer完全不现实。常见做法是先让ResNet34把图像压缩成语义特征图,再把每个空间位置当作token送入Transformer。
假设输入是256×256的灰度胸片(复制成3通道),ResNet34主干一直保留到conv5_x输出,得到[batch, 512, 8, 8]的特征图。把这8×8=64个位置展开,每个位置有512维特征,相当于一句话有64个“词”,每个词的嵌入维度是512。接着加上可学习的位置编码,因为胸片的肺野区域、心影位置、肋膈角在解剖学上固定,位置信息跟病灶定位一样重要。最后通过TransformerEncoder后取CLS token分类。完整定义如下:
import torch import torch.nn as nn import torchvision.models as models class ResNetWithTransformer(nn.Module): def __init__(self, num_classes=3, transformer_layers=4, nhead=8): super().__init__() # 加载ResNet34预训练权重,去掉全局池化和全连接层 resnet = models.resnet34(weights=models.ResNet34_Weights.IMAGENET1K_V1) self.features = nn.Sequential(*list(resnet.children())[:-2]) # 输出 [B, 512, 8, 8] # Transformer编码器:d_model必须等于特征图通道数512 encoder_layer = nn.TransformerEncoderLayer( d_model=512, nhead=nhead, dim_feedforward=2048, dropout=0.1, activation='gelu' ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=transformer_layers) # 可学习位置编码:1行对应batch共享,64个空间位置,512维 self.pos_embed = nn.Parameter(torch.randn(1, 64, 512)) # CLS令牌,用于聚合全局信息 self.cls_token = nn.Parameter(torch.zeros(1, 1, 512)) self.classifier = nn.Linear(512, num_classes) def forward(self, x): x = self.features(x) # [B, 512, 8, 8] B, C, H, W = x.shape x = x.flatten(2).permute(0, 2, 1) # [B, 64, 512] x = x + self.pos_embed # 位置编码加到每个token上 cls = self.cls_token.expand(B, -1, -1) # [B, 1, 512] x = torch.cat([cls, x], dim=1) # [B, 65, 512] x = self.transformer(x) cls_out = x[:, 0] # 取CLS token return self.classifier(cls_out)逻辑说明:self.features取自resnet34的卷积段(conv1到layer4),list(resnet.children())[:-2]去掉了最后的avgpool和fc。输入经过ResNet34下采样32倍,256×256得到8×8特征图。每个空间位置的512维向量就是一个token,这就把图像从像素空间转换成了语义空间。位置编码用randn初始化,让模型在训练中自己去学相对位置,这样比固定的正弦编码更适应胸片的解剖位置。CLS token在训练后可以理解为模型的全局决策向量,它的注意力权重对应着模型对每个区域的重视角。
参数说明:transformer_layers=4是经验值,64个token很短,4层编码器足够捕获跨肺野依赖,堆到8层会在小数据集上过拟合。nhead=8要求d_model能被8整除,512/8=64,每个头的维度合适。dim_feedforward=2048是前馈网络中间维度,如果显存紧张可以降到1024,收敛速度会慢一些,精度通常影响不大。
2.2 预训练权重的加载与通道适配:灰度图如何复用ImageNet权重
胸片是灰度图,而ResNet34官方预训练权重是在ImageNet上训练的,第一层卷积期望3通道输入。两种常见适配方式:一是把灰度图复制成三通道,继续用官方权重;二是把第一层卷积改成1通道,将预训练卷积核在通道维取平均。工程上推荐第一种,因为ImageNet预训练模型对边缘、纹理的低层响应本来就是跨颜色空间泛化的,复制通道不会破坏特征。代码如下:
import torchvision.models as models def build_model(num_classes=3, pretrained=True): resnet = models.resnet34(weights=models.ResNet34_Weights.IMAGENET1K_V1 if pretrained else None) # 截取到layer4,不要avgpool和fc backbone = nn.Sequential(*list(resnet.children())[:-2]) model = ResNetWithTransformer(num_classes=num_classes) # 将官方backbone的权重复制到我们模型的特征提取器 model.features.load_state_dict(backbone.state_dict()) return model逻辑说明:load_state_dict要求两边的层名和shape完全一致。这里直接把backbone.state_dict()装进model.features,因为二者来自同一个官方结构的前半部分。真正在训练循环里,输入图像会经过一个repeat(3,1,1)操作把单通道复制成三通道,这一步放在Dataset里,我们会在下一章数据部分看到。
参数说明:pretrained=True会从PyTorch官方hub下载权重到本地缓存。如果你的训练机没有外网,需要提前在有网的机器上下载好resnet34-333f7ec4.pth,拷贝到缓存目录,否则程序会在第一步卡住报“Connection error”。使用预训练权重是本项目能400轮收敛的核心原因,从头训练ResNet34+Transformer在这个规模的数据上至少要2000轮,而且稳定性差很多。
2.3 冻结与解冻策略:浅层冻结、深层与Transformer全量微调
医学图像与自然图像差异很大,但ResNet34的底层卷积仍然能提取通用边缘、纹理。训练中如果把整个backbone全部解冻,胸片小数据集很容易让浅层卷积被噪声带偏;如果完全冻结,深层特征又无法适配胸片的特异性纹理。一个可复制的策略是:前5轮冻结整个ResNet,只训练Transformer和分类头,让Transformer先适应ImageNet特征分布;从第6轮开始解冻layer3和layer4,浅层保持冻结。这样既防止灾难性遗忘,又能让高层特征适应肺炎病灶。
# 初始冻结backbone,只训练Transformer和classifier for name, param in model.named_parameters(): param.requires_grad = False for param in model.transformer.parameters(): param.requires_grad = True for param in model.classifier.parameters(): param.requires_grad = True # 训练到第6轮时,解冻layer3(索引6)和layer4(索引7) def unfreeze_deep_layers(model, start_epoch, current_epoch): if current_epoch == start_epoch: for i in [6, 7]: # ResNet34的layer3和layer4 for p in model.features[i].parameters(): p.requires_grad = True逻辑说明:model.features里的索引0-7分别对应conv1、bn1、relu、maxpool、layer1、layer2、layer3、layer4。冻结前5轮后,Transformer已经学会了把特征图中的信息聚合到CLS token,此时再解冻layer3和layer4,梯度可以同时调整高层卷积和Transformer,避免低层特征被破坏。
参数说明:如果你的数据集有1万张以上,可以从第1轮就全量微调;如果只有2000张左右,建议把冻结轮数延长到10轮。解冻太晚会让Transformer学习的特征分布与最终backbone输出不匹配,验证集表现会出现一次抖动,这是正常的,接着训下去会回升。
3. 数据流程与400轮训练配置:从DICOM到batch_size=32、lr=0.0001的工程落地
模型结构定下来之后,真正决定成败的是数据管道和超参数。胸部X光片的数据量通常不大,类别分布也极不均衡,稍不注意就会得到“看起来loss很低,但实际是在瞎猜”的模型。
3.1 数据加载与预处理:灰度图、resize、增强和归一化
原始胸片可能是DICOM或PNG/JPG。这里以PNG为例,读取为灰度图后做基础变换。注意输入尺寸选256而不是常见224,因为224下采样32倍是7×7=49个token,256则是8×8=64个token,后者对肺部小病灶的定位粒度更好,Transformer计算量增加不大。代码中的预处理做了灰度转三通道,与预训练权重匹配。
import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T class PneumoniaXrayDataset(Dataset): def __init__(self, image_paths, labels, train=True): self.paths = image_paths self.labels = labels self.train = train self.base = T.Compose([ T.Resize((256, 256)), T.ToTensor(), ]) self.augment = T.Compose([ T.RandomRotation(10), T.RandomHorizontalFlip(), T.ColorJitter(brightness=0.2, contrast=0.2), ]) def __len__(self): return len(self.paths) def __getitem__(self, idx): img = Image.open(self.paths[idx]).convert('L') # 灰度 img = self.base(img) if self.train: img = self.augment(img) img = img.repeat(3, 1, 1) # [1,256,256] -> [3,256,256] mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1) std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1) img = (img - mean) / std return img, self.labels[idx]逻辑说明:convert('L')将RGB转成单通道灰度。repeat(3,1,1)将单通道沿着通道维度复制成3份,这样即使原图是暗色调的低对比度胸片,也能复用ImageNet的BN统计量。归一化用的均值和标准差是ImageNet的标准值,这与预训练权重保持一致,是迁移学习里最容易漏但最重要的一步。
参数说明:RandomRotation(10)旋转角度控制在±10度。胸片的肺尖、心影、肋膈角有解剖约束,旋转超过15度模型会学到不真实的形态。ColorJitter(brightness=0.2, contrast=0.2)模拟X光机曝光差异。ColorJitter对灰度图操作时只影响亮度和对比度,不会引入异常色偏。
3.2 类别不均衡与数据划分:按患者级别拆分而不是按图像随机拆分
很多人在这一步翻车:直接对所有胸片做随机train_test_split,同一个患者的多张片子被同时分到训练集和验证集。模型其实是在“认患者”,而不是在认病灶,验证精度虚高。正确做法是先把患者ID聚合,按患者分层划分。
from sklearn.model_selection import train_test_split import pandas as pd df = pd.read_csv('metadata.csv') # 字段:patient_id, image_path, label # 每个患者只取一行的label作为该患者的类别标签,用于分层 first_label = df.groupby('patient_id')['label'].first() patients = first_label.index labels = first_label.values train_patients, val_patients = train_test_split( patients, test_size=0.2, stratify=labels, random_state=42 ) train_df = df[df['patient_id'].isin(train_patients)] val_df = df[df['patient_id'].isin(val_patients)]逻辑说明:groupby('patient_id')['label'].first()确保在分层时每个患者只贡献一个样本的标签。stratify=labels让训练集和验证集里各类别比例接近原始数据集,避免某个类别只在验证集出现。划分之后,再把原始csv过滤成train_df和val_df,之后Dataset加载时就只用这两个子集。
参数说明:test_size=0.2假设你总患者数在数百到数千。如果患者数少于200,验证集会太小,改成0.15会好一些;如果每个患者都有多张视图,这样可以保证同一患者的所有片子在同一个partition里。
3.3 400轮训练循环:分阶段解冻、梯度裁剪与余弦退火
400轮指的是epoch数,而不是iteration。用固定学习率0.0001跑400轮几乎一定会震荡到loss不降。正确的做法是用余弦退火把学习率从1e-4平滑降到1e-6,并在训练初期对backbone做分阶段解冻。下面是可运行的训练骨架:
import torch import torch.nn as nn from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR model = build_model(num_classes=3, pretrained=True) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model.to(device) # 先冻结backbone所有参数 for param in model.features.parameters(): param.requires_grad = False optimizer = AdamW( filter(lambda p: p.requires_grad, model.parameters()), lr=1e-4, weight_decay=1e-4 ) scheduler = CosineAnnealingLR(optimizer, T_max=400, eta_min=1e-6) loss_fn = nn.CrossEntropyLoss() for epoch in range(400): # 从第6轮开始解冻layer3和layer4 if epoch == 5: for i in [6, 7]: for p in model.features[i].parameters(): p.requires_grad = True optimizer.add_param_group({'params': model.features[6].parameters(), 'lr': 1e-5}) optimizer.add_param_group({'params': model.features[7].parameters(), 'lr': 1e-5}) model.train() total_loss = 0.0 for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() logits = model(imgs) loss = loss_fn(logits, labels) 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) % 10 == 0: print(f'Epoch {epoch+1:3d}/400 | loss {total_loss/len(train_loader):.4f}')逻辑说明:filter(lambda p: p.requires_grad, model.parameters())一开始只给Transformer和分类器优化器参数。到第6轮解冻layer3和layer4时,用add_param_group给新解冻的参数单独设置更小的学习率1e-5,避免大梯度冲击。clip_grad_norm_对Transformer很重要,注意力层偶尔会产生异常大的梯度,不裁剪一个step就可能让loss变成nan。CosineAnnealingLR调度器在400轮内将lr平滑降到1e-6,后期模型在损失平面平坦区域稳定下来。
参数说明:batch_size=32时,学习率1e-4是比较稳的起点。如果显存不足改到16,学习率建议先降到5e-5,因为批量变小梯度噪声变大,相同的lr会导致更新方向不稳定。weight_decay=1e-4是Transformer微调的常规值,太大容易欠拟合,太小后期会过拟合。如果你用的是冻结策略,优化器里一开始没有backbone的权重,天然会减少正则化压力。
4. 混淆矩阵评估:真正用来指导临床决策的指标怎么算
训练400轮之后,报告里不能只写“accuracy 92%”这种话。在肺炎诊断任务中,漏诊一个细菌性肺炎的后果比把正常人错判成肺炎严重得多。混淆矩阵能揭示每一类错误具体是怎么分布的,也能帮你发现模型是不是在靠类别先验猜答案。
4.1 从验证集到多分类混淆矩阵:归一化与可视化
以三类为例:Normal(正常)、Bacterial(细菌性肺炎)、Viral(病毒性肺炎)。我们需要跑一遍验证集,将所有预测结果和真实标签收集起来,用sklearn生成混淆矩阵,并按行归一化。因为类别样本数不均衡,归一化后才能看出每一类的召回率差异。
from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt import numpy as np def evaluate_confusion(model, dataloader, device, class_names): model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for imgs, labels in dataloader: imgs = imgs.to(device) logits = model(imgs) preds = logits.argmax(dim=1).cpu() all_preds.extend(preds.numpy()) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) cm_norm = cm.astype('float') / (cm.sum(axis=1, keepdims=True) + 1e-8) plt.figure(figsize=(8, 6)) sns.heatmap(cm_norm, annot=True, fmt='.2f', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.savefig('confusion_matrix_normalized.png', dpi=200) print(classification_report(all_labels, all_preds, target_names=class_names, digits=4))逻辑说明:cm.sum(axis=1)得到每行真值类别的总样本数,用行归一化后,每一行代表“真实为该类别时,被模型预测成各个类别的比例”。这样即使Viral类只有50张,Normal类有500张,也能清楚看到Viral类的召回率是否偏低。classification_report会输出每个类别的precision、recall和f1,这些值对医生来说比accuracy更有临床意义。
参数说明:fmt='.2f'显示两位小数,足够看出0.85和0.88的差别。digits=4让报告保留4位小数,避免在类别样本少的时候显示成0.0000导致误判。dpi=200保证输出图片在投影或论文里依然清晰。如果你要回到二分类(Pneumonia vs Normal),同样的代码可以直接跑,只是class_names换成两个元素。
4.2 从混淆矩阵计算敏感度、特异度与F1得分
多分类任务里,敏感度就是每个类别的召回率,特异度需要额外计算。医学报告通常需要这几个数字,下面这段代码直接给出一份可粘贴的表。
def compute_clinical_metrics(cm, class_names): recall_each = np.diag(cm) / (cm.sum(axis=1) + 1e-8) precision_each = np.diag(cm) / (cm.sum(axis=0) + 1e-8) f1_each = 2 * precision_each * recall_each / (precision_each + recall_each + 1e-8) print(f"{'Class':<12}{'Precision':<10}{'Recall':<10}{'F1-score':<10}") for i, name in enumerate(class_names): print(f"{name:<12}{precision_each[i]:<10.4f}{recall_each[i]:<10.4f}{f1_each[i]:<10.4f}") # 多分类特异度:对类别i,TN是非i样本中预测为非i的数量,FP是非i样本中预测为i的数量 tn = cm.sum() - cm.sum(axis=1) - cm.sum(axis=0) + np.diag(cm) fp = cm.sum(axis=0) - np.diag(cm) spec_each = tn / (tn + fp + 1e-8) print(f"\nMacro-Precision: {precision_each.mean():.4f}") print(f"Macro- Recall: {recall_each.mean():.4f}") print(f"Macro-Specificity: {spec_each.mean():.4f}") print(f"Macro-F1: {f1_each.mean():.4f}")逻辑说明:敏感度反映“漏诊率”,特异度反映“误诊率”。在多分类里,对类别i,所有真实不是i的样本总数是cm.sum() - cm.sum(axis=1[i]);其中预测为i且真实不是i的数量是cm.sum(axis=0[i]) - cm_diag[i]。这段代码在类别不均衡时依然有效。宏平均的F1将三个类别同等对待,比加权F1更能暴露少数类上的Performance。
参数说明:1e-8防除零。如果某个类别的precision或recall为0,宏平均会立刻拉低F1,这样比只看准确率诚实得多。如果你的类别定义是二分类,特异度公式仍然成立,只是输出的宏平均等于整体特异度。
5. 训练与评估中的常见坑:这几处让多少人白白跑完400轮
踩坑经验比模型结构更重要。下面五条都是实际训练中反复出现的典型问题,每条都按“现象 → 原因 → 解决”写清楚。
5.1 预训练权重加载报错:fc层或第一层卷积shape不匹配
现象:model.features.load_state_dict(backbone.state_dict())报错,提示“size mismatch for fc.weight”,或者“Missing key(s): fc.weight”。
原因:官方resnet34的state_dict里包含fc.weight和fc.bias,而我们自定义的model.features没有fc层。另外如果修改了第一层卷积的输入通道,conv1.weight的shape也会对不上。
解决:截取backbone时使用list(resnet.children())[:-2],这样state_dict里就已经没有fc层了。如果改了第一层卷积,就不要用load_state_dict加载conv1.weight,或者改用复制三通道输入的方式彻底绕开这个问题。还有一种保险做法是加载时带一个strict=False过滤掉不匹配的key,但建议只在你知道自己在改什么时才用。
5.2 训练loss在下降,但验证集准确率一直在50%左右
现象:每个epoch打印的loss都稳步下降,看起来模型在学习,但验证集accuracy只有五成,且分类报告里某一类的recall为0。
原因:数据划分没有按患者分层,导致同一个人不同拍摄角度的胸片出现在训练集和验证集,模型记住的是患者特征。另一个常见原因是类别极不均衡且CrossEntropyLoss没有设置类别权重,模型把所有样本都预测为多数类(Normal),loss看起来低但少数类全部被忽略。
解决:按患者ID分层划分,具体见3.2节代码。对于类别不均衡,在损失函数里传入weight,权重与样本频次成反比,例如nn.CrossEntropyLoss(weight=torch.tensor([1.0, 3.0, 5.0]).to(device))。或者使用torch.nn.functional.focal_loss(需要自己实现)。先在控制台打印每个类别的样本数,根据比例设定weight,通常能让少数类的recall有明显提升。
5.3 显存不足,把batch_size从32改到16后loss发散
现象:原配置batch=32跑得好好的,为了塞进单卡改到16,前几个step loss就冲到2.0以上,之后持续震荡不下降。
原因:学习率与批量大小高度相关。批量变小意味着每个step的梯度噪声变大,同样的学习率会让参数更新方向偏离真实梯度。直接改batch而不动lr是典型翻车操作。
解决:按线性缩放规则调整学习率,常见做法是lr_new = lr_old * (new_batch / old_batch),即从1e-4降为5e-5。如果你用的是AdamW,可以先降到7e-5观察10个epoch。更稳妥的是在优化器上单独给Transformer和backbone不同lr,backbone建议始终比Transformer低一个数量级。如果必须用batch=8,学习率降到2e-5并配合前20个step的warmup。
5.4 混淆矩阵对角线“好看”,但仔细看某一类敏感度极低
现象:整体accuracy在90%以上,但看归一化混淆矩阵时,“Bacterial”这一行的召回率只有0.55,有0.3被错判成了“Viral”。
原因:类别特征之间存在语义重叠,或者该类别训练样本过少。如果三个类别的样本数差距很大,多数类会主导梯度更新,少数类的决策边界被挤压。另一个原因是数据增强过度,比如旋转角度过大导致肺纹理方向失真。
解决:先按患者分层划分并设置class weight。若Bacterial和Viral混淆严重,可以考虑合并成统一的Pneumonia再进行二分类,然后在二分类基础上再分亚型,形成两阶段模型。这样主诊断的敏感度会高很多,亚型分类再用单独的模型负责。
5.5 Transformer注意力权重趋于平均:模型退化成线性特征拼接
现象:训练结束后可视化CLS token在最后一层的注意力权重,发现所有位置几乎均匀分布,没有明显的热点区域。验证集loss虽然正常,但Grad-CAM热力图非常分散,看不出对病灶的聚焦。
原因:位置编码没被正确加入,或者Transformer的dropout设得过高(比如0.5),导致注意力被随机化。另一个常见原因是backbone在训练早期就被全量解冻,CNN把信息压到少数字符里,Transformer偷懒,学会了忽略注意力。
解决:将Transformer的dropout降到0.1。检查self.pos_embed是否参与梯度更新,它应当在训练中发生变化。另外严格按照2.3节的冻结策略执行,让训练前5轮只调整Transformer,迫使它学会利用位置信息。如果问题依然存在,可尝试把nhead降到4,提升每个注意力头的维度,让注意力分配更有区分度。
6. 进阶验证:用Grad-CAM和注意力可视化确认模型在看肺尖还是背景
训练完400轮后,准确率和混淆矩阵只回答了“分得对不对”,回答不了“医生凭什么信你”。胸片诊断要落地,必须让人看到模型在做决策时关注哪些像素。这里推荐同时做两件验证:一是用Grad-CAM看CNN特征对分类的贡献区域,二更简单——直接取Transformer最后一层的注意力权重,映射到原图上,看CLS token在呼应哪些空间位置。
具体做法是在模型forward过程中注册一个hook,取Transformer输出的attention map。PyTorch的TransformerEncoderLayer在forward时会返回attention权重,但需要稍微侵入代码。如果你不想改模型,可以在forward里把self.transformer换成自写的编码器层,并收集每层的attn_output_weights。拿到权重后,对CLS token那一行的权重做平均或取最后一层,归一化到0-1,用matplotlib叠加在原图上。我习惯把结果和Grad-CAM并排放一起,如果两个热力图的中心都不在肺野而在纵膈区域,那说明模型学到的线索是背景伪影,需要回到数据层面去修复预处理、裁剪或更严格的标注。
对于患者级别的效果验证,还有一个很实用的技巧:对同一张胸片分别做原图预测、左右翻转预测、以及遮挡四个象限后的预测,观察logits变化。如果翻转后类别翻转、或者遮挡某块背景后预测概率明显变化,说明模型对胸片的解剖结构理解不足。我会在训练结束前,用验证集里三张最容易误判的片子跑一组这种“稳定性测试”,把结果作为模型能否提交给医生review的硬性门槛。
经过这么一轮,你会对“400轮到底够不够”有自己的判断。我的习惯是:每次调完一个超参数,就把混淆矩阵、注意力热图、几张典型误判样本存到一个以日期命名的文件夹里,下次迭代直接对比。这套流程看起来繁琐,但可以省掉很多重复跑400轮的时间。希望帮到你。
本文还有配套的精品资源,点击获取