1. 框架整体设计与目录结构
先讲讲为什么我最终会沉淀出这么一套"基于 PyTorch 的图像分类完整训练框架"。早些年做图像分类项目,基本是今天写一个脚本训 ResNet,明天复制一个脚本调 DenseNet,后天又在另一个文件夹里堆一个 EfficientNet。前几次还好,等实验数量一多,问题就来了:模型代码、数据增强、训练参数、评估逻辑全部搅在一起,换一个数据集要改七八处地方,想复现三天前的实验得靠运气。后来下决心把训练流程从具体模型里抽出来,搞成一套"模型无关、配置驱动"的训练框架,只要关注模型结构本身和数据路径,其他事情比如学习率调度、断点续训、日志记录、模型保存,全部由框架统一处理。这篇文章分享的,就是这套框架从零搭建的完整思路和可参考代码。
适合谁看?如果你准备做深度学习图像分类的入门实战,或者正在被一堆零散脚本搞得头大,又或者想改造自己的训练代码但没想清楚怎么拆模块,这套东西可以直接抄作业。我不会贴一个巨大的完整工程,而是把每个模块的设计原因、关键代码、踩坑记录都讲明白,你照着拼起来就能用。
1.1 需求分析:训练脚本到底在解决什么问题
在写任何代码之前,先想清楚一个图像分类训练脚本由哪些基本动作组成。拆开来看,任何训练流程都绕不开这么几件事:加载数据、定义模型、计算损失、反向传播、更新参数、定时评估、保存权重。这套流程是固定的,会变的只是具体的数据集路径、模型种类、超参数值。所以框架的核心思路,就是把"固定流程"和"可变配置"彻底分离。
固定流程沉淀成代码,也就是 train.py 里的训练循环;可变配置收敛到一个配置文件里,包括数据集路径、图片尺寸、batch size、初始学习率、训练轮数、优化器类型、模型名称。这样每次开新实验,只需要复制一份配置文件改改参数就行,训练主流程一行都不用动。
这种设计还有一个隐藏好处:当你的训练逻辑有 bug 时,只改 train.py 就能让所有历史实验受益;而当你的模型效果不好时,只调 config 就能快速对比多组超参。职责单一,排查问题也快不少。
1.2 完整目录结构:从 config 到 checkpoint
这套框架的最终目录结构如下,我实际项目里就是这么组织的:
project/ ├── configs/ │ ├── __init__.py │ └── resnet18_cifar10.py ├── data/ │ ├── __init__.py │ ├── dataset.py │ └── transforms.py ├── models/ │ ├── __init__.py │ └── classifier.py ├── utils/ │ ├── __init__.py │ ├── logger.py │ ├── lr_scheduler.py │ └── checkpoint.py ├── checkpoints/ ├── logs/ ├── train.py ├── infer.py └── requirements.txt各模块的职责很清晰:configs 放所有实验配置;data 目录放数据集封装和数据增强;models 目录放模型定义;utils 放日志、学习率、断点保存这些横切工具;checkpoints 和 logs 是运行时自动生成的目录,分别存模型权重和训练日志。train.py 是入口脚本,infer.py 是推理脚本。
一个容易忽略的点是每个目录下的__init__.py,很多人写小脚本时省掉它,导致后面 import 路径一团乱麻。建议从第一天就把每个目录都当成包来组织,后面改起来会舒服很多。
2. 环境搭建与依赖选择:PyTorch 基础框架的安装坑
这套框架最底层的东西,就是 PyTorch 本身。环境搭不好,后面所有代码都跑不起来。这一节我结合自己的经验,把安装过程中最常见的问题一次性讲清楚。
2.1 PyTorch 版本与 CUDA 匹配:先认清自己的显卡
安装 PyTorch 之前,先搞清楚你到底需要 GPU 版还是 CPU 版。如果你只是想先跑通代码、或者显卡是核显级别的,直接装 CPU 版完全够用,代码一行都不用改,PyTorch 会自动在 CPU 上执行。但如果你要训练真实的图像分类模型,尤其是 ResNet、EfficientNet 这类深度网络,建议还是用 GPU。
GPU 版本这里有个最容易踩的坑:CUDA 版本不匹配。很多人的习惯是去显卡驱动面板看版本号,然后照着装 PyTorch,结果装完torch.cuda.is_available()返回 False。原因是 PyTorch 要求的不是"显卡驱动版本",而是"CUDA 运行时版本"。你的驱动版本只要不低于某个门槛,就能支持 PyTorch 内置的 CUDA 运行时,不需要单独安装完整的 CUDA Toolkit。
判断方法很简单,命令行执行:
nvidia-smi看右上角的 "CUDA Version",比如显示 12.4。然后在 PyTorch 官网选一个 CUDA 版本号不大于 12.4 的安装命令,比如:
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118这里的 cu118 代表 CUDA 11.8 运行时。选低一点的版本完全没问题,PyTorch 会自带对应的 CUDA 依赖库。
补充一个经验,别盲目追新版本。PyTorch 官方 2024 年后热门趋势已经明显偏向于稳定渠道,我实测下来 cu118 或者 cu121 这类装机量大的版本兼容性最好,网上踩坑资料也最多。新版本刚发布时,经常会遇到某个配套库(比如 torchvision)还没跟上的情况。
2.2 下载慢的终极解法:国内镜像源与安装后自检
安装 PyTorch 时最折磨人的就是下载速度。特别是用默认的官方源拉取几个 GB 的安装包时,速度经常只有几十 KB/s,挂一晚上都未必能装完。有网友说"手机开了热点下载依然很慢",其实根因不是网络波动,而是 PyTorch 官方 CDN 在部分区域就是慢。解决方案是换国内镜像源。
以 pip 为例,推荐用清华源或者阿里源:
pip install torch torchvision -i https://pypi.tuna.tsinghua.edu.cn/simple但这里有个细节必须提醒:如果你需要指定 CUDA 版本,最好的方式是先从官方源下载 whl 文件,再本地安装。不过实际操作中我见过很多人直接对官方 CUDA 版命令加-i参数,结果 pip 还是走了默认源,因为 PyTorch 官方--index-url的优先级高于-i,两者会互相干扰。所以稳妥的做法是:先把 whl 文件下载到本地(可以用浏览器或 wget),然后pip install xxx.whl本地安装。
如果是用 conda 管理环境,也一样可以配置国内 conda 镜像:
conda config --add channels https://mirrors.tuna.tsinghua.edu.cn/anaconda/cloud/pytorch/ conda config --set show_channel_urls yes装完后别急着写代码,先做一个三行自检:
import torch print(torch.__version__) print(torch.cuda.is_available())如果返回 True,说明 GPU 版本已经正常工作。如果返回 False,优先检查 PyTorch 版本和 CUDA 版本是否匹配。
2.3 虚拟环境管理:为什么建议每个项目单独建环境
很多新手在图省事,直接把 PyTorch 装进 base 环境,所有项目共用一套包。前几个月没事,等做第二个项目时,发现 A 项目需要 PyTorch 2.0,B 项目还在用 1.12,版本一冲突,整个环境直接废掉。我的建议是每个项目一个独立 conda 环境,反正创建环境的成本很低:
conda create -n classify python=3.8 conda activate classify pip install torch torchvision -i https://pypi.tuna.tsinghua.edu.cn/simple选 Python 版本时也不用纠结 3.8 还是 3.10,PyTorch 对 Python 版本的兼容性一直很好。除非你后续要接一些老旧的第三方库,否则 Python 3.8 以上都可以。
3. 数据管线与图像预处理:训练框架的地基
数据读取是整个训练流程里最容易拖慢速度、又最容易被忽视的环节。很多人模型写得挺规范,数据加载却用最原始的 for 循环一张张读,训练速度直接掉一个量级。这一节讲清楚 PyTorch 数据管线的正确打开方式。
3.1 自定义 Dataset:从文件夹到样本对
图像分类任务最常见的数据组织方式是:训练集和验证集各有一个文件夹,里面按类别分子文件夹。这种情况下,PyTorch 自带的torchvision.datasets.ImageFolder可以直接用,不用自己写 Dataset。
但真实项目中,数据集往往没有这么规整:有的是 CSV 文件标注图片路径和标签,有的图片存在多个目录里需要过滤,有的还需要做样本均衡。这时候就需要自己写一个 Dataset 类。模板如下:
import torch from torch.utils.data import Dataset from PIL import Image import os class ImageClassificationDataset(Dataset): def __init__(self, image_paths, labels, transform=None): self.image_paths = image_paths self.labels = labels self.transform = transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image_path = self.image_paths[idx] image = Image.open(image_path).convert("RGB") label = self.labels[idx] if self.transform is not None: image = self.transform(image) return image, label几个关键细节:
- 用 PIL 而不是 cv2 读取图片,因为
torchvision.transforms的输入类型是 PIL Image,用 PIL 省去类型转换。 convert("RGB")一定要加,很多灰度图或 RGBA 图不统一,不转换后面会报通道数错误。- 不要在
__getitem__里做复杂的预处理,比如重 Resize 大图,会拖慢数据加载。
3.2 数据增强策略:训练集和验证集的区别对待
图像分类场景下,数据增强是提升模型泛化能力性价比最高的手段。经典的组合是:随机裁剪加缩放(RandomResizedCrop)、随机水平翻转(RandomHorizontalFlip)、颜色抖动(ColorJitter)、归一化。写到代码里:
from torchvision import transforms train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.08, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.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]) ])注意训练集和验证集的增强策略是必须不同的。训练集要多样性,做随机扰动;验证集要确定性,只做 resize 和 center crop,保证每次评估结果可复现。前几年有一个热门讨论说"为什么训练深度神经网络这么困难,问题可能不在梯度消失而在于退化",这确实和增强策略的设计息息相关,粗暴的增强会让模型在训练集上的 loss 居高不下,从而看起来像"梯度消失了"。
归一化的 mean 和 std 直接采用 ImageNet 的统计值即可,这是一套被验证过有效的默认参数,不需要自己算。但如果你的数据集和 ImageNet 差异很大(比如医学影像),后面的实验优化方向可以考虑重新计算数据集的均值和方差。
3.3 DataLoader 参数细节:num_workers 与 pin_memory 的真相
数据加载在 GPU 训练时是最容易成为瓶颈的环节。在torch.utils.data.DataLoader里,有这么几个参数值得重点关注:
num_workers:决定用几个子进程预取数据。把这个值设成 0 会在主进程里同步加载数据,GPU 每算一个 batch 就要等数据读完,训练速度慢得离谱。一般设成 CPU 核心数的一半左右,比如 8 核 CPU,设 4 或 8 都行。pin_memory:设成 True,把数据固定在锁页内存里,GPU 拷贝数据时会快不少。这个参数在 CPU 训练时没有意义,但在 GPU 训练时几乎是白捡的性能提升。drop_last:当数据集大小不能被 batch size 整除时,最后一个 batch 会很小,某些 BN 层的统计值会受到影响。训练集建议设成 True,把不完整的 batch 丢掉;验证集设成 False。
标准写法示例:
train_loader = DataLoader( train_dataset, batch_size=64, shuffle=True, num_workers=4, pin_memory=True, drop_last=True ) val_loader = DataLoader( val_dataset, batch_size=64, shuffle=False, num_workers=4, pin_memory=True )验证集不要 shuffle,因为评估时不需要随机性。如果验证集特别大,也可以把 batch size 调大一点,反正不需要反向传播,显存占用更小。
4. 模型构建与训练循环核心实现:从 ResNet18 到自定义分类头
有了数据和环境,接下来进入核心代码部分。这一节把 model 定义、训练循环、验证循环、模型保存整个链路完整过一遍。
4.1 模型初始化:PyTorch 基础框架下的分类头替换
图像分类最常用的套路是:用 ImageNet 上预训练的骨干网络做特征提取,替换最后一层全连接,让它输出自己数据集的类别数。使用torchvision.models可以非常方便地完成:
import torch import torch.nn as nn from torchvision import models def build_model(num_classes=10, model_name="resnet18", pretrained=True): if model_name == "resnet18": model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1 if pretrained else None) in_features = model.fc.in_features model.fc = nn.Linear(in_features, num_classes) elif model_name == "resnet50": model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2 if pretrained else None) in_features = model.fc.in_features model.fc = nn.Linear(in_features, num_classes) else: raise ValueError(f"Unknown model: {model_name}") return model这里有两个经验点:
第一,torchvision新版本里pretrained=True的写法已经被弃用,推荐用weights=models.ResNet18_Weights.IMAGENET1K_V1,虽然代码长一点,但更明确,而且不容易遇到版本升级后的警告或报错。
第二,替换全连接层时,先通过model.fc.in_features拿到原始输入维度,而不是硬编码成 512 或 2048。因为不同模型的 fc 层输入维度不一样,硬编码换模型时必踩坑。
4.2 训练循环详解:为什么要 zero_grad、为什么 loss.item()
训练循环的骨架如下:
def train_one_epoch(model, train_loader, criterion, optimizer, device, epoch): model.train() total_loss = 0.0 correct = 0 total = 0 for batch_idx, (images, labels) in enumerate(train_loader): images = images.to(device) labels = labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() if batch_idx % 50 == 0: print(f"Epoch [{epoch}], Batch [{batch_idx}], Loss: {loss.item():.4f}") avg_loss = total_loss / total accuracy = 100.0 * correct / total return avg_loss, accuracy为什么每次 backward 前要执行optimizer.zero_grad()?因为 PyTorch 的梯度默认是累加的。如果你不手动清零,下一次backward()会把新算出的梯度加到旧梯度上,导致参数更新方向完全错乱。这是新手最常犯的错误之一。
loss.item()的用法也值得说明。loss是一个包含梯度信息的张量,如果直接total_loss += loss,会导致计算图一直被保留,显存越占越多,最后 OOM。.item()把标量从计算图里取出来,变成普通 Python 数字,既省显存又方便打印。
4.3 验证循环与模型保存:只在验证集上做决策
训练集上的准确率没有参考价值,因为模型本来就在拟合这些数据。真正的决策依据是验证集上表现。验证循环不计算梯度,用torch.no_grad()包起来,节省显存和计算资源:
@torch.no_grad() def validate(model, val_loader, criterion, device): model.eval() total_loss = 0.0 correct = 0 total = 0 for images, labels in val_loader: images = images.to(device) labels = labels.to(device) outputs = model(images) loss = criterion(outputs, labels) total_loss += loss.item() * images.size(0) _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() avg_loss = total_loss / total accuracy = 100.0 * correct / total return avg_loss, accuracy模型保存这里我有一个推荐做法:不只是保存 epoch 结束后的模型,而是保存验证集准确率最高的一次,这样即使后面过拟合了,也能找到最好的那个权重。每次验证完,如果 acc 比历史最高还高,就覆盖保存,这就是常说的"best model"。
def save_checkpoint(state, filename): torch.save(state, filename) # 训练循环里 best_acc = 0.0 for epoch in range(epochs): train_loss, train_acc = train_one_epoch(...) val_loss, val_acc = validate(...) if val_acc > best_acc: best_acc = val_acc save_checkpoint({ "epoch": epoch, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "val_acc": val_acc, }, f"checkpoints/best_model.pth")只有模型状态没有优化器状态时,加载后能推理但不能继续训练;想断点续训,一定要连优化器的 state_dict 一起保存。
4.4 学习率调度:为什么手动衰减不是好主意
学习率是训练过程中最敏感的超参数。固定学习率从头训到尾,前期下降慢,后期又容易来回震荡。更合理的做法是先用较大的学习率快速下降,训练到中后期再把学习率调小,让损失在局部最小值附近继续精调。
PyTorch 提供了多个现成的调度器,我最常用的是ReduceLROnPlateau和CosineAnnealingLR。前者是看指标,验证集 loss 连续 N 个 epoch 不下降就衰减;后者无脑按余弦曲线衰减,不用管指标,省心。
from torch.optim.lr_scheduler import ReduceLROnPlateau optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9, weight_decay=5e-4) scheduler = ReduceLROnPlateau(optimizer, mode="min", factor=0.5, patience=5) # 每个 epoch 之后 scheduler.step(val_loss)如果你的优化器选了 Adam,学习率初始值建议从小到大试,常见范围 1e-3 到 3e-4;如果使用 SGD+momentum,初始值一般 0.01 到 0.1 之间。
还有一个很实用的技术是"冻结部分模型"。当你的预训练模型要在小数据集上做迁移学习时,前面几层学到的是基础纹理、边缘特征,这些特征非常通用,不需要在目标数据集上重新学习。可以先冻结 backbone,只训练新换的分类头,等分类头收敛后再解冻全部层微调。实现方式:
for param in model.parameters(): param.requires_grad = False for param in model.fc.parameters(): param.requires_grad = True不过 torchvision 的优化器会默认更新所有 requires_grad=True 的参数,所以冻结后,优化器自然只更新 fc 层。
5. 训练日志与断点续训:让实验不再"失忆"
训练一个完整模型动辄几个小时甚至几天,如果不记录日志、不支持断点续训,一次意外断电机就能让所有工作白费。这一节把训练过程中容易被忽略的工程化细节讲清楚。
5.1 用 TensorBoard 还是自定义日志
训练过程可视化,最早的方案是 TensorBoard,虽然它源自 TensorFlow,但 PyTorch 的torch.utils.tensorboard可以直接调用。后来又流行起来了 W&B(Weights & Biases),可视化能力强,还能在网页上对比多次实验。两者怎么选?
我的选择标准是:个人开发或者公司内网实验,优先 TensorBoard,免费、无需联网、够用;需要团队协作、大量跑实验对比,才考虑 W&B。
TensorBoard 的基础用法:
from torch.utils.tensorboard import SummaryWriter writer = SummaryWriter("logs/experiment_01") writer.add_scalar("Loss/train", train_loss, epoch) writer.add_scalar("Loss/val", val_loss, epoch) writer.add_scalar("Acc/val", val_acc, epoch) writer.close()运行:
tensorboard --logdir logs浏览器打开http://localhost:6006就能看到训练曲线。如果同时训练多个模型,就把 log 放到不同子目录下,TensorBoard 会自动叠加对比,用起来很舒服。
5.2 日志记录实现:print 的替代方案
用 print 打印训练信息不是不行,但问题很明显:输出被终端缓冲区截断、无法同时输出到文件和屏幕、信息太乱没法按等级过滤。我的做法是直接用 Python 自带的logging模块,训练脚本开头统一配置一下:
import logging logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s", handlers=[ logging.FileHandler(f"logs/train_{timestamp}.log", encoding="utf-8"), logging.StreamHandler() ] ) logger = logging.getLogger(__name__) logger.info(f"Epoch [{epoch}/{epochs}] train_loss: {train_loss:.4f}, val_acc: {val_acc:.2f}%")这样一条日志同时进文件和终端,训练完翻日志也比较方便,尤其是程序崩溃时,可以从日志里看到最后一步做了什么。
5.3 断点续训:从 .pth 加载模型和优化器
训练到一半因为各种原因中断,是家常便饭。断电、显存不够被 kill、甚至手滑关掉终端,都可能导致训练中断。如果从头开始训,等于浪费之前所有算力。断点续训的实现其实很简单,因为 4.3 节保存 checkpoint 时已经把模型参数、优化器参数、epoch 都存进去了,加载时反过来恢复就行:
checkpoint = torch.load("checkpoints/last_model.pth", map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"]) optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) start_epoch = checkpoint["epoch"] + 1加载优化器 state_dict 后,继续之前的学习率调度状态也最好恢复。如果你手动设置了scheduler,同样保存scheduler_state_dict并在加载后调用scheduler.load_state_dict(...)。
关于模型加载还一个常见问题:只保存了model_state_dict的模型文件,加载时如果模型定义里num_classes和之前不一样,会报维度不匹配。解决办法是检查state_dict里最后一层 fc 的out_features大小,或者干脆不要加载最后两层,如下:
pretrained_dict = torch.load("model.pth") model_dict = model.state_dict() pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict and v.shape == model_dict[k].shape} model_dict.update(pretrained_dict) model.load_state_dict(model_dict)这种"只加载形状匹配的层"的技巧,在做迁移学习或者微调开源权重时非常实用。
6. 单卡训练全流程跑通:以 CIFAR-10 为例
前面几节把模块拆开讲了,这一节把完整的训练流程串起来,从命令行入口到最终保存模型,给出一份可以直接跑通的训练脚本参考。我用 CIFAR-10 作为示例数据集,因为 torchvision 自带下载,零成本复现。
6.1 train.py 主函数:从 config 到 checkpoint 的完整串联
整体代码结构如下,我把主流程拆成几个函数,方便理解每一步在做什么:
import argparse import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms def load_config(config_path): import importlib.util spec = importlib.util.spec_from_file_location("config", config_path) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) return module.config def build_data(config): # 参考第 3 节的数据加载部分 transform_train = ... transform_val = ... train_dataset = datasets.CIFAR10(root=config["data_root"], train=True, download=True, transform=transform_train) val_dataset = datasets.CIFAR10(root=config["data_root"], train=False, download=True, transform=transform_val) train_loader = DataLoader(train_dataset, batch_size=config["batch_size"], shuffle=True, num_workers=config["num_workers"], pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=config["batch_size"], shuffle=False, num_workers=config["num_workers"], pin_memory=True) return train_loader, val_loader def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", type=str, default="configs/resnet18_cifar10.py") args = parser.parse_args() config = load_config(args.config) device = torch.device(config["device"] if torch.cuda.is_available() else "cpu") train_loader, val_loader = build_data(config) model = build_model(num_classes=config["num_classes"], model_name=config["model_name"]) model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.SGD(model.parameters(), lr=config["lr"], momentum=0.9, weight_decay=5e-4) best_acc = 0.0 for epoch in range(1, config["epochs"] + 1): train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device, epoch) val_loss, val_acc = validate(model, val_loader, criterion, device) print(f"Epoch {epoch}/{config['epochs']}, Train Loss: {train_loss:.4f}, " f"Train Acc: {train_acc:.2f}%, Val Acc: {val_acc:.2f}%") if val_acc > best_acc: best_acc = val_acc torch.save(model.state_dict(), "checkpoints/best_model.pth") if __name__ == "__main__": main()配置文件configs/resnet18_cifar10.py长这样:
config = { "data_root": "./data", "num_classes": 10, "model_name": "resnet18", "batch_size": 64, "epochs": 50, "lr": 0.01, "num_workers": 4, "device": "cuda", }6.2 训练效果评估:loss 和 acc 怎么读
CIFAR-10 上,用 ResNet18 预训练模型+SGD,50 个 epoch 的正常结果大概是:训练集 acc 90% 以上,验证集 acc 85% 到 90% 之间。验证集 acc 和训练集 acc 的差距控制在 5 个百分点以内,基本可以接受;差距超过 10 个百分点,就要反思是不是过拟合了。
训练开始的前几个 epoch,loss 不降反升是正常的。因为预训练模型一开始在 ImageNet 的特征空间,换到 CIFAR-10 的分类头需要适应新数据分布,先让 loss 震荡一两轮再说。如果 5 个 epoch 后 loss 依然纹丝不动,那才需要担心。
6.3 推理脚本 infer.py:加载模型并预测单张图片
训练完模型后,通常写一个简单的推理脚本,加载权重并对单张图片预测:
import torch from PIL import Image from torchvision import transforms def predict(image_path, model, class_names, device="cuda"): 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]) ]) image = Image.open(image_path).convert("RGB") image_tensor = transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): outputs = model(image_tensor) probs = torch.softmax(outputs, dim=1) top_prob, top_class = probs.topk(1, dim=1) return class_names[top_class.item()], top_prob.item()注意推理时不要忘了torch.no_grad(),model.eval()也很重要,它会关闭 dropout 和 BN 层的训练行为,让推理结果的随机性降为零。这个坑我真踩过:忘记eval(),同一个模型跑两次推理结果都不一样。
7. 常见问题与排查技巧实录:训练框架的避坑指南
框架写好后,真正跑起来时会遇到各种奇奇怪怪的问题。这一节把我在实际使用中遇到的高频问题整理成速查表,每个都是真实经历。
7.1 环境与安装类问题
Q1:PyTorch 装好了,但 torch.cuda.is_available() 返回 False
排查顺序:先nvidia-smi看驱动是否正常,再看驱动 CUDA version 是否 >= 你安装的 PyTorch CUDA 版本。如果驱动正常,确定安装的是 GPU 版而不是 CPU 版。很多人用 pip 换源时不小心装成了 CPU 版,因为 PyTorch CPU 版的包名不带 cuXX 后缀,遇到这个问题重装一次 GPU 版即可。
Q2:官方源下载太慢怎么办
换国内镜像源下载纯 CPU 版最省心;GPU 版建议先获取官方 whl 的直链,下载到本地后再安装。不建议在官方命令后面直接加-i,因为--index-url会覆盖-i的配置,导致镜像失效。
Q3:conda 创建虚拟环境时卡在 Solving environment
conda 在处理包依赖时经常很慢,这种情况建议换用mamba或者直接用 pip+venv。特别是 PyTorch 这类依赖数很多的包,pip 的解析速度通常比 conda 快很多。如果坚持用 conda,先把 conda 镜像和 pip 镜像都配好,能显著缩短时间。
7.2 训练过程类问题
这部分是整个框架的核心痛点,我单独展开讲。
Q4:显存 OOM(Out of Memory)
图片尺寸越大、batch size 越大、模型越深,显存占用越高。真的 OOM 了,最直接的解法是减小 batch size,一次别喂那么多图。如果调小 batch size 后精度掉得厉害,可以试试梯度累积,思路是先攒几个 batch 的梯度再更新一次参数,效果上相当于大 batch,显存占用却不变:
accumulation_steps = 4 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): outputs = model(images) loss = criterion(outputs, labels) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()另一个容易忽略的点是,验证阶段也要记得torch.no_grad(),否则验证循环同样会构建计算图,显存峰值会翻倍。
Q5:训练 loss 不下降
先确认数据加载没问题:打印几个 batch 看看图片和标签是否对应。然后确认模型是否train()模式,有些新手在循环外调了model.eval()忘了调回来,BN 和 dropout 全部失效,模型根本学不进去。再查学习率是否合适:学习率太大 loss 会震荡,太小 loss 龟速下降。建议用 1e-3 作为初始值快速验证一次,再根据现象调整。
Q6:训练集 acc 高、验证集 acc 低,过拟合了怎么办
过拟合在小数据集上特别常见。优先级排序:先加数据增强,再看是否需要加 dropout 或 weight_decay,最后考虑换更小的模型。数据增强是最温和的手段,几乎所有图像分类任务都能从中受益。weight_decay 一般从 1e-4 到 5e-4 之间取值,调大一些能显著抑制过拟合,代价是训练集 acc 也会降一点,这个取舍是正常的。
Q7:不同 epoch 的结果波动大
如果验证集 acc 忽高忽低,大概率是验证集太小,评估结果受随机性影响大。解决办法是增大验证集,或者把验证集评估多跑几次取平均。还有一种可能是学习率太大,后期在局部最优附近震荡,调小学习率或者换余弦退火调度器能缓解。
Q8:从 .pth / .pt / .bin 加载模型时维度不匹配
这个在迁移学习场景里几乎一定会碰到。建议用 5.3 节的"只加载形状匹配的层"方案,其实更省心的做法是在保存模型时就把num_classes记到配置里,加载前先确认类别数一致。如果是第三方权重文件格式比较特殊(比如某些项目保存成.bin或.pth.tar),先用torch.load打印一下 state_dict 的 keys 和 shapes,快速判断里面存的是什么结构,再决定怎么加载。
7.3 数据加载类问题
Q9:训练时 GPU 利用率为 0,CPU 快跑满
典型的瓶颈在数据加载。优先把num_workers调大;如果改了还不行,检查是否把图片直接放在机械硬盘上,这种场景 IO 会拖死训练,把数据提前复制到 SSD 或内存里能显著提速。另外transforms里如果有大量 CPU 预处理,尽量简化,把缩放、裁剪这类操作放到 GPU 上做不现实,但可以减少重复计算。
Q10:图片读取时出现损坏文件
真实数据集里混入个别损坏图片很常见,PIL.Image.open会直接抛异常导致训练中断。在 Dataset 里加上异常保护和降级策略,用白名单或者过滤掉打不开的图片:
def __getitem__(self, idx): for _ in range(10): try: image_path = self.image_paths[idx] image = Image.open(image_path).convert("RGB") break except Exception: idx = (idx + 1) % len(self.image_paths) ...这个方案简单粗暴,能保证训练不中断,但对特别脏的数据集来说,还是建议先离线清洗一遍再开训。
8. 从单卡到多卡:PyTorch 训练框架的常见扩展方向
框架搭起来后,下一步自然是想着怎么训得更快、更稳。这一节简单聊聊几个常见的扩展方向,以及我个人实际用下来的感受。
8.1 混合精度训练:白捡的性能提升
如果你用的是 Volta 及其之后的 NVIDIA 显卡(包括 Turing、Ampere、Ada Lovelace 架构),GPU 里都有专门的 Tensor Core 单元,PyTorch 提供了torch.cuda.amp模块可以实现自动混合精度训练。核心代码改动很小,只需要在训练循环里加一个 GradScaler:
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for images, labels in train_loader: images = images.to(device) labels = labels.to(device) with autocast(): outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度训练在保持精度几乎不变的前提下,通常能把训练速度提升 1.5 到 2 倍,显存占用也能降 30% 左右。这个优化不改变模型结构和数据流,接入成本很低,值得作为框架默认选项。
8.2 多卡训练:DataParallel 与 DistributedDataParallel 的选择
如果你的机器有多张显卡,自然会想到用多卡加速。PyTorch 提供了两种方式,nn.DataParallel和nn.DistributedDataParallel。前者只需要一行代码model = nn.DataParallel(model),但没有线程安全问题,多卡通信效率也低;后者配置复杂一些,但性能明显更好,是官方推荐的方式。
关于分布式训练,我给的建议很直接:单机多卡用 DDP,无脑上DistributedDataParallel。如果你只需要在单卡场景跑实验,干脆先别上多卡,把单卡流程跑到极致后再考虑,否则环境配置带来的额外复杂度只会消耗热情。
以下是单机多卡 DDP 的极简模板:
import torch.distributed as dist import torch.multiprocessing as mp def train_worker(rank, world_size): dist.init_process_group("nccl", rank=rank, world_size=world_size) torch.cuda.set_device(rank) model = model.to(rank) model = nn.parallel.DistributedDataParallel(model, device_ids=[rank]) ... if __name__ == "__main__": world_size = torch.cuda.device_count() mp.spawn(train_worker, args=(world_size,), nprocs=world_size)8.3 实验管理:多次实验结果的对比与追溯
训练框架稳定之后,最值得投入的反而不是代码本身,而是实验管理机制。我见过太多人跑了十几组实验,最后根本分不清哪组用了什么参数。我的做法是:每组实验用单独的时间戳目录名,把配置、日志、最优模型的 checkpoint 全部放在一起:
runs/ ├── 20250112_0930_resnet18_b64_lr001/ │ ├── config.py │ ├── train.log │ ├── best_model.pth │ └── events.out.tfevents... ├── 20250113_1030_efficientnet_b32_lr0003/ │ ├── config.py │ ├── train.log │ ├── best_model.pth │ └── events.out.tfevents...这样每个实验目录内聚,TensorBoard 也能直接指向 runs 目录对比多组实验。遇到效果好的实验,直接复制整个目录就能复现,不需要从记忆里拼凑信息。
9. 写在最后:训练框架的设计心得
这套基于 PyTorch 的图像分类完整训练框架,前前后后被我迭代了很多版本,从最初的一个 train.py 到现在 config 驱动、模型与流程分离、带日志和断点续训的工程,中间踩过的坑都写在上面了。回头来看,整个设计里最重要的不是某个具体的技巧,而是"把固定流程和可变配置分离"这个原则。只要守住这个原则,后续加新模型、新数据集、新训练技巧,都只是增加一个配置项或一个类的事,不会让代码变成屎山。
在实际操作中,我最想提醒大家的一点是:别一上来就追求完美的框架。先把自己手头的实验跑通,哪怕代码丑一点、逻辑乱一点都没关系。等跑通了两三个实验,你自然会发现有些代码在反复复制,有些函数在频繁改动,那时候再动手重构,方向会准确得多。我这套框架也不是凭空设计的,是跑了十几个实验之后才慢慢抽象成现在的样子。
最后分享一个我保存模型的小习惯:除了保存 best model,每个 epoch 结束也可以保留最近一次的权重作为 last model。因为有的实验在验证集最高点之后可能还继续涨,记录 last checkpoint 能让你在发现验证集 acc 还在上升的时候,反手从 last checkpoint 继续训练,而不是只能从 best 重新开始。多一个存档,总比少一个好。