资讯动态

ResNet + PyTorch:医学图像分类从论文到代码的完整实践

发布时间:2026/9/7 14:21:03 来源:尧图企业网站定制
医学图像分类这个方向这两年最大的变化不是模型翻新有多快而是“怎么把一篇经典论文读进工程代码里”成了真正的分水岭。很多人拿着 ResNet 训练脚本就跑却发现医学数据集上准确率虚高、迁移学习权重不适用、验证集划分还泄露了患者信息。这些坑论文里不会直接写教程里也常常一句带过。这篇文章的选择很简单以 ResNet 论文为线索以 PyTorch 为工具完整走一遍医学图像分类从数据处理、模型搭建、训练验证到常见排查的流程。不是只贴代码也不只是复述论文而是把两者对照起来告诉你论文里哪些设计至今管用哪些地方在医学场景下需要调整。读完之后你既能看懂 ResNet 的残差结构为什么能解决网络退化问题也能拿到一个可以直接跑的医学图像二分类示例并且知道下一步该往哪个方向深入。1. 这篇文章真正要解决的问题如果你正在做医学图像相关的项目大概率遇到过下面几类问题数据集只有几千张甚至几百张图用 ResNet 从头训练验证集准确率始终上不去。直接用 ImageNet 预训练权重做迁移学习发现效果时好时坏还没法解释原因。把同一个患者的若干张切片随机分进训练集和验证集导致验证准确率虚高模型真实泛化能力被严重高估。遇到类别不均衡模型把多数类全猜对了整体准确率很高但少数类几乎全部漏检。这些问题和 ResNet 的结构本身关系不大更多是“论文思路”和“工程落地”之间的落差。ResNet 论文解决的核心问题是深度网络难以训练但在医学图像场景里你还要额外处理小样本、标注噪声、类不均衡和数据划分的伦理与有效性。这篇文章就是站在“论文带读 代码复现”的角度把 ResNet 的核心原理和 PyTorch 实现放到医学图像分类的真实约束下重新讲一遍。读完后你会有一个清晰的判断ResNet 依然是医学图像分类里最值得优先尝试的骨干网络之一但指望“直接套模型”就能解决临床问题是不现实的。2. ResNet 论文核心思想带读残差学习解决了什么问题ResNet 论文的题目是 Deep Residual Learning for Image Recognition发表于 2015 年。它的出发点是当卷积网络深度不断增加时训练误差反而会上升。这不是过拟合因为训练误差本身就变高了。论文把这个现象叫做“退化问题”并指出深层网络难以通过恒等映射来保持性能。通俗解释一个 20 层的网络理论上至少不应该比 10 层网络差因为前 10 层可以学一样的东西后 10 层可以学成“什么都不做”。但实际训练时深层网络很难让后面那些层学会“什么都不做”梯度在反向传播中也更容易出问题。于是论文提出了残差学习。假设网络某一层希望拟合的潜在映射是 H(x)残差块不直接让这一层去学 H(x)而是去学 F(x) H(x) - x最后的输出是 F(x) x。这个过程用公式表达很简单普通网络输出 H(x)残差网络输出 F(x) x这里的 x 会通过一个 shortcut connection快捷连接直接传到后面的层。如果 F(x) 趋近于 0输出就约等于 x网络就能轻易地学习“恒等映射”。这相当于给深层网络的训练兜了一条底梯度也能通过 shortcut 更顺畅地回传。论文还设计了两种残差块BasicBlock两个 3x3 卷积适合 ResNet18/34 这种较浅的网络。Bottleneck1x1 卷积降维、3x3 卷积、1x1 卷积升维适合 ResNet50/101/152 这种深层网络。Bottleneck 的核心是降低计算量。比如输入是 256 维如果直接做两个 3x3 卷积计算量很大先用 1x1 降到 64 维做完 3x3 再升回 256 维参数量和计算量明显下降。论文还强调shortcut 连接在输入输出维度一致时不需要额外参数维度不一致时有两种处理方法一种是补零另一种是用 1x1 卷积投影。在工程实现里torchvision 的 ResNet50 选择的是 1x1 卷积投影步长为 2 的时候还会在 shortcut 里加上 stride2 的下采样。不少新手容易忽略的一个细节是ResNet 的潜力来自“深度”但深度只有配合预训练和足够数据才有意义。在 ImageNet 上ResNet152 比 ResNet34 有明显优势但在医学小数据集上ResNet34 未必输给 ResNet152甚至可能更好训练。3. 医学图像分类的数据特点与处理策略医学图像分类和自然图像分类的最大区别不是模型而是数据。如果不对数据有清醒认知后面所有代码都可能跑出一个“看起来不错但实际没用”的模型。3.1 小样本与迁移学习医学数据集通常比较小因为标注需要专业医生成本很高。几千张图已经算是中等规模很多公开数据集只有几百到一两千张。在这种规模下从零训练 ResNet50 很容易过拟合。常见做法是使用 ImageNet 预训练权重做迁移学习并在训练时根据数据量决定是否冻结前几层。如果数据量特别少比如每类只有一两百张可以考虑使用预训练权重只训练最后的全连接层。做更强的数据增强。使用较小的网络比如 ResNet18。如果有条件用医学领域预训练模型而不是 ImageNet 预训练模型。3.2 类别不均衡医学图像里阳性样本往往远少于阴性样本比如罕见病筛查、病灶区域分类。此时直接使用准确率作为指标很容易产生误导。假设 95% 是阴性、5% 是阳性模型全预测阴性也能有 95% 的准确率看起来很好实际上毫无临床价值。处理方式包括使用加权交叉熵损失给少数类更高的权重。在评估时关注召回率、特异度、F1-score、AUC。调整分类阈值而不是默认使用 0.5。3.3 数据泄露问题这是医学图像分类里容易被忽视、后果却很严重的坑。一个患者可能有多张图像比如多张肺窗 CT 切片、多张病理视野图。如果这些图像被随机分到训练集和验证集模型实际上是在记忆患者特征而不是学习疾病特征。验证集准确率可能非常高换个患者队列就崩了。正确做法是按患者 ID 划分数据保证同一个患者的图像全部属于同一个集合。通常建议按照“训练集、验证集、测试集”三层划分测试集只用来做最终评估不参与任何调参。3.4 数据增强策略医学图像增强需要谨慎。常见的几何增强如随机旋转、水平翻转是安全的但也要结合任务判断。比如手部 X 光片左右翻转就没有太大问题但病理图像中如果组织结构有方向性翻转可能引入不合理样本。除此之外还可以考虑对比度调整、亮度调整、随机裁剪这些对泛化能力有实际帮助。4. 环境准备PyTorch 安装与依赖库在开始代码之前先准备好运行环境。本文的代码基于 PyTorch推荐使用 Anaconda 管理 Python 环境。不同操作系统下安装方式略有区别下面给出通用流程。4.1 创建虚拟环境conda create -n medical_resnet python3.10 -y conda activate medical_resnetPython 版本以实际环境兼容为准。如果机器上没有 conda也可以直接用 venv 或系统 Python但建议统一使用虚拟环境避免不同项目依赖互相干扰。4.2 安装 PyTorchPyTorch 的安装命令根据 CUDA 版本不同而变化。CPU 环境用于测试没有问题但训练 ResNet50 建议使用 GPU。安装前先查看显卡驱动支持的 CUDA 版本然后到 PyTorch 官网选择对应命令。CPU 版本pip install torch torchvisionGPU 版本具体 CUDA 版本号以官网为准pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121安装完成后验证python -c import torch; print(torch.__version__); print(torch.cuda.is_available())如果 torch.cuda.is_available() 返回 False优先检查安装的 PyTorch 版本与 CUDA 驱动是否匹配而不是怀疑显卡坏了。4.3 安装其他依赖除了 torch还需要用到 torchvision、Pillow、numpy、matplotlib、scikit-learn 等库。pip install pillow numpy matplotlib scikit-learntorchvision 会随 PyTorch 一起安装如果不确定可以单独指定版本。版本兼容性以实际环境为准重点是 torch 和 torchvision 的大版本尽量对应否则可能报算子不匹配的错误。5. 基于 ResNet50 的医学图像分类代码实现下面以“肺部 X 光图像二分类”为例演示完整的训练流程。任务目标是把图像分为“正常”和“肺炎”两个类别。数据集目录结构如下data/ ├── train/ │ ├── normal/ │ └── pneumonia/ ├── val/ │ ├── normal/ │ └── pneumonia/ └── test/ ├── normal/ └── pneumonia/这里特意讲一下如果数据是同一个患者的多个文件请务必先按患者 ID 划分好目录再进入下面的训练流程不要把所有图像混在一起随机划分。5.1 数据加载与增强PyTorch 推荐使用 torchvision.datasets.ImageFolder 读取这种目录结构。我们先用 transforms 定义训练集和验证集的数据处理方式。# 文件路径data_loader.py import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 训练集增强 train_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(10), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) # 验证集和测试集不做随机增强 val_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) train_dataset datasets.ImageFolder(rootdata/train, transformtrain_transforms) val_dataset datasets.ImageFolder(rootdata/val, transformval_transforms) test_dataset datasets.ImageFolder(rootdata/test, transformval_transforms) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse, num_workers4) print(类别映射:, train_dataset.class_to_idx) print(训练集图像数:, len(train_dataset)) print(验证集图像数:, len(val_dataset)) print(测试集图像数:, len(test_dataset))这里使用了 ImageNet 数据集的均值和标准差做标准化。因为我们要加载 ImageNet 预训练权重所以必须沿用它的标准化参数否则预训练模型的特征分布会被打乱。5.2 搭建 ResNet50 模型torchvision 提供了现成的 resnet50 模型我们只需要把最后一层全连接改成二分类输出。# 文件路径model.py import torch.nn as nn from torchvision import models def get_resnet50(num_classes2, pretrainedTrue): model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1 if pretrained else None) # 获取全连接层的输入维度 in_features model.fc.in_features # 替换最后一层为二分类 model.fc nn.Linear(in_features, num_classes) return model注意如果使用旧版 torchvisionweights 参数写法可能是 pretrainedTrue新版中推荐用 weights 枚举方式避免未来的兼容性问题。实际以你安装的版本为准。如果需要冻结前面层可以这样设置# 冻结前 5 个块 for param in list(model.parameters())[:-3]: param.requires_grad False是否冻结层没有绝对标准。数据量越少冻结越多数据量足够大可以全部微调。常见策略是先冻结观察训练效果再逐步解冻更多层做微调。5.3 定义损失函数和优化器二分类任务使用交叉熵损失即可。由于存在类别不均衡问题这里演示如何计算类别权重并传给损失函数。# 文件路径train_utils.py import torch import torch.nn as nn import torch.optim as optim def get_class_weights(dataset): 根据样本数量计算类别权重样本越少权重越高 targets dataset.targets class_counts torch.bincount(torch.tensor(targets)) total sum(class_counts) class_weights total / (len(class_counts) * class_counts.float()) return class_weights class_weights get_class_weights(train_dataset) print(类别权重:, class_weights) model get_resnet50(num_classes2, pretrainedTrue) criterion nn.CrossEntropyLoss(weightclass_weights.to(cuda if torch.cuda.is_available() else cpu)) optimizer optim.Adam(model.parameters(), lr1e-4)使用类别权重的效果是少数类损失被放大模型会更重视少数类。不过权重也不能设得太大否则会导致多数类大量误判。具体权重可以结合验证集结果调整。5.4 训练循环训练循环由四个主要部分构成前向传播、计算损失、反向传播、参数更新。下面给出一个精简版本。# 文件路径train.py import torch import torch.nn as nn from tqdm import tqdm def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for inputs, labels in tqdm(dataloader, descTraining): inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * inputs.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc correct / total return epoch_loss, epoch_acc def evaluate(model, dataloader, criterion, device): model.eval() running_loss 0.0 correct 0 total 0 with torch.no_grad(): for inputs, labels in tqdm(dataloader, descEvaluating): inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) loss criterion(outputs, labels) running_loss loss.item() * inputs.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() epoch_loss running_loss / total epoch_acc correct / total return epoch_loss, epoch_acc训练主函数# 文件路径main.py import torch from data_loader import train_loader, val_loader from model import get_resnet50 from train_utils import get_class_weights from train import train_one_epoch, evaluate device torch.device(cuda if torch.cuda.is_available() else cpu) model get_resnet50(num_classes2, pretrainedTrue).to(device) criterion nn.CrossEntropyLoss(weightget_class_weights(train_loader.dataset).to(device)) optimizer torch.optim.Adam(model.parameters(), lr1e-4) num_epochs 20 best_val_acc 0.0 for epoch in range(num_epochs): train_loss, train_acc train_one_epoch( model, train_loader, criterion, optimizer, device ) val_loss, val_acc evaluate(model, val_loader, criterion, device) print(fEpoch {epoch1}/{num_epochs}) print(fTrain Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}) print(fVal Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}) # 保存验证集上表现最好的模型 if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_resnet50.pth) print(fSaved best model with val acc: {val_acc:.4f})这段代码里用了一个非常简单的模型选择策略只保存验证集准确率最高的权重。实际项目中建议同时关注验证集 F1 或 AUC尤其当类别不均衡时准确率最高并不代表模型最优。5.5 测试与预测训练完成后我们需要在测试集上评估最终模型并写一个单图预测函数。# 文件路径predict.py import torch from PIL import Image from torchvision import transforms from model import get_resnet50 device torch.device(cuda if torch.cuda.is_available() else cpu) model get_resnet50(num_classes2, pretrainedFalse) model.load_state_dict(torch.load(best_resnet50.pth, map_locationdevice)) model.to(device) model.eval() # 保持和训练一致的预处理 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) def predict_image(image_path): image Image.open(image_path).convert(RGB) image_tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): output model(image_tensor) prob torch.softmax(output, dim1) pred torch.argmax(output, dim1).item() return pred, prob.squeeze().cpu().numpy() pred, prob predict_image(data/test/normal/example.jpg) print(f预测类别: {pred}, 各类别概率: {prob})这里的模型在加载权重时没有传入预训练权重因为权重已经以文件形式保存不需要再次从网上下载。如果在新环境中没有 GPUtorch.load 需要指定 map_locationcpu否则会报 CUDA 不可用的错误。6. 运行结果与效果验证在训练脚本中每轮训练会打印当前 epoch 的训练损失、训练准确率、验证损失、验证准确率。你需要重点观察以下信号训练损失是否持续下降。如果训练损失不降先检查数据预处理是否合理学习率是否过大或过小。验证损失是否在某个点开始上升。如果训练损失下降、验证损失上升说明过拟合已经出现。验证准确率是否稳定。如果忽高忽低可能是 batch size 太小或学习率太高。一个常见的健康训练曲线应该是前几个 epoch 训练损失快速下降验证准确率同步提升随后提升速度变慢最后进入平台区。运行完成后最终模型保存在 best_resnet50.pth 中。接下来在测试集上计算最终指标才算是任务完成。# 文件路径evaluate_test.py from data_loader import test_loader import torch from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix device torch.device(cuda if torch.cuda.is_available() else cpu) model get_resnet50(num_classes2, pretrainedFalse).to(device) model.load_state_dict(torch.load(best_resnet50.pth, map_locationdevice)) model.eval() all_preds [] all_labels [] with torch.no_grad(): for inputs, labels in test_loader: inputs, labels inputs.to(device), labels.to(device) outputs model(inputs) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) print(Accuracy:, accuracy_score(all_labels, all_preds)) print(Precision:, precision_score(all_labels, all_preds)) print(Recall:, recall_score(all_labels, all_preds)) print(F1:, f1_score(all_labels, all_preds)) print(Confusion Matrix:) print(confusion_matrix(all_labels, all_preds))判断模型是否真正可用不能只凭 Accuracy 一个指标。在医学场景里Recall 往往更关键因为漏诊一个阳性病例的代价通常高于误诊一个阴性病例。如果 Accuracy 很高但 Recall 很低说明模型在少数类上表现很差需要针对性调整。如果测试结果不理想优先检查数据划分是否真的按患者级别隔离。数据增强是否过于激进导致训练集和真实测试分布不一致。类别权重是否设置合理是否过度惩罚了多数类。预训练权重的 ImageNet 统计量和你的医学图像分布差距是否太大需要更多微调轮数。7. 常见问题与排查思路问题现象可能原因排查方式解决方案torch.cuda.is_available() 返回 FalsePyTorch 版本与 CUDA 驱动不匹配运行 nvidia-smi 查看驱动支持的 CUDA 版本打印 torch.version.cuda安装与驱动匹配的 PyTorch 版本或更新显卡驱动训练损失不下降学习率过大或过小、数据预处理错误、标签噪声输出第一个 batch 的损失值尝试将学习率调低 10 倍检查图像标准化参数使用学习率预热或余弦退火修正数据预处理验证准确率虚高但测试准确率骤降数据划分时未按患者隔离检查 train/val 目录中是否出现同一患者的多张切片按患者 ID 重新划分数据类别不均衡导致模型偏向多数类使用交叉熵损失但未设置权重打印各类别样本数查看混淆矩阵中少数类召回率使用加权交叉熵损失调整分类阈值加载模型时报错缺少 state_dict 键保存的是完整模型而不是 state_dict或者模型类别数不一致打印保存文件内容检查 model.fc.out_features统一使用 model.state_dict() 保存和加载显存不足 OOMbatch size 太大或图像分辨率太高降低 batch size查看单张图像占用显存使用梯度累积模拟更大 batch size使用更小的输入尺寸排查问题时有个基本原则先看数据再看模型结构最后看训练策略。很多看似是模型的问题最后都出在数据划分或预处理上。8. 最佳实践与工程建议下面这些建议来自医学图像分类项目中比较通用的经验虽然不会直接出现在论文里但对真实项目稳定性影响很大。8.1 按患者级别划分数据再次强调这一点因为它太重要了。如果数据集包含多个患者且同一患者有多张图像必须把同一患者的图像全部放进同一个集合。否则模型会学习“这个患者看起来有点像训练集中的样子”而不是“这个病变的特征是什么”。这是医学图像训练中最常见的隐性错误。8.2 使用分层五折交叉验证数据量不大时单次划分的验证结果波动很大。更可靠的做法是做分层五折交叉验证把数据按患者划分成 5 份每次用 4 份训练、1 份验证最终取 5 次的平均指标。这样做的好处是更能反映模型在不同子集上的稳定性也减少了随机划分带来的偶然性。8.3 保存完整训练日志与配置代码跑通不算什么能复现才是工程价值。建议训练时记录以下信息数据集划分方式与患者 ID 映射。随机种子。数据增强策略。学习率、batch size、优化器参数。每次 epoch 的训练损失、验证损失、验证指标。模型权重文件的保存路径和评价指标。这些日志可以用 CSV 或 JSON 保存在复现和调试时能节省大量时间。8.4 重视可解释性验证医学图像模型的临床可信度离不开可解释性分析。常见的做法是在测试集的代表性样本上做 Grad-CAM 热力图观察模型关注的是病灶区域还是无关背景。如果模型重点关注的位置与医生判断区域不一致即使准确率很高也需要谨慎使用。PyTorch 生态中已有一些库可以快速生成 Grad-CAM比如 pytorch-grad-cam。如果在生产环境使用还需要让医生参与评估热力图的合理性而不是只看数值指标。8.5 安全与合规提醒医学图像涉及患者隐私和伦理问题训练前必须确认数据来源合规去除个人身份信息并确保实验行为符合医疗机构或数据提供方要求。涉及真实临床环境部署时还需要额外的模型验证、临床评估和相关审批流程。不建议把未经充分验证的模型直接用于诊断或治疗决策。8.6 生产环境部署前检查如果把训练好的模型部署到线上建议在部署前增加以下检查在独立的测试集上复算指标确认结果稳定。保存模型时同步保存预处理参数和类别映射。设计输入校验逻辑拒绝损坏图像或尺寸异常图像。在灰度发布阶段设置人工复核环节对比模型预测与医生判断。9. 总结与后续学习方向这篇文章做的事情可以概括为三件带读了 ResNet 论文的核心思想解释了医学图像数据的特点与常见陷阱给出了一个基于 PyTorch ResNet50 的完整医学图像分类代码流程。下一步的学习路径可以这样安排如果你还不太熟悉 PyTorch 基础先跑通本文的代码理解 Dataset、DataLoader、模型定义、训练循环、评估流程这几块内容。如果你想提升分类效果可以往两个方向深入一是数据增强和数据策略二是模型改进。后者包括使用更轻量的 EfficientNet、引入注意力机制、或者尝试 Vision Transformer 这类结构。如果你关心可解释性建议学习 Grad-CAM 的实现原理和工具用法这对医学图像项目尤其重要。如果你关注产线落地可以继续学习 ONNX 模型导出、推理服务部署和灰度验证方案。ResNet 是一个起点但也是理解现代卷积网络的最佳教材。把它的残差思想读懂很多后续的网络改进读起来都会顺畅很多。希望这篇“论文带读 代码复现”能帮你少踩一些坑建议收藏备用在项目需要时按步骤实操一遍。

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

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

免费获取报价