用PyTorch训练ResNet:从数据准备到调参避坑全流程
手头攒了一批自己的图片数据想用深度模型做个分类又不想从零搭网络这时候 PyTorch 加 ResNet 基本是最稳的组合。可是真到了“用自己的数据训练 ResNet”这个环节你会发现网上教程大多停在跑通官方 demo一旦换成自己的数据立刻冒出各种奇怪问题环境装不上、数据格式不对、loss 不降、显存爆掉、验证集准确率上不去。这篇博文就围绕这个场景把我实际踩过的坑和一套能直接照做的流程完整写出来。这篇文章适合两类人一类是刚学完 PyTorch 基础、想拿自己的数据练手的学生或转行者另一类是在实际工作中需要快速做图像分类基线模型的工程师。我会从环境搭建开始一路讲到数据组织、模型改造、训练调参、问题排查尽量让你照着操作就能跑通少走弯路。1. 环境准备与基础选型先把坑填平再动手1.1 用 Anaconda 隔离一套干净的 PyTorch 环境很多人第一步就栽在环境上。我自己见过太多案例电脑里 Python 版本乱七八糟装过的 torch 互相冲突最后连import torch都报错。所以我建议不管在 Windows 还是 Linux 上都先用 Anaconda 建一个独立环境把 PyTorch 和系统其它 Python 包隔离干净。conda create -n pytorch python3.10 conda activate pytorchPython 版本我推荐 3.9 或 3.10兼容性好老代码新代码都能跑。建完环境后再装 PyTorch不要直接用pip install torch因为默认源装的是 CPU 版速度慢不说后续训练只要数据集稍微大一点就完全跑不动。正确的做法是去 PyTorch 官网的安装页根据自己的 CUDA 版本复制对应的安装命令。如果你不确定 CUDA 版本在命令行执行nvidia-smi看右上角的 CUDA Version那个数字表示你显卡驱动支持的最高 CUDA 版本。按这个数字去官网选对应版本的 PyTorch 就行。这里注意一个容易混淆的点驱动支持的 CUDA 版本是向下兼容的也就是说驱动是 CUDA 12.0你装 PyTorch 配 CUDA 11.8 也完全没问题所以没必要追求版本号完全一致只要 PyTorch 要求的 CUDA 不高于驱动支持的版本就行。CPU 版的 PyTorch 我一般不推荐除非你只有笔记本集成显卡或者纯粹想练语法。但即使是练语法我也建议在 GPU 环境里写因为后续做真实验证时CPU 训练 ResNet 的速度会让人怀疑人生。实测用 CPU 跑一个几百张图片的小数据集一个 epoch 都要几分钟而 GPU 几秒钟就完事。1.2 预训练权重ResNet 微调的关键前提环境装好之后还有一个前置资源很容易出问题预训练权重。我们用自己的数据训练 ResNet绝大多数情况不是从零开始随机初始化而是加载在 ImageNet 上预训练好的权重然后在自己的数据上微调。这背后的逻辑很朴素ImageNet 上有 1000 类、上千万张图片模型在这么大数据集上学到的底层特征边缘、纹理、形状对绝大多数图像任务都是通用的。你的数据量再小也能站在这些通用特征的基础上快速收敛。PyTorch 里加载预训练权重很简单我用得最多的是这一行import torchvision.models as models model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT)这里有个经验之谈weights参数是 PyTorch 新版 API 的写法老版本里用的是pretrainedTrue。如果你在网上查到旧代码运行时报参数错误别慌多半就是 API 版本问题。新 API 里DEFAULT等价于IMAGENET1K_V1意思就是加载 ImageNet 上预训练好的默认权重。下载权重的时候还经常遇到一个问题从官网下载太慢或者超时。这种情况建议先手动下载权重文件配置好环境变量TORCH_HOME把权重文件放到对应目录下然后设置环境变量离线加载。还有一种更省心的方式用torch.hub.load_state_dict_from_url下载时如果失败会有明显的报错提示直接把报错信息里的 URL 复制到浏览器里手动下载再把文件放到C:\Users\你的用户名\.cache\torch\hub\checkpoints或 Linux 下的对应目录就能避过下载超时。这个坑几乎每个人都会踩一次。2. 数据组织与预处理自己的数据要先“懂事”2.1 目录结构直接决定代码复杂度用自己的数据训练模型第一件事不是写代码而是把数据整理成模型能直接读取的格式。PyTorch 里最省事的方式是使用torchvision.datasets.ImageFolder它要求数据目录按“类别文件夹”的方式组织data/ train/ dog/ dog_001.jpg dog_002.jpg cat/ cat_001.jpg bird/ bird_001.jpg val/ dog/ cat/ bird/这个结构的含义是train下面每个文件夹的名字就是类别名label文件夹里的图片就是该类别的样本。ImageFolder会自动扫描所有子文件夹按文件夹名的字母顺序分配类别索引。这个顺序很重要因为后续计算准确率、输出混淆矩阵时你要知道索引 0 对应的是哪个类别不然结果出来了对着索引一脸懵。如果你手头的数据是一堆散乱的文件名或者存在 Excel 表格的标注里那就需要先写一个数据整理脚本。这里我给一个我常用的简单划分脚本思路遍历所有图片按比例随机划分到 train 和 val用shutil.copy复制或移动文件到目标目录。划分比例我习惯 8:2 或 9:1。这里强调一点验证集一定要独立出来绝对不能和训练集混在一起否则后面评估的准确率是虚高的模型泛化能力会被高估。数据量小到几十张甚至十几张每类时可以先把训练和验证合并用交叉验证或者留一法来评估不然验证集太小准确率波动会非常剧烈。这种情况通常还有一个更好的选择直接用预训练模型做特征提取不训练整个网络这里到第 3 部分再展开说。2.2 数据增强与标准化参数数据整理好了加载的时候还不能直接原始图片输入网络。ResNet 这类网络对输入有固定要求尺寸要统一、像素值范围要标准化。PyTorch 里通常在torchvision.transforms中完成这些操作。我个人推荐的训练集 transform 组合如下from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])解释一下每个操作的作用。RandomResizedCrop(224)会随机裁剪图片再缩放到 224x224相当于让模型看到同一张图的不同局部和不同尺度这能显著提升模型对目标位置和尺寸变化的鲁棒性。RandomHorizontalFlip是随机水平翻转对大多数自然图像任务都有效但要注意如果你的任务跟方向强相关比如识别文字方向、区分左右手这个增强就不能用。ColorJitter是随机调整亮度、对比度、饱和度能增强模型对光照变化的适应性。验证集不用做数据增强但也不能直接Resize(224)就完事。标准做法是先缩放到 256再中心裁剪 224这样能保留图片中间区域的信息比直接拉伸变形效果好得多。Normalize 里的 mean 和 std 数值是 ImageNet 数据集的统计值因为预训练模型就是在这些统计值下训练的所以微调时我们也要用同样的数值对图片做标准化否则输入分布不一致预训练特征的优势就发挥不出来。这里再强调一个细节ToTensor()必须在Normalize()之前因为它会把 PIL 图片转成张量并把像素值从 0-255 缩放到 0-1 区间。如果你顺序写反了会得到一堆奇怪的值训练基本起不来。2.3 类别不均衡的处理策略自己的数据很少像教科书那样完美均衡。最常见的场景是某些类别图片特别多某些类别特别少。如果你直接拿这种数据训练模型会倾向于把所有样本都预测成数量多的类别整体准确率看着不错但少数类的 recall 几乎为零。处理类别不均衡我按推荐程度排个序最简单粗暴如果少数类样本实在少得可怜比如少于 20 张直接考虑去掉这个类别或者改用更小的任务目标。加权损失函数给CrossEntropyLoss传weight参数少数类权重设大多数类权重设小实现简单效果好。过采样在 DataLoader 里用WeightedRandomSampler让少数类被抽中的概率更高。这个方式我也常用效果比加权损失更直接。代码示意from torch.utils.data import WeightedRandomSampler sample_weights [] for cls_idx in dataset.targets: sample_weights.append(1.0 / class_sample_count[cls_idx]) sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue) dataloader DataLoader(dataset, batch_size32, samplersampler)class_sample_count就是每个类别的样本数量。replacementTrue表示允许重复采样这样每个 epoch 里多数类样本会被跳过一部分少数类样本会被重复抽取整体来看每个类别的曝光次数更均衡。3. 模型搭建怎么把 ResNet 改成你自己的分类器3.1 ResNet 的基本结构和平替选择ResNet 的核心思想是残差连接简单说就是让网络在学习恒等映射时很容易——如果某个卷积层没有用处模型可以通过残差连接跨过它从而解决深层网络退化的问题。这也是 ResNet 能在 2015 年之后成为视觉任务主力骨架的根本原因。实操中你只需要在 ResNet 系列里做选择。我用的最多的是 ResNet18 和 ResNet50ResNet18层数少、参数少、速度快适合数据量不大、计算资源有限的情况。我自己做工业项目时经常先用 ResNet18 跑通基线。ResNet50层数深、表达能力更强但需要更多数据和更长训练时间。如果你的数据量在上万张级别ResNet50 能带来明显的准确率提升。ResNet101 和更深的版本在小数据集上反而容易过拟合而且训练很慢一般不建议作为首选。如果你的输入不是 224x224而是更高分辨率比如 384 或 512可以在加载模型时设置model models.resnet18()并在构造后修改model.conv1的stride和padding或者直接调整输入尺寸。不过这是进阶玩法新手阶段建议就老老实实用 224x224。3.2 修改分类头与冻结策略加载预训练 ResNet 之后最后一层全连接层是 1000 个输出对应 ImageNet 的 1000 类。你的任务如果是 10 类、20 类就要把这个全连接层替换掉import torch.nn as nn num_classes 10 # 换成你自己的类别数 model.fc nn.Linear(model.fc.in_features, num_classes)这一行是最核心的改动。model.fc.in_features会自动读取原来全连接层的输入维度ResNet18 是 512ResNet50 是 2048不用自己硬编码方便以后换模型。接下来是微调策略的选择。这里有个关键问题要不要冻结前面的层只训练最后的分类头我的建议是分场景数据量少每类几百张以下冻结骨干网络只训练新加的 fc 层。因为数据少微调全部参数很容易过拟合。数据量够每类上千张或者总量上万全部参数一起微调但可以给骨干网络和分类头设置不同学习率。冻结骨干网络的代码for param in model.parameters(): param.requires_grad False for param in model.fc.parameters(): param.requires_grad True这里一个常见误区是只改fc层结构之后如果你要全量微调什么都不用做模型所有参数默认都是可训练的。只有当你想要冻结部分层时才需要手动设置requires_grad。3.3 完整训练脚本关键代码把数据加载和模型改造拼接起来最核心的训练骨架大概是下面这样。注意我这里刻意省去了日志和保存的细节那些放到第 4 部分说先让训练能跑起来。from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder from torchvision import models import torch.nn as nn import torch.optim as optim train_dataset ImageFolder(data/train, transformtrain_transform) val_dataset ImageFolder(data/val, transformval_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4) model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT) model.fc nn.Linear(model.fc.in_features, len(train_dataset.classes)) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) for epoch in range(10): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) print(fEpoch {epoch1}, Loss: {running_loss / len(train_dataset):.4f})这段代码有几个细节值得单独拎出来讲。ImageFolder加载后会有一个classes属性它是一个类别名列表排序跟模型训练时的 label 索引一一对应这个列表在后续做推理时一定要保存下来。model.train()和model.eval()的状态切换也不能忘因为train()模式下 BatchNorm 层会更新均值和方差验证和推理时要用固定统计量。新手最容易犯的错就是验证时忘了切eval()模式导致结果不稳定。optimizer.zero_grad()这行也容易漏。PyTorch 默认会累积梯度如果不每次清零梯度会不断累加loss 曲线会非常奇怪。我见过好几个初学者在这里卡了半天最后就是少了这一行。4. 训练配置与调参让 Loss 真正降下去4.1 损失函数和优化器选择分类任务的损失函数基本就是CrossEntropyLoss它内部把LogSoftmax和NLLLoss合在了一起所以你不需要在模型输出层再手动加 softmax可以直接把网络原始的 logits 传进去PyTorch 会自动计算。优化器这一块我个人的路线是小数据集上微调用 Adam追求更高精度时换 SGD。Adam 的优点是无脑、收敛快、对学习率不敏感适合快速迭代验证想法。SGD 加 Momentum 和 Weight Decay 在训练后期往往能得到更好的泛化效果这是被大量实验验证过的经验——Adam 虽然收敛快但有时会停在比较尖锐的极小值点SGD 则更平滑。optimizer optim.Adam(model.parameters(), lr1e-3) # 换成 SGD 时我常用的配置 optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay1e-4)关于初始学习率我的经验值全量微调用 1e-4 到 1e-3 之间只训练新的 fc 层可以用 1e-3 到 1e-2。如果你用的是 SGD初始学习率建议从 0.01 起步再根据 loss 曲线调整。Batch size 影响学习率的选择batch size 越大学习率可以相应调大因为梯度估计更稳定。显存不够时先减 batch size不要硬扛常见取值 16、32、64超内存就减半。4.2 学习率策略与早停训练深度模型学习率策略比很多人想象的更重要。固定学习率训练到底往往初期挺快后期 loss 就卡住不动了。我常用的两个策略ReduceLROnPlateau当模型在验证集上的指标连续若干个 epoch 不提升时自动把学习率降低一个量级。CosineAnnealingLR学习率按照余弦曲线从初始值逐渐衰减到接近 0。适合训练轮数固定的场景。scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.1, patience3)配合早停Early Stopping是更完整的方案设定一个 patience 值比如连续 5 个 epoch 验证集准确率没有刷新纪录就停止训练并恢复最优模型权重。这里要注意保存的应该是验证集上效果最好的模型而不是最后一个 epoch 的模型。训练后期的模型往往已经开始过拟合直接拿来推理效果反而不如中间的好。4.3 验证与模型保存验证过程的核心代码和训练类似但要加上torch.no_grad()和model.eval()def evaluate(model, val_loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return correct / totaltorch.max(outputs, 1)是取每个样本在类别维度上的最大概率索引也就是预测的类别。为什么不先把 logits 过 softmax 再取最大值因为softmax是单调函数不会改变 argmax 的位置直接对 logits 取最大值就行省一次运算。模型保存我建议只保存权重不保存整个模型对象torch.save(model.state_dict(), best_model.pth)加载的时候需要先实例化一个同样结构的模型再load_state_dict。如果你连模型结构也想一起保存也可以用torch.save(model, full_model.pth)但这种方式对代码版本和训练脚本文件的依赖很重换个环境或者改了代码就容易被版本兼容问题卡住。我工作中通常还会把优化器状态、epoch、类别名列表一并保存成一个字典这样中断训练后可以从断点继续torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), class_names: train_dataset.classes, }, checkpoint.pth)5. 高频问题排查直接把报错抄走的速查表5.1 WinError 1114 这类环境报错热词里反复出现OSError: [WinError 1114] 动态链接库(DLL)初始化例程失败这个报错我见过太多次了。它表面上是c10.dll加载失败实际原因通常有几种第一种是缺少 Visual C 运行库。PyTorch 编译出来的 DLL 依赖系统里的 VC 运行库如果你电脑上没装或者版本太旧就会出现这个错误。解决办法是去微软官网下载并安装最新的 Visual C Redistributable装完重启就能解决。第二种是 torch 和 torchvision 版本不匹配或者安装的是 CPU 版本和 GPU 版本混在一起。这里的典型场景是先装了 CPU 版后来又pip install torchvision时自动装了一个不匹配的依赖版本导致 DLL 冲突。解决方案很简单把当前环境里的 torch 和 torchvision 全部卸载干净pip uninstall torch torchvision torchaudio -y然后严格按照官网命令重新安装装完用几行代码验证python -c import torch; print(torch.__version__, torch.cuda.is_available())第三种可能是 conda 环境本身混乱某些依赖库被升级或降级导致不兼容。这种情况最省事的办法是直接删掉重建环境conda deactivate conda remove -n pytorch --all conda create -n pytorch python3.10再走一遍安装流程。这种重装大法看起来粗暴但实测解决问题的效率最高比在一个坏掉的环境里反复调试 DLL 依赖快得多。5.2 训练过程中的典型问题训练时最容易遇到的是显存溢出Out of Memory, OOM。解决办法优先级减小 batch size - 减小输入图片尺寸 - 换更小的模型ResNet18 换掉 ResNet50- 检查是不是验证集也把梯度存了进去。还有一个经常被忽略的点PyTorch 默认在反向传播后不会自动释放计算图所以每个 step 里一定要调用optimizer.zero_grad()否则梯度累积会导致显存暴涨。另一个高频问题loss 一直不降。如果 loss 停留在接近类别数的对数值附近比如 10 类分类loss 卡在 2.3 左右这通常是模型根本没学到东西。先检查数据预处理是否正确尤其是ToTensor和Normalize的顺序再检查标签是否对齐ImageFolder的类别索引是否和真实类别对应。还有一个快速验证方法用几十张图片过拟合到 100% 准确率如果这个都做不到说明模型或数据代码有问题如果做到了说明整体流程没问题只是训练策略需要调整。验证集准确率上不去、训练集却很高这是典型的过拟合。处理手段按优先级排序增加数据增强强度 - 加 dropout 或 weight decay - 减少模型层数 - 冻结更多层。我用得最多的是增强数据增强和调 weight decay双管齐下。5.3 推理部署时的细节模型训练完真正放到业务里做推理时还有几个容易忽略的细节。第一个是类别顺序。ImageFolder在训练时按文件夹名的字母顺序生成类别索引比如文件夹名是dog和cat那么索引 0 是cat索引 1 是dog。推理时如果不做映射直接输出索引 1 就当成 dog会出大问题。所以我每次拿到训练好的模型都会顺手把class_names存一份推理时严格按照同一份索引解码。第二个是输入预处理要跟训练时保持一致。推理时用的 transform 必须跟验证集一致Resize(256) - CenterCrop(224) - ToTensor - Normalize。很多人训练时好好的部署时忘了 Normalize或者随手用了个别的尺寸准确率直接崩掉。第三个是模型切换 eval 模式。推理时务必执行model.eval()否则 BatchNorm 层在误开 train 模式的情况下可能会用当前批数据的统计量导致输出不稳定。我见过一个真实的线上故障就是有人忘了eval()同一个输入两次推理结果不一样查了半天才发现是 BatchNorm 在作怪。最后再分享一个我在实际项目里的体会用自己的数据训练 ResNet最靠谱的路径不是一上来就追求高精度而是先用 ResNet18 加少量 epoch 快速跑通整个流程确认数据、代码、评估都没问题之后再逐步加数据增强、换更大的模型、精细调学习率。这个过程看起来绕了一圈实际上比直接怼一个大模型然后被各种环境问题和数据问题困住要快得多。后续如果你想继续扩展还可以试试把 ResNet 换成 EfficientNet 或者 Vision Transformer 做对比实验训练代码几乎不用改但你会发现不同模型的自有数据适配情况差别还挺大这个对比本身就很有价值。