1. 先搞清楚卷积神经网络解决图像分类的什么痛点
图像分类是计算机视觉里最常见、也最适合入门的一类任务。简单说,就是给模型一张图,让它判断“这张图属于哪一类”。传统做法是人工提取特征,比如颜色直方图、边缘检测、纹理描述,再交给分类器。这套流程在简单场景下能跑,但一到真实世界就变得脆弱:光照一变、背景一换、物体角度一偏,特征就不稳定了。
卷积神经网络(CNN)解决的核心问题,是把“特征提取”和“分类决策”一起交给模型自己学。你不再需要人为设计特征规则,只需要提供足够多的标注图像,模型就会自动从像素里学习到有用信息。早期卷积层学习边缘、纹理,深层卷积层学习部件、形状,最后通过全连接层输出每个类别的概率。这个能力让图像分类从“拼特征工程”变成了“拼数据和调参”。
这篇文章要做的,是带你用 PyTorch 从零跑通一个图像分类实战任务。包括环境安装、数据准备、卷积网络搭建、训练与评估、保存与加载模型,以及我实际运行中遇到的常见问题排查。适合已经会基本 Python 语法、想正式进入深度学习视觉方向的人。如果你是第一次接触 PyTorch,也没关系,我会把每一步为什么这么做讲清楚。
最值得关注的点是:用 PyTorch 做图像分类,真正的难点不在“写模型结构”,而在数据流程、训练节奏、资源分配和问题定位。很多人拿着别人代码跑通了 MNIST,换个数据集就报错,原因就是没有理解每一步之间的依赖关系。这篇文章会重点把这条链路打通。
2. 环境准备:PyTorch 安装和硬件选型是第一步,也是很多人卡住的地方
2.1 CPU 和 GPU 怎么选,先想清楚
图像分类任务能不能跑、跑得快不快,首先取决于你手里的计算资源。官方给出的定位很明确:PyTorch 支持 CPU 版本和 CUDA 加速版本。CPU 版本适合入门验证、跑小数据集、改代码逻辑,但训练稍大一点的模型就会很慢。GPU 版本能把训练时间从几小时压到十几分钟,但需要满足硬件和驱动条件。
我的建议是:第一轮实战不要纠结 GPU,先让代码跑通。即使只有 CPU,也可以完成本篇文章的全部步骤,只是训练时间会更长。第一次跑通的价值在于理解流程,而不是追求速度。如果后续想认真做实验,再考虑 GPU 环境和线上算力资源。
2.2 用 Anaconda 创建环境,避免依赖冲突
Python 深度学习的依赖关系非常敏感,尤其是 PyTorch、CUDA、Python 版本之间的匹配关系。直接往系统环境里安装,很容易出现一个项目需要旧版本、另一个项目需要新版本的问题。最省心的做法是使用 Anaconda 或 Miniconda 创建独立环境。
创建环境时,建议先明确 Python 版本。PyTorch 对 Python 版本的支持范围比较广,但也不建议直接上最高版本,使用官方文档推荐的稳定版本更稳妥。创建命令类似:
conda create -n pytorch_cnn python=3.10 conda activate pytorch_cnn进入环境之后,再安装 PyTorch。安装命令不要自己去拼版本,直接去 PyTorch 官网的安装页选择对应系统、包管理工具和 CUDA 版本,复制生成的命令即可。官网会根据你的选择自动给出兼容的下载命令,这是避免版本错配最可靠的方式。
2.3 CPU 和 GPU 的安装差异
如果你打算用 CPU 训练,官网会提供类似这样的安装命令:
pip install torch torchvision torchaudio如果你打算用 GPU,则通常需要先安装 NVIDIA 驱动,然后选择合适的 CUDA 版本,再安装对应的 PyTorch 版本:
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118这里的cu118代表 CUDA 11.8。实际版本号要以你机器上的驱动和官网推荐为准,不建议照抄。驱动支持的 CUDA 版本和 PyTorch 对应版本如果不匹配,就可能出现“PyTorch 无法使用 GPU”的警告。
安装完之后,有一个必须做的验证步骤:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"如果输出True,说明 GPU 版本生效。如果输出False,说明当前装的 PyTorch 是 CPU 版本,或者 CUDA 环境没有配置好。这个验证很关键,因为它直接决定后续训练能否用到 GPU。
注意:不要因为网上很多帖子里写着某个 CUDA 版本,就直接照搬。不同驱动、不同显卡、不同 PyTorch 版本之间的匹配关系很复杂,最稳妥的做法是用官网生成安装命令。
2.4 其他依赖和项目结构准备
除了 PyTorch,还需要安装torchvision。它提供了常见数据集接口和预训练模型,也包含图像处理变换工具。这会让实验代码简单不少。
我一般还会安装matplotlib用于画训练曲线,安装numpy用于数据操作。这些用pip install matplotlib numpy即可。
另外,建议一开始就把项目目录规划好:
cnn_classify/ ├── data/ # 存放下载的数据集 ├── checkpoints/ # 存放训练好的模型权重 ├── logs/ # 存放训练日志 └── scripts/ # 存放训练和测试脚本这个结构看起来简单,但实际作用很大。后面训练过程中,你会频繁修改输入路径、保存输出和查看结果。提前固定目录,可以省掉很多路径错误。
3. 数据准备:图像分类不是把图片塞给模型就行,先要统一尺寸和数值范围
3.1 选一个适合实战的数据集
常见入门图像分类数据集有:
- MNIST:手写数字,单通道 28x28,适合验证流程。
- FashionMNIST:衣物类图片,同样是单通道 28x28,比数字更有区分度。
- CIFAR-10:10 类彩色图,32x32,是入门到进阶之间很好的过渡。
- Food-101、CIFAR-100、ImageNet 子集:类别更多、图像更复杂,适合进阶练习。
本篇文章以 CIFAR-10 为例。它包含飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车十个类别,训练集 50000 张,测试集 10000 张,图片尺寸 32x32,三通道彩色图。这个数据集不大,但比 MNIST 接近真实场景,用来学习 CNN 图像分类很合适。
3.2 数据变换:Resize、ToTensor、Normalize
很多新手拿到图像数据集后,直接就把原始图片丢给模型训练,这经常会出问题。因为不同数据集的图片尺寸、通道数、像素范围不一致,模型输入是固定维度的,图片必须先做统一处理。
PyTorch 里处理图像通常用torchvision.transforms。CIFAR-10 本身就很小,但为了兼容不同输入大小,我习惯在训练流程中加上Resize,把图片缩放到统一尺寸。如果数据是真实拍摄的图片,这一步基本是必须的。
核心处理步骤有两个:
ToTensor():把 PIL.Image 或 numpy 数组转成 PyTorch 张量,同时把像素值从 [0, 255] 缩放到 [0.0, 1.0]。Normalize():把每个通道的数值做标准化,让均值接近 0,方差接近 1。这样可以让梯度更新更稳定,训练更容易收敛。
CIFAR-10 的常用标准化参数是mean=(0.485, 0.456, 0.406)和std=(0.229, 0.224, 0.225)。这些数值来自 ImageNet 数据集的统计值,在很多视觉任务里被直接复用。如果你的数据集差异很大,可以自己统计均值方差再传入。
代码示例:
from torchvision import transforms transform = transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)) ])3.3 数据加载:Dataset 与 DataLoader 的关系
PyTorch 里有两个容易混淆的接口:Dataset负责“如何取一条数据”,DataLoader负责“如何分批取出多条数据并打乱顺序”。
torchvision.datasets.CIFAR10内置了数据下载接口,第一次运行会自动下载到root参数指定的目录。如果下载过慢,可以手动下载后解压到对应目录。
加载训练集和测试集:
from torchvision import datasets train_dataset = datasets.CIFAR10( root='data', train=True, transform=transform, download=True ) test_dataset = datasets.CIFAR10( root='data', train=False, transform=transform, download=True )然后创建 DataLoader:
from torch.utils.data import DataLoader batch_size = 64 train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)这里有几个值得注意的点:
- 训练集设置了
shuffle=True,目的是每个 epoch 内样本顺序都被打乱,避免模型记住固定顺序。 - 测试集设置
shuffle=False,方便后续评估时按顺序核对样本标签。 batch_size不宜一开始就拉大。CPU 环境建议 16 或 32,GPU 显存足够再上调到 64 或 128。
3.4 先可视化一批数据再训练
我在每次启动训练前,都会先抽一批数据做可视化。这不是浪费时间,而是最直接的输入检查。把图像和标签打印出来,确认通道顺序、尺寸、类别名称是否正确。
PyTorch 的张量默认是[C, H, W]结构,而matplotlib显示图像需要[H, W, C],所以需要转换维度。此外,因为前面做了标准化,直接显示图片会偏暗,需要把均值方差加回去再显示。
这一步的目的很简单:先不要让模型面对你看不懂的数据。输入一旦出错,后面的结果必然混乱,而且很难定位问题源头。
4. 构建 CNN 模型:先搭一个能跑出结果的网络,再谈效果
4.1 CNN 的基本结构拆解
一个用于图像分类的卷积神经网络,通常由下面几个部分组成:
- 卷积层(Conv2d):提取局部特征。通过卷积核在图像上滑动,学习到边缘、颜色、纹理等信息。
- 激活函数(ReLU):增加非线性。如果没有激活函数,多层卷积叠加完还是一个线性变换,学习能力很弱。
- 池化层(MaxPool2d):降低特征图尺寸,减少计算量,同时保留主要特征。
- 全连接层(Linear):将特征图展平后映射到类别概率。
- 输出层:一般使用
LogSoftmax或直接输出一个未归一化的 logits 向量,交给损失函数处理。
这个结构不是唯一选择,但适合作为第一版模型。先把基础结构跑通,再逐步替换成残差网络、带注意力机制的模块。
4.2 用 PyTorch 实现一个简单的 CNN
下面是一个适配 CIFAR-10 的简单 CNN 示例。输入是三通道 32x32 图像,输出是 10 个类别的得分:
import torch import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes=10): super().__init__() self.features = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=2, stride=2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=2, stride=2), nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=2, stride=2), ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(128 * 4 * 4, 256), nn.ReLU(inplace=True), nn.Linear(256, num_classes) ) def forward(self, x): x = self.features(x) x = self.classifier(x) return x4.3 每一步参数为什么这么设
首先看nn.Conv2d(3, 32, kernel_size=3, padding=1)。第一个参数3是输入通道数,对应 RGB 三个通道。第二个参数32是输出通道数,也就是这一层提取 32 种不同特征。kernel_size=3表示使用 3x3 的卷积核。padding=1表示边缘补一圈 0,保证卷积后特征图尺寸不缩小。
接着是MaxPool2d(kernel_size=2, stride=2),每次把特征图宽高减半。32x32 的图像经过三次池化后变成 4x4,所以全连接层第一层输入是128 * 4 * 4。这个数值不是随便写的,必须跟前面卷积池化的输出尺寸一致。如果你改动了池化次数或者输入尺寸,这里就要同步调整。
这里没有加入BatchNorm,因为我想先把最基础的结构跑通,减少变量干扰。加入之后通常能加速收敛,但也会让问题排查多一层复杂。第一版模型越简单越好。
4.4 为什么不直接使用预训练模型
torchvision.models里提供了 ResNet、VGG、EfficientNet 等预训练模型,用起来确实方便,效果也比自己搭的简单网络好很多。但既然标题是实战,我建议第一遍还是自己搭一个简单 CNN。
原因是:预训练模型把“设计网络结构”这件事遮住了。遇到输入尺寸不对,你知道要改最后全连接层,但不清楚前面为什么有那么多适配步骤。自己搭一遍网络,你能直观看到每个参数的变化如何影响特征图尺寸、参数量和最终输出。
等第一遍训练结束,再切换到resnet18或mobilenet_v2,你会更容易理解迁移学习的意义。
5. 训练模型:损失函数、优化器和训练循环是三个独立概念
5.1 损失函数:交叉熵是怎么判断预测好坏的
图像分类中,最常用的损失函数是交叉熵损失。PyTorch 里对应的类是:
criterion = nn.CrossEntropyLoss()你不需要手动计算 softmax,CrossEntropyLoss内部已经包含 softmax 操作。它做的事情是:把模型输出的原始分数(logits)转成概率分布,然后与真实标签做交叉熵计算。预测越接近真实标签,损失值越小。
这里有一个常见误区:模型最后一层输出的不是概率,而是每个类别的“未归一化分数”。很多人只看输出里哪个值最大,就直接当置信度用,这不够严谨。要得到概率,需要先经过 softmax。在训练过程中,CrossEntropyLoss已经处理了这一步,所以不需要额外加nn.Softmax。但在推理阶段,如果你想输出概率,就需要手动加:
probabilities = torch.softmax(logits, dim=1)5.2 优化器:Adam 和 SGD 怎么取舍
优化器决定了模型参数如何根据损失更新。常见选择是 SGD 和 Adam。
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)Adam 适合刚入门时使用,因为它对学习率不那么敏感,默认参数在很多任务上都能收敛。SGD 需要更精细的学习率设置和动量参数,但有些场景下最终收敛效果更好。
我给你的建议是:第一版先用 Adam,跑通流程。后续想提升精度时,再换 SGD + Momentum 或者对学习率做衰减。不要一上来就在优化器上纠结太久。
5.3 完整训练循环
PyTorch 的训练过程通常包含以下几个步骤:
- 模型设为训练模式:
model.train() - 遍历 DataLoader 中的每个 batch。
- 将数据传入模型得到输出。
- 计算损失。
- 清空梯度:
optimizer.zero_grad() - 反向传播:
loss.backward() - 更新参数:
optimizer.step() - 记录损失,观察训练是否正常。
对应代码如下:
from tqdm import tqdm device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = SimpleCNN(num_classes=10).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=0.001) num_epochs = 10 for epoch in range(num_epochs): model.train() running_loss = 0.0 for images, labels in tqdm(train_loader): images = images.to(device) labels = labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss += loss.item() avg_loss = running_loss / len(train_loader) print(f"Epoch [{epoch+1}/{num_epochs}] Loss: {avg_loss:.4f}")这里每行代码都有明确目的,尤其是三步顺序:
zero_grad()必须先调用,否则 PyTorch 会默认累积梯度,导致参数更新错误。loss.backward()计算每个参数的梯度。optimizer.step()使用梯度更新参数。
顺序一旦颠倒,训练结果就会出问题,而且这类错误不会直接报错,只会表现为损失不下降或乱跳。
5.4 训练过程中要关注什么
我一般会在训练时盯三个指标:
- loss 是否在下降。如果是,说明模型在学习。
- loss 波动幅度。波动大不一定有问题,但如果是持续上升,就要检查学习率是不是太大。
- 每个 epoch 的耗时。如果耗时突然飙升,多半是数据处理或日志记录环节出了问题。
另外,不要指望第一个 epoch 就有很高的准确率。CIFAR-10 使用简单 CNN 从零训练,通常训练 5 到 10 个 epoch 后,准确率才会明显提升。如果只跑一两个 epoch 就下结论,说明还没等到模型真正开始收敛。
6. 评估模型:准确率是最基础的指标,但还不够
6.1 计算测试集准确率
训练完成后,必须用模型从未见过的测试集来评估,这样才能判断模型的泛化能力。测试代码和训练代码略有不同:
- 不需要计算梯度。
- 不需要更新参数。
- 模型需要切换为评估模式:
model.eval()
评估代码示例:
def evaluate(model, dataloader): model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in dataloader: images = images.to(device) labels = labels.to(device) outputs = model(images) _, predicted = torch.max(outputs, dim=1) total += labels.size(0) correct += (predicted == labels).sum().item() accuracy = 100.0 * correct / total return accuracy test_acc = evaluate(model, test_loader) print(f"Test Accuracy: {test_acc:.2f}%")torch.no_grad()的作用是关闭梯度计算,减少显存和内存占用,同时加速推理。
6.2 为什么要同时看训练集准确率
只看测试集准确率还不够。如果测试集准确率很低,你无法判断是模型没学到东西,还是过拟合导致的。所以我会同时计算训练集准确率。
如果训练集准确率高、测试集准确率低,说明模型过拟合了。此时需要对模型做正则化,比如增加数据增强、加入 Dropout、减小模型复杂度。
如果训练集和测试集准确率都低,说明模型欠拟合。此时应该增大模型容量、增加训练轮数、调整学习率,或者检查数据预处理是否出错。
6.3 用混淆矩阵观察具体错误
准确率只能告诉你整体效果,不能告诉你模型在哪类样本上出错。对于 CIFAR-10,比如汽车和卡车很容易混淆,鸟和飞机有时也会看错。用混淆矩阵可以直观看到错误集中在哪些类别之间。
混淆矩阵的实现在 sklearn 中可以直接调用:
from sklearn.metrics import confusion_matrix all_preds = [] all_labels = [] model.eval() with torch.no_grad(): for images, labels in test_loader: images = images.to(device) outputs = model(images) _, predicted = torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) print(cm)拿到混淆矩阵后,重点不是打印它,而是找出第 i 类被误分类为第 j 类最多的组合,去查看对应样本。我见过太多人只盯着准确率数字,从不看实际错误样本,结果一直找不到改进方向。
6.4 预测结果显示的常见坑
在显示预测结果时,图像是经过标准化的,直接用 matplotlib 显示会不正常。需要先反标准化:
def denormalize(tensor): mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1) std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1) tensor = tensor * std + mean return tensor如果不做这一步,图片看起来亮度偏低甚至发绿,这不是模型错,而是显示问题。
7. 保存和加载模型:训练完不保存,相当于白跑
7.1 保存完整模型与保存状态字典的区别
PyTorch 保存模型有几种方式,最推荐的是保存state_dict,也就是模型的权重参数。这种方式体积小、兼容性好,加载时还需要你重新定义模型结构。
保存方式:
checkpoint_path = "checkpoints/simple_cnn_cifar10.pth" torch.save(model.state_dict(), checkpoint_path)加载方式:
model = SimpleCNN(num_classes=10) model.load_state_dict(torch.load(checkpoint_path, map_location=device)) model.to(device)使用map_location可以解决在 GPU 上训练、在 CPU 上加载的问题。如果不加这个参数,加载时可能会报设备不存在的错误。
7.2 保存更多信息:优化器状态、epoch、损失历史
如果是长期训练任务,建议不只是保存权重,还要保存优化器状态、当前 epoch、训练损失历史等。这样即使训练中断,也可以从断点继续训练。
checkpoint = { "epoch": epoch + 1, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "loss": avg_loss, } torch.save(checkpoint, "checkpoints/checkpoint_last.pth")继续训练时:
checkpoint = torch.load("checkpoints/checkpoint_last.pth") model.load_state_dict(checkpoint["model_state_dict"]) optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) start_epoch = checkpoint["epoch"]很多初学者只保存权重,后来想用同一模型继续训练时,才发现优化器状态全部丢失,只能从零开始重新调整学习率。提前多存几个字段,是节省时间的好习惯。
7.3 保存多个模型文件,而不是只覆盖最后一个
我会按不同策略保存模型:
checkpoint_last.pth:最新状态,用于断点续训。checkpoint_best.pth:测试集准确率最高时的模型。checkpoint_epoch_xx.pth:每隔几个 epoch 保存一次,便于复盘。
这样训练崩溃或效果不理想时,你还有历史版本可以回溯。不要只保留一个文件。
8. 数据增强与正则化:想让分类准确率更高,这两步不可跳过
8.1 数据增强让模型更鲁棒
数据增强不是蒙骗模型,而是从已有数据中生成更多样本来,让模型看到更多位置、角度、亮度变化。这样模型就不容易记住训练样本的某个固定外观。
常见的适用于 CIFAR-10 的数据增强:
- 随机水平翻转:
transforms.RandomHorizontalFlip() - 随机旋转:
transforms.RandomRotation(10) - 随机裁剪:
transforms.RandomCrop(32, padding=4) - 颜色抖动:
transforms.ColorJitter(brightness=0.2, contrast=0.2)
把这些操作加入 transform:
train_transform = transforms.Compose([ transforms.RandomCrop(32, padding=4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)) ])需要注意:测试集和验证集通常不做数据增强,只做尺寸调整、转张量和标准化。因为测试时你希望看到模型对每一张图的稳定判断,而不是随机变化后的结果。
数据增强也不是越多越好。增强太强,比如对真实图片做极端颜色偏移,反而可能让训练集和测试集分布差异变大,导致模型学不到有效特征。
8.2 Dropout 和 BatchNorm 怎么配合使用
Dropout 是一种正则化手段,做法是在训练时随机让一部分神经元不参与计算,减少神经元之间固定的依赖关系,降低过拟合风险。
在 PyTorch 中,Dropout 通常放在全连接层后面:
self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(128 * 4 * 4, 256), nn.ReLU(inplace=True), nn.Dropout(0.5), nn.Linear(256, num_classes) )0.5表示训练时随机丢弃一半神经元。推理时,Dropout 层不会生效,模型会使用全部神经元,所以不需要手动切换。
BatchNorm 的通常作用是缓解内部协变量偏移,让每一层的数据分布更稳定,训练时可使用较大学习率,收敛更快。但要注意,BatchNorm 在model.train()和model.eval()模式下的行为不同。训练时它使用当前 batch 的统计量,评估时使用训练阶段累积的统计量。
如果你加了 BatchNorm,一定要记得及时切换model.train()和model.eval(),否则结果会不稳定。
8.3 从零训练好,还是用迁移学习好
对于 CIFAR-10 这种中等规模数据集,从零训练一个简单 CNN 有一定效果,但很难达到接近最先进模型的精度。原因很简单:数据量不够大,模型学习到的特征不够通用。
如果换成真实项目场景,我更推荐使用预训练模型做迁移学习。比如 ResNet18 在 ImageNet 上预训练过,已经学会了很多基础特征,你只需要修改最后一层全连接输出为 10 类,然后在目标数据集上微调。
import torchvision.models as models model = models.resnet18(pretrained=True) num_features = model.fc.in_features model.fc = nn.Linear(num_features, 10)迁移学习的优势在于:即使你的数据集只有几千张图片,模型也能借助预训练特征获得不错效果。但从实战学习的角度,我仍然建议你先自己训练一个简单 CNN,理解参数变化,再切换到迁移学习。
9. 训练技巧与调参:学习率、批量大小、训练轮次怎么配合
9.1 学习率是最需要反复调的参数
学习率控制每次参数更新的步长。学习率过大,损失会在最小值附近震荡甚至发散;学习率过小,训练收敛速度慢,耗时很长。
一个常见经验是:先用固定的学习率1e-3跑几个 epoch,观察损失曲线,再决定是否降低。如果损失下降很慢,可以调大一点;如果损失上下剧烈跳动,就调小一点。
我通常在训练后期使用学习率衰减,让模型在接近最优解时用更小的步长精细调整。PyTorch 提供了一些现成调度器:
from torch.optim.lr_scheduler import StepLR scheduler = StepLR(optimizer, step_size=5, gamma=0.1)每个step_size轮之后,学习率乘以gamma。同样,这组参数不是绝对的,要根据训练效果调整。
9.2 批量大小与梯度更新的关系
批量大小(batch size)决定了每次前向传播和反向传播用多少张图。批量越大,单次更新越“平滑”,但占用的显存或内存也越高。批量越小,更新越频繁,训练波动也越大。
如果你的损失曲线严重抖动,可以把 batch size 调大一些。但如果硬件资源有限,batch size 不能随意增大,另一种做法是降低学习率。
CPU 训练时,batch size 不宜过大,因为大 batch 会快速占满内存,同时 CPU 的计算瓶颈会更明显。建议 CPU 环境从 16 开始。
9.3 Epoch 数量:从过拟合和训练时间两个角度判断
训练轮次过多会增加过拟合风险,还会浪费计算资源。训练轮次太少,模型还没收敛,准确率偏低。
判断方法很直接:每跑完一个 epoch,同时计算训练集和测试集准确率。当测试集准确率连续多个 epoch 不再提升,甚至开始下降,说明继续训练的意义已经不大了。
新手常常犯的一个错是:只盯着训练集损失,认为 loss 越低越好,结果模型记住了训练集的噪声,换到新图片上效果很差。一定记得用测试集或验证集来判断训练是否停止。
9.4 使用 TensorBoard 观察训练过程
PyTorch 可以配合 TensorBoard 记录损失和指标曲线。如果你的任务比较复杂,手动打印 Loss 不够直观,可以这样接入:
from torch.utils.tensorboard import SummaryWriter writer = SummaryWriter("logs") # 在训练循环中 writer.add_scalar("Loss/train", avg_loss, epoch) writer.add_scalar("Acc/test", test_acc, epoch) # 训练结束后 writer.close()然后在命令行启动:
tensorboard --logdir logs浏览器访问提示的本地地址,就能看到曲线。这个问题可能不是模型结构问题,而是某个 batch 的数据标签错位导致的。某个类别数量特别少,模型在局部类别权重调整不到位,也会引起跳变。
10.3 训练速度特别慢
排除数据加载慢的因素后,最可能的原因是:
- 在 CPU 上训练复杂模型。
- 没有使用 GPU 版本的 PyTorch。
- 没有在训练循环中把数据移动到 GPU。
- 日志打印过于频繁,占用大量时间。
我建议你在训练前打印一行确认信息:
print(f"Using device: {device}")如果显示cpu,说明你虽然在有 GPU 的机器上,但当前没有使用 GPU 计算。另外,数据加载可以用num_workers参数开启多个进程提高吞吐:
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2)Windows 系统下如果设置num_workers大于 0 遇到主进程报错,可以在 DataLoader 里加if __name__ == "__main__":保护,或者把num_workers改为 0。这个在跨平台运行时特别常见。
10.4 GPU 显存不足怎么办
显存不足报错通常会显示CUDA out of memory。常见原因包括 batch size 太大、图像分辨率太高、模型中同时保存了多个版本的大张量。
处理顺序:
- 减小 batch size。
- 降低图片分辨率。
- 检查是否在
with torch.no_grad():内执行推理。 - 检查代码中是否持有大量无用的历史张量。
- 使用
torch.cuda.empty_cache()清空缓存。
不要一上来就买新显卡,大部分情况下调小 batch size 就能解决。
11. 代码组织优化:从能跑通到方便复用
11.1 把训练流程封装成函数
第一次跑通时,代码可以全部写在脚本里。但第二次、第三次就会发现问题:只改一个数据集或网络结构,就要重读一遍代码,很容易改漏。
一个较清晰的模块结构是:
cnn_classify/ ├── config.py # 超参数、路径、类别数 ├── dataset.py # 构建 Dataset 和 DataLoader ├── model.py # 网络结构定义 ├── train.py # 训练脚本 ├── evaluate.py # 评估脚本 └── utils.py # 可视化、日志、保存检查点这不是最复杂的工程结构,但对个人项目和中小型教学任务足够用了。把超参数集中到一个配置文件,改参数时不用在代码里到处搜索。
11.2 用配置文件统一管理参数
例如config.py:
class Config: num_classes = 10 batch_size = 64 learning_rate = 1e-3 num_epochs = 20 data_root = "data" checkpoint_dir = "checkpoints" log_dir = "logs"使用配置文件的好处是:模型、数据、训练逻辑互相解耦。想对比不同学习率时,不需要复制多份训练脚本,只需要修改配置并重新运行。
11.3 记录实验日志,方便回头看
每次训练增加一个日志记录,内容可以包括:
- 实验时间。
- 使用的数据集、模型结构。
- 超参数。
- 训练集和测试集准确率。
- 遇到的问题和解法。
这个习惯在你做多个实验对比时尤其有用。很多项目跑了十几次,最后只能靠文件名区分,根本分不清哪个是好模型。
12. 项目实战中的经验边界:哪些情况不能照搬这套流程
12.1 小数据集和类别不均衡问题
本篇文章以 CIFAR-10 为例,各类别数据量均衡,图片尺寸统一,标签完整。如果你的数据是自采的图片,可能有以下几个问题:
- 不同类别图片数量差异很大。
- 图片尺寸、比例、清晰度不一致。
- 标签存在错误。
- 同一张图片重复出现在训练集和测试集。
处理策略也很明确:先做类别分布统计;对样本少的类别做数据增强;使用分层抽样划分训练集、验证集;在训练前人工检查一批硬标签或错标签。
12.2 是否所有图像分类任务都需要 CNN
不一定。如果图片特征在传统特征空间中已经很好分,比如工业缺陷检测里的固定角度、固定光照场景,使用 SVM、随机森林等传统方法也能得到不错效果。CNN 的优势在于自动学习特征,但需要更多数据和计算资源。小样本业务场景中,传统方法有时反而更稳。
12.3 从 CIFAR-10 到真实项目的距离
CIFAR-10 是学术界常用的入门实验基准,它帮你验证模型、训练流程是否正常。真实项目往往没有这么规范的输入数据。你需要自己负责爬取或采集图像、清洗噪声样本、标注、划分数据、处理类别不均衡、设计评估指标等。
这也是我建议你完成本篇文章实验后,找一个小型真实数据集练手的原因。形式与 CIFAR-10 一致,但数据是无序、混乱和真实的,你会碰到更多问题,也会更理解这个流程。
12.4 生产环境和实验环境的区别
如果在生产环境部署图像分类模型,不能只关注准确率。还要考虑推理速度、模型大小、并发数、日志采集、异常输入处理、模型更新流程等。简单 CNN 模型可能准确率不够,但推理速度快;ResNet 效果好但体积大。部署时需要根据业务场景做选择性取舍。
如果模型要部署到移动端或者前端浏览器,还需要考虑量化、剪枝、ONNX 转换等后续操作。这里先把实验跑通,生产化部署是下一个阶段的事。
13. 延伸:从简单 CNN 到更现代的网络结构
13.1 ResNet 的核心思想
如果你已经完成了上面的简单 CNN,下一步建议看 ResNet。ResNet 通过残差连接(跳过连接)让梯度更容易从深层传到浅层,解决了深层网络训练困难的问题。
用torchvision.models调用 ResNet18 几乎是零成本:
model = models.resnet18(pretrained=False, num_classes=10)不需要自己实现残差结构,也能正常使用。但如果你对学习有更高要求,建议读源码,查看BasicBlock里conv1、bn1、relu、conv2、bn2和最后加回输入的过程。
13.2 数据标准化在预训练模型里更关键
如果使用在 ImageNet 上预训练的模型,输入图像的 mean 和 std 必须严格按照预训练时的统计值来设置。torchvision.models官方文档里推荐的 transform 就是前面提到的(0.485, 0.456, 0.406)和(0.229, 0.224, 0.225)。
如果用错 mean 和 std,预训练模型的特征提取能力会大大削弱。
13.3 多标签分类和检测任务的区别
本篇文章讨论的是单标签分类,也就是每张图片只属于一个类别。但如果你的业务场景是一张图同时包含猫和狗,那就不是单标签分类问题了,需要使用多标签分类,损失函数也要从CrossEntropyLoss改成BCEWithLogitsLoss。
再进一步,如果还要输出物体位置,那问题就变成目标检测。常用的框架有 Faster R-CNN、YOLO、SSD 等。它们往往也基于卷积神经网络做特征提取。学好基础分类,对后续学习检测很有帮助。
14. 一次完整的最小实验流程回顾
为了让你有一个整体感,我把从环境到结果的完整流程浓缩成下面的清单。你可以把这个当作自己复现时的参考:
- 创建 Anaconda 环境,安装 CPU 或 GPU 版 PyTorch。
- 创建项目目录,规划 data、checkpoints、logs、scripts。
- 用 torchvision 加载 CIFAR-10,定义 transform。
- 可视化第一批数据,确认图像和标签对得上。
- 定义简单 CNN 模型,确认输出维度是 10。
- 定义 CrossEntropyLoss 和 Adam 优化器。
- 写训练循环,注意 zero_grad、backward、step 的顺序。
- 每轮打印平均损失,观察是否下降。
- 跑 5 到 10 个 epoch,切换到 GPU 或调小 batch size 再继续。
- 在测试集上计算准确率。
- 保存模型权重。
- 写一个独立的预测脚本,加载模型,对单张图片做推理。
- 显示预测结果,检查反标准化是否做好。
- 如果效果不行,加数据增强、Dropout,或者换成预训练 ResNet。
这套流程不是最复杂的,也不涉及超大规模数据,但它是一个可以稳定复现、逐步扩展的基础框架。后续所有更复杂的图像分类项目,基本都在这个流程上做增改。
15. 写在最后的实操建议
在真实跑这个项目时,我最想强调的几点经验是:
第一,宁可先跑小数据,也不要一上来完任务全量。可以先从 CIFAR-10 中取 2000 张图片做一个小实验,验证代码流程没问题,再使用全量数据训练。这样能更早发现代码错误,节省排队等待时间。
第二,不要同时改动太多参数。很多人在训练效果不好时,今天调学习率,明天加数据增强,后天换网络结构,最后完全不知道是哪个改动生效。更稳妥的做法是,一次只改一个变量,并记录对应结果。
第三,训练时不要只关心 Loss,要同时关注输入数据、日志和资源占用。模型表现异常,不一定是模型问题,也可能是某个 batch 里的图片全是噪声、标签顺序出错、GPU 显存碎片化、数据集路径错误。
第四,遇到报错先读 Python 回溯信息,再找人问。很多报错信息已经提示了出错文件和行号,直接定位到那一行,往往能自己解决。不要上来就把整段代码复制进社区提问,别人也没法替你看环境。
第五,保存模型时,要把评估指标、epoch、配置一起写入检查点。这会让以后的复现和对比轻松很多。
图像分类是视觉任务里最基础、也是最适合建立信心的一类问题。用 PyTorch 从零实现一次,你对数据流、模型结构、训练过程、评估方式的理解都会更扎实。后面再接触目标检测、图像分割、图像生成,都会更顺畅。