医学图像分类这个方向,这两年最大的变化不是模型翻新有多快,而是“怎么把一篇经典论文读进工程代码里”成了真正的分水岭。很多人拿着 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 这种较浅的网络。
- Bottleneck:1x1 卷积降维、3x3 卷积、1x1 卷积升维,适合 ResNet50/101/152 这种深层网络。
Bottleneck 的核心是降低计算量。比如输入是 256 维,如果直接做两个 3x3 卷积,计算量很大;先用 1x1 降到 64 维,做完 3x3 再升回 256 维,参数量和计算量明显下降。
论文还强调,shortcut 连接在输入输出维度一致时不需要额外参数,维度不一致时有两种处理方法:一种是补零,另一种是用 1x1 卷积投影。在工程实现里,torchvision 的 ResNet50 选择的是 1x1 卷积投影,步长为 2 的时候还会在 shortcut 里加上 stride=2 的下采样。
不少新手容易忽略的一个细节是: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 python=3.10 -y conda activate medical_resnetPython 版本以实际环境兼容为准。如果机器上没有 conda,也可以直接用 venv 或系统 Python,但建议统一使用虚拟环境,避免不同项目依赖互相干扰。
4.2 安装 PyTorch
PyTorch 的安装命令根据 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(p=0.5), transforms.RandomRotation(10), transforms.ColorJitter(brightness=0.2, contrast=0.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(root='data/train', transform=train_transforms) val_dataset = datasets.ImageFolder(root='data/val', transform=val_transforms) test_dataset = datasets.ImageFolder(root='data/test', transform=val_transforms) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4) 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_classes=2, pretrained=True): model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1 if pretrained else None) # 获取全连接层的输入维度 in_features = model.fc.in_features # 替换最后一层为二分类 model.fc = nn.Linear(in_features, num_classes) return model注意:如果使用旧版 torchvision,weights 参数写法可能是 pretrained=True,新版中推荐用 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_classes=2, pretrained=True) criterion = nn.CrossEntropyLoss(weight=class_weights.to('cuda' if torch.cuda.is_available() else 'cpu')) optimizer = optim.Adam(model.parameters(), lr=1e-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, desc='Training'): 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, desc='Evaluating'): 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_classes=2, pretrained=True).to(device) criterion = nn.CrossEntropyLoss(weight=get_class_weights(train_loader.dataset).to(device)) optimizer = torch.optim.Adam(model.parameters(), lr=1e-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(f'Epoch {epoch+1}/{num_epochs}') print(f'Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}') print(f'Val 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(f'Saved 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_classes=2, pretrained=False) model.load_state_dict(torch.load('best_resnet50.pth', map_location=device)) 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, dim=1) pred = torch.argmax(output, dim=1).item() return pred, prob.squeeze().cpu().numpy() pred, prob = predict_image('data/test/normal/example.jpg') print(f'预测类别: {pred}, 各类别概率: {prob}')这里的模型在加载权重时没有传入预训练权重,因为权重已经以文件形式保存,不需要再次从网上下载。如果在新环境中没有 GPU,torch.load 需要指定 map_location='cpu',否则会报 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_classes=2, pretrained=False).to(device) model.load_state_dict(torch.load('best_resnet50.pth', map_location=device)) 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() 返回 False | PyTorch 版本与 CUDA 驱动不匹配 | 运行 nvidia-smi 查看驱动支持的 CUDA 版本;打印 torch.version.cuda | 安装与驱动匹配的 PyTorch 版本,或更新显卡驱动 |
| 训练损失不下降 | 学习率过大或过小、数据预处理错误、标签噪声 | 输出第一个 batch 的损失值;尝试将学习率调低 10 倍;检查图像标准化参数 | 使用学习率预热或余弦退火;修正数据预处理 |
| 验证准确率虚高但测试准确率骤降 | 数据划分时未按患者隔离 | 检查 train/val 目录中是否出现同一患者的多张切片 | 按患者 ID 重新划分数据 |
| 类别不均衡导致模型偏向多数类 | 使用交叉熵损失但未设置权重 | 打印各类别样本数;查看混淆矩阵中少数类召回率 | 使用加权交叉熵损失;调整分类阈值 |
| 加载模型时报错缺少 state_dict 键 | 保存的是完整模型而不是 state_dict,或者模型类别数不一致 | 打印保存文件内容,检查 model.fc.out_features | 统一使用 model.state_dict() 保存和加载 |
| 显存不足 OOM | batch 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 是一个起点,但也是理解现代卷积网络的最佳教材。把它的残差思想读懂,很多后续的网络改进读起来都会顺畅很多。希望这篇“论文带读 + 代码复现”能帮你少踩一些坑,建议收藏备用,在项目需要时按步骤实操一遍。