基于Transformer与ResNet34的医学影像肺炎诊断系统实战 📅 发布时间:2026/9/3 9:44:13 👁 浏览次数: 简介本资源是一套面向医学影像AI初学者与临床辅助诊断研究者的胸部X光肺炎智能识别系统基于Transformer与ResNet34双主干融合设计解决小样本医学图像分类中特征提取不足与泛化性弱的典型问题。压缩包共13个文件9个Python核心脚本、1个说明文档、1个类别索引JSON、1个Markdown说明及1个TXT配置指南总大小仅55KB轻量易部署其中train2.py与model_MSG.py实现Transformer模块嵌入confusion_matrix_resnet.py和confusion_matrix_MSG.py分别支持双模型评估predict.py提供端到端推理接口utils.py与my_dataset.py封装数据增强与加载逻辑。已有54人学习下载配套附赠资源.docx含项目背景与使用流程说明文件.txt详述PyTorch环境配置、训练参数复现要点及混淆矩阵解读方法README.md则梳理了完整目录结构与模块调用关系助用户快速理解架构设计并开展迁移实验。1. 项目概述当Transformer遇见医学影像最近在做一个挺有意思的活儿一个基于Transformer架构的胸部X光肺炎诊断系统。这项目听起来挺唬人但说白了就是用深度学习模型让计算机学会看X光片判断病人有没有得肺炎。肺炎诊断这事儿在临床上挺关键的尤其是对于儿童和老年人早发现早治疗能省不少事儿。传统的诊断依赖放射科医生肉眼读片费时费力不说还容易因为疲劳或经验差异导致误判。所以用AI来辅助甚至部分自动化这个流程就成了一个很有价值的方向。这个项目标题信息量不小直接把核心框架、骨干网络、训练策略和评估方法都点出来了。它用上了这两年火得不能再火的Transformer架构但又不是直接拿Vision Transformer那种庞然大物硬上而是结合了经典的ResNet34作为特征提取的“地基”加载了预训练权重这算是非常务实且高效的组合拳。训练了整整400轮批量大小Batch Size设成32学习率Learning Rate是0.0001这些参数一看就是经过了精心调校不是随便填的。最后还用混淆矩阵做了评估这是分类任务里检验模型“真功夫”的硬指标。整个项目基于PyTorch框架实现生态成熟复现起来也方便。如果你是刚接触医学影像AI的开发者或者对Transformer如何应用到视觉任务特别是这种关键的医疗诊断场景感兴趣那这个项目会是一个绝佳的切入点。它没有一味追求最前沿、最复杂的模型而是在效果和效率之间做了很好的平衡里面的很多设计思路和调参经验直接抄作业都能学到不少东西。2. 核心思路与架构设计解析2.1 为什么是“Transformer ResNet”的混合架构看到“基于Transformer架构”很多人第一反应可能是直接用ViTVision Transformer或者Swin Transformer这类纯Transformer的视觉模型。但在医疗影像尤其是数据量可能不那么庞大的特定任务如某个医院的X光数据集上纯Transformer模型有两个明显的挑战一是对数据量要求高二是局部特征捕捉能力在早期层可能不如卷积神经网络CNN那么直接和高效。所以这个项目采用了一个非常聪明的策略用CNNResNet34作为特征提取器用Transformer编码器部分作为特征增强与关系建模器。你可以把它想象成一条流水线特征提取ResNet34输入的X光图片比如224x224像素首先经过ResNet34。ResNet34是一个深度残差网络它在ImageNet这样的大规模数据集上预训练过已经学会了识别边缘、纹理、形状等通用视觉特征。这一步相当于让一个经验丰富的“低年资医生”先快速扫一眼片子找出所有可疑的病灶区域和生理结构。特征转换与序列化从ResNet34的最后一层卷积层出来的特征图Feature Map其形状通常是[batch_size, channels, height, width]。Transformer处理的是序列Sequence所以我们需要把这个二维的特征图“拍平”成一个一维的序列。具体做法是把height和width两个维度合并将每个空间位置共height * width个的特征向量长度为channels视为序列中的一个“词”Token。这样我们就得到了一个长度为height * width每个词维度为channels的序列。关系建模与增强Transformer Encoder这个序列被送入Transformer的编码器。Transformer的核心是自注意力机制Self-Attention它能让序列中的每一个“位置”都去关注序列中的所有其他“位置”。对应到我们的X光片这意味着模型可以学习到“左下肺叶的某个阴影”与“右肺门区域的淋巴结”之间的潜在关联或者判断某个高亮区域是血管纹理还是炎症浸润。这种全局的、动态的关系建模能力是传统CNN通过堆叠卷积层逐步扩大感受野所难以媲美的尤其适合分析需要整体把握的医学影像。分类头Classifier Head经过Transformer编码器增强后的序列通常会取一个特殊的[CLS]标记在序列开头添加对应的输出或者对所有位置的输出进行全局平均池化得到一个固定维度的、融合了全局信息的特征向量。最后这个向量通过一个全连接层Linear Layer映射到分类数例如2类正常 vs 肺炎完成诊断。这种混合架构的优势在于它结合了CNN强大的局部特征提取能力和Transformer卓越的全局上下文建模能力。ResNet34提供了高质量、稠密的初始特征降低了Transformer直接处理原始像素的学习难度而Transformer则赋予了模型“纵观全局、联系思考”的医生般的思维能力。2.2 关键超参数设置背后的考量标题里给出的几个超参数不是随便选的每一个都有其道理批量大小Batch Size 32在GPU内存允许的范围内较大的Batch Size能使梯度估计更准确训练更稳定。32是一个在常用显卡如RTX 3080/4090 显存12G-24G上处理224x224图像比较均衡的值。太小如8则噪声大、训练慢太大如128可能超出显存且可能损害模型泛化性能。学习率Learning Rate 0.0001这是一个相对较小的学习率。因为我们在使用预训练权重Pre-trained Weights。ResNet34在ImageNet上预训练的权重已经包含了丰富的视觉知识我们的任务肺炎诊断可以看作是在此基础上的“微调”Fine-tuning。使用较小的学习率是为了避免在微调初期就用太大的步子“破坏”掉这些宝贵的预训练特征让模型能够温和地适应新任务。0.0001即1e-4是微调任务中非常经典和常用的初始学习率。训练轮数Epochs 400400轮听起来很多但在深度学习尤其是医学图像分析中并不罕见。原因有二一是数据集可能本身不是特别大例如几千到几万张需要更多轮次让模型充分学习二是我们使用了较小的学习率收敛速度会相对慢一些需要更长的训练周期来达到最优性能。当然在实际操作中我们一定会配合早停法Early Stopping即当模型在验证集上的性能不再提升时就停止训练防止过拟合400可以看作是设置的一个上限。注意直接训练400轮而不加任何监控是危险的。务必在训练循环中每训练完一个Epoch就在一个独立的验证集上评估模型性能如计算准确率、F1分数并保存验证集上表现最好的模型权重。当连续多个Epoch如10或20个验证集性能不再提升时就应触发早停。3. 数据准备与预处理实战3.1 数据集获取与结构组织对于肺炎X光片一个公开且常用的基准数据集是“Chest X-Ray Images (Pneumonia)”数据集它通常包含“正常”Normal和“肺炎”Pneumonia两类图像并已划分好训练集、验证集和测试集。假设我们下载的数据集原始结构如下chest_xray/ ├── train/ │ ├── NORMAL/ │ │ ├── normal_image1.jpeg │ │ └── ... │ └── PNEUMONIA/ │ ├── pneumonia_image1.jpeg │ └── ... ├── val/ │ ├── NORMAL/ │ └── PNEUMONIA/ └── test/ ├── NORMAL/ └── PNEUMONIA/在代码中我们通常使用torchvision.datasets.ImageFolder来加载这种按类别分文件夹存储的数据它会自动根据文件夹名生成标签。3.2 图像预处理与增强流水线医学影像的预处理至关重要直接影响到模型的学习效率和最终性能。我们使用torchvision.transforms来构建一个处理流水线。对于训练集我们需要进行增强Augmentation以增加数据多样性模拟各种拍摄情况提升模型鲁棒性from torchvision import transforms train_transform transforms.Compose([ transforms.Resize((256, 256)), # 首先缩放到稍大尺寸 transforms.RandomCrop(224), # 随机裁剪到模型输入尺寸224x224 transforms.RandomHorizontalFlip(p0.5), # 随机水平翻转肺炎通常左右对称增强有效 transforms.RandomRotation(degrees10), # 小幅随机旋转应对拍摄角度差异 transforms.ColorJitter(brightness0.1, contrast0.1), # 微调亮度和对比度模拟不同设备差异 transforms.ToTensor(), # 转换为Tensor并归一化像素值到[0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], # ImageNet的均值 std[0.229, 0.224, 0.225]) # ImageNet的标准差 ])Resize到(256, 256)再RandomCrop到(224, 224)是标准操作既提供了随机裁剪的增强效果又确保了输入尺寸统一。RandomHorizontalFlip和RandomRotation是常用的空间增强。对于X光片水平翻转是合理的生理对称性增强。ColorJitter用于增强模型对图像亮度、对比度变化的鲁棒性因为不同医院、不同设备的X光片在灰度表现上可能有差异。Normalize使用的均值和标准差是ImageNet数据集的统计值。这是因为我们的ResNet34是在ImageNet上预训练的其权重适应了这种分布。保持输入数据分布与预训练时一致能更快更好地进行微调。对于验证集和测试集我们不需要随机性只需要进行确定性的预处理val_test_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), # 中心裁剪确保评估一致性 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])3.3 构建数据加载器DataLoader使用ImageFolder和DataLoader来高效地加载和批处理数据from torchvision.datasets import ImageFolder from torch.utils.data import DataLoader # 创建数据集对象 train_dataset ImageFolder(rootchest_xray/train, transformtrain_transform) val_dataset ImageFolder(rootchest_xray/val, transformval_test_transform) test_dataset ImageFolder(rootchest_xray/test, transformval_test_transform) # 创建数据加载器 batch_size 32 train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_sizebatch_size, shuffleFalse, num_workers4, pin_memoryTrue) test_loader DataLoader(test_dataset, batch_sizebatch_size, shuffleFalse, num_workers4)shuffleTrue仅用于训练集在每个Epoch开始时打乱数据顺序有助于模型学习更泛化的特征。num_workers设置用于数据加载的子进程数可以加速数据读取。根据你的CPU核心数调整通常设为4或8。pin_memoryTrue当使用GPU时将数据固定到页锁定内存可以加速从CPU到GPU的数据传输。4. 模型构建混合架构的PyTorch实现4.1 骨干网络加载预训练的ResNet34我们截取ResNet34的前面部分通常到最后一个卷积块结束丢弃其原有的全局平均池化层和全连接分类头将其作为一个特征提取器。import torch import torch.nn as nn from torchvision import models class HybridPneumoniaModel(nn.Module): def __init__(self, num_classes2, embed_dim512, num_heads8, num_layers3): super(HybridPneumoniaModel, self).__init__() # 1. 加载预训练的ResNet34骨干网络 backbone models.resnet34(weightsmodels.ResNet34_Weights.IMAGENET1K_V1) # 移除最后的全连接层和平均池化层 self.feature_extractor nn.Sequential(*list(backbone.children())[:-2]) # 此时对于输入224x224的图像self.feature_extractor的输出是 # [batch_size, 512, 7, 7] (512个通道7x7的空间大小) # 获取ResNet最终特征图的通道数 self.in_channels 512 self.feature_size 7 # 假设输入224经过ResNet34后特征图大小为7x7 # 2. 将特征图转换为序列 (Flatten spatial dimensions) self.seq_length self.feature_size * self.feature_size # 49 self.embed_dim embed_dim # 一个线性层将ResNet的通道数512映射到Transformer期望的嵌入维度 self.projection nn.Linear(self.in_channels, self.embed_dim) # 3. Transformer编码器 encoder_layer nn.TransformerEncoderLayer( d_modelself.embed_dim, nheadnum_heads, dim_feedforward2048, dropout0.1, activationrelu, batch_firstTrue # 输入输出形状为 (batch, seq_len, embed_dim) ) self.transformer_encoder nn.TransformerEncoder(encoder_layer, num_layersnum_layers) # 4. 分类头 self.classifier nn.Sequential( nn.LayerNorm(self.embed_dim), nn.Linear(self.embed_dim, 256), nn.ReLU(), nn.Dropout(0.3), nn.Linear(256, num_classes) ) def forward(self, x): # 阶段1: CNN特征提取 cnn_features self.feature_extractor(x) # [B, 512, 7, 7] # 阶段2: 重塑并投影到嵌入空间 B, C, H, W cnn_features.shape # 将空间维度展平为序列 cnn_features_flat cnn_features.view(B, C, -1).permute(0, 2, 1) # [B, 49, 512] # 投影到Transformer的嵌入维度 projected_features self.projection(cnn_features_flat) # [B, 49, embed_dim] # 阶段3: Transformer编码 transformer_features self.transformer_encoder(projected_features) # [B, 49, embed_dim] # 阶段4: 全局特征聚合与分类 # 方法1: 使用序列第一个位置的特征类似[CLS] token这里我们取平均 global_feature transformer_features.mean(dim1) # [B, embed_dim] # 方法2: 也可以额外添加一个可学习的[CLS] token这里为简化取平均 # 阶段5: 分类 output self.classifier(global_feature) # [B, num_classes] return output关键点解析models.resnet34(weights...)使用torchvision中带预训练权重的ResNet34。IMAGENET1K_V1代表在ImageNet上训练的权重。list(backbone.children())[:-2]children()返回模型的所有一级子模块。ResNet34的结构大致是卷积层 - BN层 - ReLU - 池化层 - 4个layer每个layer包含多个残差块 - 自适应平均池化层AdaptiveAvgPool2d - 全连接层Linear。[:-2]去掉了最后的平均池化层和全连接层保留了特征提取部分。self.projection由于ResNet34最后一层特征通道数是512而Transformer编码器期望的嵌入维度embed_dim我们可以自定义例如512这个线性层负责进行维度对齐。nn.TransformerEncoderLayer和nn.TransformerEncoder这是PyTorch内置的标准Transformer编码器实现。d_model是嵌入维度nhead是多头注意力的头数dim_feedforward是前馈网络中间层的维度dropout用于防止过拟合。在forward函数中cnn_features.view(B, C, -1).permute(0, 2, 1)这一步是关键。它将[B, 512, 7, 7]的特征图重塑为[B, 512, 49]然后转置为[B, 49, 512]使得49个空间位置7x7成为序列长度每个位置是一个512维的特征向量。最终分类时我们对Transformer输出的所有位置的特征取平均mean(dim1)得到一个全局特征向量再送入分类头。4.2 模型初始化与设备配置# 初始化模型 model HybridPneumoniaModel(num_classes2, embed_dim512, num_heads8, num_layers3) # 设备配置 device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) # 打印模型参数量 total_params sum(p.numel() for p in model.parameters()) trainable_params sum(p.numel() for p in model.parameters() if p.requires_grad) print(fTotal parameters: {total_params:,}) print(fTrainable parameters: {trainable_params:,})实操心得在微调预训练模型时有时我们并不想更新所有参数。例如可以冻结ResNet34特征提取器前几层的权重只训练后面的层和Transformer部分。这可以通过设置param.requires_grad False来实现。对于小数据集冻结前面层有助于防止过拟合对于大数据集解冻全部层进行微调可能效果更好。这是一个需要根据实际情况调整的超参数。5. 训练循环与优化策略5.1 损失函数与优化器选择对于二分类任务正常/肺炎我们使用带Logits的二元交叉熵损失BCEWithLogitsLoss。虽然我们的输出是两个节点但BCEWithLogitsLoss配合sigmoid激活内置在损失函数中可以很好地处理。当然使用CrossEntropyLoss需要模型输出未经激活的logits并将标签设置为0和1也是完全等价的这里选择前者。优化器选择AdamW它是Adam优化器的一个变种加入了权重衰减Weight Decay的正则化通常能获得更好的泛化性能。import torch.optim as optim from torch.nn import BCEWithLogitsLoss criterion BCEWithLogitsLoss() # 模型输出不需要sigmoid损失函数内部处理 # 或者 criterion nn.CrossEntropyLoss() 此时模型最后一层不需要sigmoid/softmax optimizer optim.AdamW(model.parameters(), lr0.0001, weight_decay1e-4) # 学习率调度器在训练过程中动态调整学习率有助于后期收敛 scheduler optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.5, patience5, verboseTrue) # 这里监控验证集准确率max当其在5个epoch内未提升时学习率减半。5.2 完整的训练与验证循环训练循环是项目的核心引擎。我们需要记录损失和准确率并在验证集上定期评估。def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (images, labels) in enumerate(dataloader): images, labels images.to(device), labels.to(device).float().unsqueeze(1) # BCEWithLogitsLoss需要标签为float且形状匹配 # 清零梯度 optimizer.zero_grad() # 前向传播 outputs model(images) loss criterion(outputs, labels) # 反向传播与优化 loss.backward() optimizer.step() # 统计 running_loss loss.item() * images.size(0) # 计算准确率将sigmoid输出0.5的视为正类肺炎 predicted (torch.sigmoid(outputs) 0.5).float() total labels.size(0) correct (predicted labels).sum().item() # 可选每N个batch打印一次进度 if batch_idx % 50 0: print(f Batch [{batch_idx}/{len(dataloader)}], Loss: {loss.item():.4f}) epoch_loss running_loss / total epoch_acc 100. * correct / total return epoch_loss, epoch_acc def validate(model, dataloader, criterion, device): model.eval() running_loss 0.0 correct 0 total 0 with torch.no_grad(): # 关闭梯度计算节省内存和计算 for images, labels in dataloader: images, labels images.to(device), labels.to(device).float().unsqueeze(1) outputs model(images) loss criterion(outputs, labels) running_loss loss.item() * images.size(0) predicted (torch.sigmoid(outputs) 0.5).float() total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc 100. * correct / total return epoch_loss, epoch_acc # 主训练循环 num_epochs 400 best_val_acc 0.0 patience_counter 0 patience 15 # 早停耐心值 train_losses, train_accs [], [] val_losses, val_accs [], [] for epoch in range(num_epochs): print(fEpoch [{epoch1}/{num_epochs}]) # 训练阶段 train_loss, train_acc train_one_epoch(model, train_loader, criterion, optimizer, device) train_losses.append(train_loss) train_accs.append(train_acc) print(fTrain Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%) # 验证阶段 val_loss, val_acc validate(model, val_loader, criterion, device) val_losses.append(val_loss) val_accs.append(val_acc) print(fVal Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%) # 学习率调度 scheduler.step(val_acc) # 保存最佳模型 if val_acc best_val_acc: best_val_acc val_acc torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), val_acc: val_acc, }, best_pneumonia_model.pth) print(f - Best model saved with Val Acc: {val_acc:.2f}%) patience_counter 0 # 重置早停计数器 else: patience_counter 1 print(f - No improvement for {patience_counter} epoch(s).) # 早停判断 if patience_counter patience: print(fEarly stopping triggered at epoch {epoch1}) break print(Training finished.)训练循环要点model.train()和model.eval()在训练和验证/测试前必须切换模式这会影响Dropout、BatchNorm等层的行为。optimizer.zero_grad()在每个batch开始前清零梯度防止梯度累积。loss.backward()和optimizer.step()标准的反向传播和参数更新。with torch.no_grad()在验证和测试时使用禁用自动求导大幅减少内存消耗并加速计算。早停Early Stopping这是防止过拟合的关键技术。我们监控验证集准确率如果连续patience个epoch这里设15都没有提升就认为模型已经过拟合停止训练。模型保存只保存验证集上性能最好的模型权重而不是最后一个epoch的。6. 模型评估与混淆矩阵分析训练完成后我们需要在独立的测试集上评估模型的最终性能以检验其泛化能力。混淆矩阵Confusion Matrix是评估分类模型最直观的工具之一。6.1 在测试集上进行最终评估首先加载我们保存的最佳模型进行测试。# 加载最佳模型 checkpoint torch.load(best_pneumonia_model.pth) model.load_state_dict(checkpoint[model_state_dict]) model.eval() # 切换到评估模式 # 在测试集上评估 test_loss, test_acc validate(model, test_loader, criterion, device) print(fFinal Test Results - Loss: {test_loss:.4f}, Accuracy: {test_acc:.2f}%)6.2 生成并可视化混淆矩阵混淆矩阵能告诉我们模型在每一类上的具体表现比如把多少正常样本误判为肺炎假阳性又把多少肺炎样本漏掉了假阴性。这对于医疗诊断至关重要因为假阴性漏诊的后果可能比假阳性误诊更严重。import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns def evaluate_with_cm(model, dataloader, device): model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in dataloader: images images.to(device) outputs model(images) # 获取预测类别 (0或1) preds (torch.sigmoid(outputs) 0.5).int().squeeze().cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.int().cpu().numpy()) return np.array(all_labels), np.array(all_preds) # 获取测试集的真实标签和预测标签 y_true, y_pred evaluate_with_cm(model, test_loader, device) # 计算混淆矩阵 cm confusion_matrix(y_true, y_pred, labels[0, 1]) # 假设0正常(NORMAL), 1肺炎(PNEUMONIA) print(Confusion Matrix:) print(cm) # 打印详细的分类报告 print(\nClassification Report:) print(classification_report(y_true, y_pred, target_names[NORMAL, PNEUMONIA])) # 可视化混淆矩阵 plt.figure(figsize(8,6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabels[NORMAL, PNEUMONIA], yticklabels[NORMAL, PNEUMONIA]) plt.ylabel(True Label) plt.xlabel(Predicted Label) plt.title(Confusion Matrix on Test Set) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi300) plt.show()混淆矩阵解读 假设输出如下Confusion Matrix: [[300 25] [ 15 410]]这是一个2x2的矩阵左上角 (300)真阴性True Negative, TN。模型正确预测为“正常”的正常样本数。右上角 (25)假阳性False Positive, FP。模型错误地将“正常”样本预测为“肺炎”。误诊左下角 (15)假阴性False Negative, FN。模型错误地将“肺炎”样本预测为“正常”。漏诊临床风险高右下角 (410)真阳性True Positive, TP。模型正确预测为“肺炎”的肺炎样本数。从这个矩阵我们可以计算出更丰富的指标准确率Accuracy (TPTN) / Total (410300) / (3002515410) ≈ 94.7%精确率/查准率Precision TP / (TPFP) 410 / (41025) ≈ 94.3% 在所有预测为肺炎的样本中真正患肺炎的比例召回率/查全率Recall TP / (TPFN) 410 / (41015) ≈ 96.5% 在所有实际患肺炎的样本中被模型找出来的比例F1分数 2 * (Precision * Recall) / (Precision Recall) ≈ 95.4% 精确率和召回率的调和平均classification_report会直接给出这些指标。在医疗场景下我们往往更关注召回率Recall/Sensitivity即尽可能少地漏诊肺炎病例。如果召回率偏低可能需要调整分类阈值从0.5调低或者使用代价敏感学习给漏诊更高的惩罚权重。6.3 可视化注意力可选但推荐为了增强模型的可解释性我们可以尝试可视化Transformer的自注意力权重看看模型在做出诊断时“关注”了图像的哪些区域。这有助于医生理解AI的判断依据建立信任。 一种常见的方法是使用Grad-CAM梯度加权类激活映射或其变种但针对Transformer模型我们可以直接提取自注意力层的注意力图。# 这是一个简化的示例展示如何获取最后一层Transformer编码器的注意力权重 def get_attention_maps(model, image_tensor, device): model.eval() image_tensor image_tensor.unsqueeze(0).to(device) # 增加batch维度 # 前向传播并保留中间注意力 with torch.no_grad(): # 假设我们修改了模型使其在forward中返回注意力权重 # 这里需要根据你的模型具体实现来调整 cnn_features model.feature_extractor(image_tensor) B, C, H, W cnn_features.shape cnn_features_flat cnn_features.view(B, C, -1).permute(0, 2, 1) projected_features model.projection(cnn_features_flat) # 手动运行Transformer编码器层以获取注意力 # 注意这是一个简化示例实际中可能需要修改模型结构 x projected_features for layer in model.transformer_encoder.layers: # 这里需要调用layer的自注意力前向传播并保存注意力权重 # 通常需要修改源码或使用hook机制 pass # 返回注意力图 (序列长度 x 序列长度) 或聚合后的空间注意力图 (H x W) # 具体实现略取决于模型细节 return attention_map # 获取一张测试图像的注意力图 sample_image, sample_label test_dataset[0] attn_map get_attention_maps(model, sample_image, device) # 将attn_map (49, 49) 或聚合后的 (7,7) 上采样回原图尺寸并与原图叠加显示实现完整的注意力可视化需要深入模型内部可能涉及注册前向钩子forward hook来捕获注意力权重。这能提供宝贵的模型洞察但实现复杂度较高。7. 部署推理与常见问题排查7.1 单张图像推理脚本训练好的模型最终要用于实际预测。下面是一个简单的推理脚本。def predict_single_image(model_path, image_path, transform, device, class_names[NORMAL, PNEUMONIA]): 对单张胸部X光图像进行肺炎诊断预测。 参数: model_path: 保存的模型权重路径 (.pth) image_path: 待预测的图片路径 transform: 与训练时验证集相同的预处理变换 device: CPU或CUDA设备 class_names: 类别名称列表 # 1. 加载模型 model HybridPneumoniaModel(num_classes2, embed_dim512, num_heads8, num_layers3) checkpoint torch.load(model_path, map_locationdevice) model.load_state_dict(checkpoint[model_state_dict]) model.to(device) model.eval() # 2. 加载并预处理图像 from PIL import Image image Image.open(image_path).convert(RGB) # 确保是三通道 image_tensor transform(image).unsqueeze(0).to(device) # 增加batch维度 # 3. 预测 with torch.no_grad(): outputs model(image_tensor) # 使用sigmoid得到概率 probability torch.sigmoid(outputs).squeeze().item() # 根据阈值判断类别 pred_class_idx 1 if probability 0.5 else 0 pred_class_name class_names[pred_class_idx] print(f图像: {image_path}) print(f预测为 {pred_class_name} 的概率: {probability:.4f}) print(f诊断结果: {pred_class_name}) if pred_class_idx 1: print(提示: 该影像提示肺炎可能性建议结合临床进一步检查。) else: print(提示: 该影像未见明确肺炎征象。) # 注意此结果仅为AI辅助参考不能作为最终诊断依据。 return pred_class_name, probability # 使用示例 transform val_test_transform # 使用验证/测试时的预处理 result predict_single_image(best_pneumonia_model.pth, path_to_your_xray.jpg, transform, device)7.2 训练与推理中的常见问题与解决方案在实际操作中你几乎一定会遇到下面这些问题。这里是我踩过坑后总结的经验问题现象可能原因排查与解决方案训练损失Loss不下降准确率徘徊在50%左右随机猜测1.学习率过大或过小。2.模型未正确训练如某些层被冻结但本不该冻结。3.数据标签错误或预处理出错如归一化参数不对。4.损失函数或优化器设置错误。1. 尝试调整学习率如1e-3, 1e-4, 1e-5。使用学习率查找器LR Finder是个好方法。2. 检查model.parameters()中requires_grad为True的参数数量确保关键层如Transformer、分类头是可训练的。3. 可视化几张预处理后的图像检查是否正常。确保标签0/1对应正确的文件夹。4. 检查criterion和optimizer是否与模型输出、任务类型匹配。对于二分类确保标签是float且形状为[batch_size, 1]。验证集准确率远低于训练集过拟合严重1.模型过于复杂或训练数据太少。2.数据增强不够。3.训练时间太长没有使用早停。1. 增加数据量如果可能。尝试简化模型如减少Transformer层数num_layers。增加Dropout率或权重衰减weight_decay。2. 增强数据增强策略如添加随机裁剪、色彩抖动、模糊等。3.务必使用早停。降低patience值使其更早停止。训练时GPU内存溢出OOM1.批量大小Batch Size太大。2.图像尺寸太大。3.模型太大。1. 减小batch_size如从32降到16。可以使用梯度累积Gradient Accumulation来模拟大batch效果。2. 减小输入图像尺寸如从224降到192。3. 使用更小的骨干网络如ResNet18或减少Transformer的embed_dim和num_layers。混淆矩阵显示某一类如肺炎召回率极低1.类别严重不平衡数据集中肺炎样本远少于正常样本。2.分类阈值0.5不适合当前任务。1. 使用加权损失函数给样本少的类别更高的权重。BCEWithLogitsLoss可以通过pos_weight参数实现。或在数据加载时使用加权随机采样WeightedRandomSampler。2.调整决策阈值。不一定要用0.5可以根据验证集上的PR曲线或F1分数选择一个能平衡精确率和召回率的最佳阈值。推理速度慢1. 模型参数量大计算复杂。2. 未使用GPU或数据加载慢。1. 考虑模型轻量化知识蒸馏、剪枝、量化PyTorch支持INT8量化。2. 确保推理时model.eval()和torch.no_grad()。使用torch.jit.trace或torch.jit.script将模型转换为TorchScript可能提升效率。对于生产环境可考虑使用ONNX导出并用TensorRT等推理引擎加速。注意力可视化图看起来是噪声没有聚焦到病灶1. 模型可能没有学到有意义的注意力。2. 注意力提取或聚合的方式不对。3. 需要更专业的可视化方法如Grad-CAM for Transformer。1. 首先确认模型分类性能是否达标。如果模型本身准确率低注意力图无意义是正常的。2. 检查提取的注意力权重是哪个头、哪一层的。通常最后一层或多头注意力的平均图更有意义。尝试对多个头的注意力图进行平均。3. 研究针对Transformer的可解释性工具如BertViz适配视觉任务需修改或Captum库中的集成梯度等方法。7.3 项目扩展与优化方向这个项目是一个强大的基线但还有很大的优化和扩展空间更先进的架构尝试纯Transformer模型如Swin Transformer它在多个视觉任务上超越了ResNet。或者使用更高效的混合架构如ConvNeXt或MobileViT。更丰富的数据集与任务当前是二分类正常/肺炎。可以扩展到多分类区分细菌性肺炎、病毒性肺炎、结核等。也可以引入更大型、更多样的胸部X光数据集如CheXpert或MIMIC-CXR。处理类别不平衡医疗数据中正常样本往往远多于患病样本。深入研究Focal Loss、代价敏感学习Cost-Sensitive Learning或过采样/欠采样技术如SMOTE。不确定性估计对于AI辅助诊断知道模型“有多不确定”和知道诊断结果同样重要。可以引入蒙特卡洛DropoutMC Dropout或深度集成Deep Ensembles来估计预测的不确定性。模型可解释性除了注意力图可以系统性地使用LIME、SHAP或Grad-CAM系列工具生成热力图向医生展示模型判断所依据的影像区域这对于临床采纳至关重要。部署优化使用PyTorch Mobile或ONNX Runtime将模型部署到移动设备或边缘设备如便携式X光机实现离线、低延迟的诊断辅助。这个基于Transformer和ResNet34的肺炎诊断系统从架构设计到训练调优再到评估分析涵盖了一个完整深度学习项目的核心流程。其中关于混合架构的权衡、超参数设置的考量、以及针对医疗场景的评估重点如召回率和混淆矩阵都是跨领域项目通用的宝贵经验。代码和思路已经相当完整你可以直接用它作为起点替换自己的数据集或者按照上述优化方向进行探索相信能做出更有意思也更有价值的东西。本文还有配套的精品资源点击获取