PyTorch 入门实战(二):手写 CNN 实现 CIFAR-10 彩色图像分类
前言
本文会用PyTorch从零搭建一个卷积神经网络,在CIFAR-10彩色图像数据集上完成 10 分类任务。代码量不到 150 行,但覆盖了深度学习项目的完整流程:数据加载 → 模型搭建 → 训练 → 评估 → 可视化。
读完你会掌握:
- CIFAR-10 数据集的结构与预处理
- 如何用
nn.Conv2d+nn.MaxPool2d搭建一个基础 CNN - 为什么彩色图像输入通道是 3,全连接层维度怎么算
- 完整的训练循环写法
- 模型评估与预测结果可视化
一、CIFAR-10 数据集简介
CIFAR-10 是深度学习入门的经典 benchmark,由 60000 张 32×32 的彩色图像组成,共 10 个类别:
飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车- 训练集:50000 张
- 测试集:10000 张
- 每张图:3 通道 × 32 像素 × 32 像素
由于 CIFAR-10 的图片尺寸只有 32×32,即使简单的 CNN 也能在 CPU 上快速训练,非常适合用来练手。
️ 二、环境配置
你需要安装以下 Python 库:
pipinstalltorch torchvision matplotlib numpyPyTorch 建议去官网根据你的 CUDA 版本选择合适的安装命令。如果没有 GPU,CPU 版也能跑,只是稍慢一点。
三、第一步:导入库 & 检测设备
importtorchimporttorch.nnasnnimporttorch.optimasoptimfromtorchvisionimportdatasets,transformsimportmatplotlib.pyplotaspltimportnumpyasnp# 解决中文乱码plt.rcParams['font.sans-serif']=['SimHei']plt.rcParams['axes.unicode_minus']=False# 自动选择 GPU 或 CPUdevice=torch.device("cuda"iftorch.cuda.is_available()else"cpu")print(f" 使用设备:{device}")要点:torch.device会自动检测是否有可用的 GPU,有就用 CUDA,没有就回退到 CPU。有了这个 device 变量之后,所有张量和模型都通过.to(device)统一迁移,代码完全硬件无关。
四、第二步:加载 CIFAR-10 数据
4.1 数据预处理
transform=transforms.Compose([transforms.ToTensor(),# PIL → Tensor,像素值缩放到 [0, 1]transforms.Normalize((0.4914,0.4822,0.4465),# 按通道减均值(0.2023,0.1994,0.2010))# 按通道除标准差])为什么用这组均值和标准差?这是 CIFAR-10 官方推荐的数值,能让每个通道的数据变成近似零均值、单位方差的分布,从而加速收敛、稳定训练。简单说——用了比不用训得更快。
4.2 加载数据
train_dataset=datasets.CIFAR10(root='./data',train=True,download=True,transform=transform)test_dataset=datasets.CIFAR10(root='./data',train=False,download=True,transform=transform)train_loader=torch.utils.data.DataLoader(train_dataset,batch_size=64,shuffle=True)test_loader=torch.utils.data.DataLoader(test_dataset,batch_size=64,shuffle=False)几个关键参数:
train=True/False:区分训练集和测试集download=True:首次运行会自动下载(约 170MB)batch_size=64:每次取 64 张图一起训练,平衡了内存消耗和梯度稳定性shuffle=True:训练时打乱顺序,测试时不需要
️ 五、第三步:搭建 CNN 网络
这是整个项目的核心。我们设计了一个轻量级网络,结构如下:
输入 (3×32×32) ↓ Conv2d(3→32, 3×3, padding=1) → ReLU → MaxPool(2×2) → 输出 (32×16×16) ↓ Conv2d(32→64, 3×3, padding=1) → ReLU → MaxPool(2×2) → 输出 (64×8×8) ↓ Flatten (64×8×8 = 4096) ↓ Linear(4096 → 256) → ReLU ↓ Linear(256 → 10)classSimpleCNN(nn.Module):def__init__(self):super(SimpleCNN,self).__init__()# 卷积层 1:输入 3 通道(彩色图),输出 32 个特征图,卷积核 3×3self.conv1=nn.Conv2d(in_channels=3,out_channels=32,kernel_size=3,padding=1)self.pool=nn.MaxPool2d(kernel_size=2,stride=2)# 卷积层 2:输入 32 通道,输出 64 个特征图self.conv2=nn.Conv2d(in_channels=32,out_channels=64,kernel_size=3,padding=1)# 全连接层:两次池化后尺寸变为 64×8×8 = 4096self.fc1=nn.Linear(64*8*8,256)self.fc2=nn.Linear(256,10)defforward(self,x):x=self.pool(torch.relu(self.conv1(x)))# Conv1 + ReLU + Poolx=self.pool(torch.relu(self.conv2(x)))# Conv2 + ReLU + Poolx=x.view(-1,64*8*8)# 展平x=torch.relu(self.fc1(x))# FC1 + ReLUx=self.fc2(x)# 输出层returnx model=SimpleCNN().to(device)print(model)关键问题:全连接层维度怎么来的?
这是新手最容易困惑的地方。我们来一步步推:
| 步骤 | 输入尺寸 | 操作 | 输出尺寸 |
|---|---|---|---|
| 原始图像 | 3 × 32 × 32 | - | 3 × 32 × 32 |
| Conv1 (3→32, 3×3, pad=1) | 3 × 32 × 32 | 卷积 | 32 × 32 × 32 |
| MaxPool (2×2) | 32 × 32 × 32 | 池化 | 32 × 16 × 16 |
| Conv2 (32→64, 3×3, pad=1) | 32 × 16 × 16 | 卷积 | 64 × 16 × 16 |
| MaxPool (2×2) | 64 × 16 × 16 | 池化 | 64 × 8 × 8 |
所以最终展平后是64 × 8 × 8 = 4096,这就是nn.Linear(4096, 256)的输入维度。
padding=1的作用:卷积核 3×3 本来会让尺寸缩小 2(两边各少 1),加了padding=1后正好抵消,输出尺寸等于输入尺寸。这样池化才是唯一缩小尺寸的操作,计算路径更清晰。
六、第四步:损失函数与优化器
criterion=nn.CrossEntropyLoss()# 多分类交叉熵optimizer=optim.Adam(model.parameters(),lr=0.001)为什么用CrossEntropyLoss?
- 它内部已经集成了Softmax + 负对数似然,不需要在模型最后一层手动加 Softmax
- 输出层直接出原始分数(logits),CrossEntropyLoss 会自动处理
- 数值稳定性更好(内部用了 log-sum-exp 技巧)
为什么选 Adam 而不是 SGD?
Adam 结合了动量(Momentum)和自适应学习率(RMSProp)的优点:收敛快、对学习率不敏感、大多数情况下开箱即用。对于入门项目,Adam 是首选。
️ 七、第五步:训练循环
epochs=10forepochinrange(epochs):running_loss=0.0fori,(images,labels)inenumerate(train_loader):images,labels=images.to(device),labels.to(device)# 前向传播outputs=model(images)loss=criterion(outputs,labels)# 反向传播optimizer.zero_grad()# 清空梯度loss.backward()# 计算梯度optimizer.step()# 更新参数running_loss+=loss.item()if(i+1)%500==0:print(f"[{epoch+1}/{epochs}] Step{i+1}, Loss:{loss.item():.4f}")print(f"Epoch{epoch+1}结束, 平均损失:{running_loss/len(train_loader):.4f}")训练循环三件套
这是每个 PyTorch 训练代码中固定的三步,理解了就不会忘:
optimizer.zero_grad()# ① 清空上一轮的梯度,否则会累加loss.backward()# ② 反向传播,计算各参数的 ∂L/∂woptimizer.step()# ③ 沿梯度反方向更新参数:w = w - lr × grad如果你忘了第 ① 步,梯度会不断累加,模型就学歪了——这是新手最常见的问题之一。
八、第六步:测试准确率
correct=0total=0withtorch.no_grad():# 不计算梯度,省内存、加速forimages,labelsintest_loader:images,labels=images.to(device),labels.to(device)outputs=model(images)_,predicted=torch.max(outputs.data,1)# 取概率最大的类别total+=labels.size(0)correct+=(predicted==labels).sum().item()print(f" 测试集准确率:{100*correct/total:.2f}%")几个关键细节:
torch.no_grad():推理时禁用梯度计算,大幅减少显存占用torch.max(outputs, 1):dim=1 表示在类别维度上取最大值,返回 (values, indices)model.eval()通常也应该调(切换到评估模式,影响 Dropout 和 BatchNorm),本项目没加但建议补上
这个简单的 CNN 训 10 个 epoch 大约能达到70% ~ 75%的准确率,作为基线模型已经不错了(随机猜是 10%)。
️ 九、第七步:随机抽取预测结果可视化
dataiter=iter(test_loader)images,labels=next(dataiter)images,labels=images.to(device),labels.to(device)outputs=model(images)_,predicted=torch.max(outputs,1)images=images.cpu().numpy()labels=labels.cpu().numpy()predicted=predicted.cpu().numpy()plt.figure(figsize=(12,6))classes=['飞机','汽车','鸟','猫','鹿','狗','青蛙','马','船','卡车']foriinrange(5):plt.subplot(1,5,i+1)img=images[i].transpose((1,2,0))# (3,32,32)→(32,32,3)img=img*np.array([0.2023,0.1994,0.2010])+np.array([0.4914,0.4822,0.4465])img=np.clip(img,0,1)plt.imshow(img)plt.title(f"真实:{classes[labels[i]]}\n预测:{classes[predicted[i]]}")plt.axis('off')plt.tight_layout()plt.show()注意transpose:PyTorch 的图像张量是(C, H, W),而plt.imshow需要(H, W, C),所以必须转置。
注意逆归一化:训练时做过的 Normalize 操作需要还原,否则图片颜色会失真(偏暗、偏蓝)。
十、完整代码
importtorchimporttorch.nnasnnimporttorch.optimasoptimfromtorchvisionimportdatasets,transformsimportmatplotlib.pyplotaspltimportnumpyasnp plt.rcParams['font.sans-serif']=['SimHei']plt.rcParams['axes.unicode_minus']=Falsedevice=torch.device("cuda"iftorch.cuda.is_available()else"cpu")print(f" 使用设备:{device}")# 数据加载transform=transforms.Compose([transforms.ToTensor(),transforms.Normalize((0.4914,0.4822,0.4465),(0.2023,0.1994,0.2010))])train_dataset=datasets.CIFAR10(root='./data',train=True,download=True,transform=transform)test_dataset=datasets.CIFAR10(root='./data',train=False,download=True,transform=transform)train_loader=torch.utils.data.DataLoader(train_dataset,batch_size=64,shuffle=True)test_loader=torch.utils.data.DataLoader(test_dataset,batch_size=64,shuffle=False)# CNN 模型classSimpleCNN(nn.Module):def__init__(self):super(SimpleCNN,self).__init__()self.conv1=nn.Conv2d(3,32,3,padding=1)self.pool=nn.MaxPool2d(2,2)self.conv2=nn.Conv2d(32,64,3,padding=1)self.fc1=nn.Linear(64*8*8,256)self.fc2=nn.Linear(256,10)defforward(self,x):x=self.pool(torch.relu(self.conv1(x)))x=self.pool(torch.relu(self.conv2(x)))x=x.view(-1,64*8*8)x=torch.relu(self.fc1(x))x=self.fc2(x)returnx model=SimpleCNN().to(device)criterion=nn.CrossEntropyLoss()optimizer=optim.Adam(model.parameters(),lr=0.001)# 训练epochs=10forepochinrange(epochs):running_loss=0.0fori,(images,labels)inenumerate(train_loader):images,labels=images.to(device),labels.to(device)outputs=model(images)loss=criterion(outputs,labels)optimizer.zero_grad()loss.backward()optimizer.step()running_loss+=loss.item()if(i+1)%500==0:print(f"[{epoch+1}/{epochs}] Step{i+1}, Loss:{loss.item():.4f}")print(f"Epoch{epoch+1}结束, 平均损失:{running_loss/len(train_loader):.4f}")# 测试correct=total=0withtorch.no_grad():forimages,labelsintest_loader:images,labels=images.to(device),labels.to(device)outputs=model(images)_,predicted=torch.max(outputs.data,1)total+=labels.size(0)correct+=(predicted==labels).sum().item()print(f" 测试集准确率:{100*correct/total:.2f}%")# 可视化classes=['飞机','汽车','鸟','猫','鹿','狗','青蛙','马','船','卡车']dataiter=iter(test_loader)images,labels=next(dataiter)images,labels=images.to(device),labels.to(device)outputs=model(images)_,predicted=torch.max(outputs,1)images,labels,predicted=images.cpu().numpy(),labels.cpu().numpy(),predicted.cpu().numpy()plt.figure(figsize=(12,6))foriinrange(5):plt.subplot(1,5,i+1)img=images[i].transpose((1,2,0))img=img*np.array([0.2023,0.1994,0.2010])+np.array([0.4914,0.4822,0.4465])img=np.clip(img,0,1)plt.imshow(img)plt.title(f"真实:{classes[labels[i]]}\n预测:{classes[predicted[i]]}")plt.axis('off')plt.tight_layout()plt.show()十一、知识回顾
看完这篇博客,不妨自测一下:
| 问题 | 你的答案 |
|---|---|
| CIFAR-10 每张图的尺寸是多少?通道数呢? | |
padding=1的作用是什么? | |
nn.Linear的输入维度64×8×8是怎么推出来的? | |
optimizer.zero_grad()忘写会怎样? | |
为什么要逆归一化再plt.imshow? | |
| 卷积层和全连接层的区别是什么? |
如果能流畅答出 5 道以上,说明你已经吃透了这篇文章。
十二、下一步可以做什么?
- 加入 Dropout:在
fc1后面加一层nn.Dropout(0.5),观察准确率和过拟合情况 - 换成 Batch Normalization:在每层卷积和激活函数之间插入
nn.BatchNorm2d - 加深网络:把
SimpleCNN变成 4 层卷积,看看准确率能提升多少 - 学习率调度:用
torch.optim.lr_scheduler.StepLR在训练后期降低学习率 - 模型保存与加载:用
torch.save(model.state_dict(), 'cnn_cifar10.pth')保存训练好的模型