PyTorch 入门实战(二):手写 CNN 实现 CIFAR-10 彩色图像分类

📅 2026/8/6 18:50:13
PyTorch 入门实战(二):手写 CNN 实现 CIFAR-10 彩色图像分类
PyTorch 入门实战二手写 CNN 实现 CIFAR-10 彩色图像分类前言本文会用PyTorch从零搭建一个卷积神经网络在CIFAR-10彩色图像数据集上完成 10 分类任务。代码量不到 150 行但覆盖了深度学习项目的完整流程数据加载 → 模型搭建 → 训练 → 评估 → 可视化。读完你会掌握CIFAR-10 数据集的结构与预处理如何用nn.Conv2dnn.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 版本选择合适的安装命令。如果没有 GPUCPU 版也能跑只是稍慢一点。三、第一步导入库 检测设备importtorchimporttorch.nnasnnimporttorch.optimasoptimfromtorchvisionimportdatasets,transformsimportmatplotlib.pyplotaspltimportnumpyasnp# 解决中文乱码plt.rcParams[font.sans-serif][SimHei]plt.rcParams[axes.unicode_minus]False# 自动选择 GPU 或 CPUdevicetorch.device(cudaiftorch.cuda.is_available()elsecpu)print(f 使用设备:{device})要点torch.device会自动检测是否有可用的 GPU有就用 CUDA没有就回退到 CPU。有了这个 device 变量之后所有张量和模型都通过.to(device)统一迁移代码完全硬件无关。四、第二步加载 CIFAR-10 数据4.1 数据预处理transformtransforms.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_datasetdatasets.CIFAR10(root./data,trainTrue,downloadTrue,transformtransform)test_datasetdatasets.CIFAR10(root./data,trainFalse,downloadTrue,transformtransform)train_loadertorch.utils.data.DataLoader(train_dataset,batch_size64,shuffleTrue)test_loadertorch.utils.data.DataLoader(test_dataset,batch_size64,shuffleFalse)几个关键参数trainTrue/False区分训练集和测试集downloadTrue首次运行会自动下载约 170MBbatch_size64每次取 64 张图一起训练平衡了内存消耗和梯度稳定性shuffleTrue训练时打乱顺序测试时不需要️ 五、第三步搭建 CNN 网络这是整个项目的核心。我们设计了一个轻量级网络结构如下输入 (3×32×32) ↓ Conv2d(3→32, 3×3, padding1) → ReLU → MaxPool(2×2) → 输出 (32×16×16) ↓ Conv2d(32→64, 3×3, padding1) → 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.conv1nn.Conv2d(in_channels3,out_channels32,kernel_size3,padding1)self.poolnn.MaxPool2d(kernel_size2,stride2)# 卷积层 2输入 32 通道输出 64 个特征图self.conv2nn.Conv2d(in_channels32,out_channels64,kernel_size3,padding1)# 全连接层两次池化后尺寸变为 64×8×8 4096self.fc1nn.Linear(64*8*8,256)self.fc2nn.Linear(256,10)defforward(self,x):xself.pool(torch.relu(self.conv1(x)))# Conv1 ReLU Poolxself.pool(torch.relu(self.conv2(x)))# Conv2 ReLU Poolxx.view(-1,64*8*8)# 展平xtorch.relu(self.fc1(x))# FC1 ReLUxself.fc2(x)# 输出层returnx modelSimpleCNN().to(device)print(model)关键问题全连接层维度怎么来的这是新手最容易困惑的地方。我们来一步步推步骤输入尺寸操作输出尺寸原始图像3 × 32 × 32-3 × 32 × 32Conv1 (3→32, 3×3, pad1)3 × 32 × 32卷积32 × 32 × 32MaxPool (2×2)32 × 32 × 32池化32 × 16 × 16Conv2 (32→64, 3×3, pad1)32 × 16 × 16卷积64 × 16 × 16MaxPool (2×2)64 × 16 × 16池化64 × 8 × 8所以最终展平后是64 × 8 × 8 4096这就是nn.Linear(4096, 256)的输入维度。padding1的作用卷积核 3×3 本来会让尺寸缩小 2两边各少 1加了padding1后正好抵消输出尺寸等于输入尺寸。这样池化才是唯一缩小尺寸的操作计算路径更清晰。六、第四步损失函数与优化器criterionnn.CrossEntropyLoss()# 多分类交叉熵optimizeroptim.Adam(model.parameters(),lr0.001)为什么用CrossEntropyLoss它内部已经集成了Softmax 负对数似然不需要在模型最后一层手动加 Softmax输出层直接出原始分数logitsCrossEntropyLoss 会自动处理数值稳定性更好内部用了 log-sum-exp 技巧为什么选 Adam 而不是 SGDAdam 结合了动量Momentum和自适应学习率RMSProp的优点收敛快、对学习率不敏感、大多数情况下开箱即用。对于入门项目Adam 是首选。️ 七、第五步训练循环epochs10forepochinrange(epochs):running_loss0.0fori,(images,labels)inenumerate(train_loader):images,labelsimages.to(device),labels.to(device)# 前向传播outputsmodel(images)losscriterion(outputs,labels)# 反向传播optimizer.zero_grad()# 清空梯度loss.backward()# 计算梯度optimizer.step()# 更新参数running_lossloss.item()if(i1)%5000:print(f[{epoch1}/{epochs}] Step{i1}, Loss:{loss.item():.4f})print(fEpoch{epoch1}结束, 平均损失:{running_loss/len(train_loader):.4f})训练循环三件套这是每个 PyTorch 训练代码中固定的三步理解了就不会忘optimizer.zero_grad()# ① 清空上一轮的梯度否则会累加loss.backward()# ② 反向传播计算各参数的 ∂L/∂woptimizer.step()# ③ 沿梯度反方向更新参数w w - lr × grad如果你忘了第 ① 步梯度会不断累加模型就学歪了——这是新手最常见的问题之一。八、第六步测试准确率correct0total0withtorch.no_grad():# 不计算梯度省内存、加速forimages,labelsintest_loader:images,labelsimages.to(device),labels.to(device)outputsmodel(images)_,predictedtorch.max(outputs.data,1)# 取概率最大的类别totallabels.size(0)correct(predictedlabels).sum().item()print(f 测试集准确率:{100*correct/total:.2f}%)几个关键细节torch.no_grad()推理时禁用梯度计算大幅减少显存占用torch.max(outputs, 1)dim1 表示在类别维度上取最大值返回 (values, indices)model.eval()通常也应该调切换到评估模式影响 Dropout 和 BatchNorm本项目没加但建议补上这个简单的 CNN 训 10 个 epoch 大约能达到70% ~ 75%的准确率作为基线模型已经不错了随机猜是 10%。️ 九、第七步随机抽取预测结果可视化dataiteriter(test_loader)images,labelsnext(dataiter)images,labelsimages.to(device),labels.to(device)outputsmodel(images)_,predictedtorch.max(outputs,1)imagesimages.cpu().numpy()labelslabels.cpu().numpy()predictedpredicted.cpu().numpy()plt.figure(figsize(12,6))classes[飞机,汽车,鸟,猫,鹿,狗,青蛙,马,船,卡车]foriinrange(5):plt.subplot(1,5,i1)imgimages[i].transpose((1,2,0))# (3,32,32)→(32,32,3)imgimg*np.array([0.2023,0.1994,0.2010])np.array([0.4914,0.4822,0.4465])imgnp.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()注意transposePyTorch 的图像张量是(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]Falsedevicetorch.device(cudaiftorch.cuda.is_available()elsecpu)print(f 使用设备:{device})# 数据加载transformtransforms.Compose([transforms.ToTensor(),transforms.Normalize((0.4914,0.4822,0.4465),(0.2023,0.1994,0.2010))])train_datasetdatasets.CIFAR10(root./data,trainTrue,downloadTrue,transformtransform)test_datasetdatasets.CIFAR10(root./data,trainFalse,downloadTrue,transformtransform)train_loadertorch.utils.data.DataLoader(train_dataset,batch_size64,shuffleTrue)test_loadertorch.utils.data.DataLoader(test_dataset,batch_size64,shuffleFalse)# CNN 模型classSimpleCNN(nn.Module):def__init__(self):super(SimpleCNN,self).__init__()self.conv1nn.Conv2d(3,32,3,padding1)self.poolnn.MaxPool2d(2,2)self.conv2nn.Conv2d(32,64,3,padding1)self.fc1nn.Linear(64*8*8,256)self.fc2nn.Linear(256,10)defforward(self,x):xself.pool(torch.relu(self.conv1(x)))xself.pool(torch.relu(self.conv2(x)))xx.view(-1,64*8*8)xtorch.relu(self.fc1(x))xself.fc2(x)returnx modelSimpleCNN().to(device)criterionnn.CrossEntropyLoss()optimizeroptim.Adam(model.parameters(),lr0.001)# 训练epochs10forepochinrange(epochs):running_loss0.0fori,(images,labels)inenumerate(train_loader):images,labelsimages.to(device),labels.to(device)outputsmodel(images)losscriterion(outputs,labels)optimizer.zero_grad()loss.backward()optimizer.step()running_lossloss.item()if(i1)%5000:print(f[{epoch1}/{epochs}] Step{i1}, Loss:{loss.item():.4f})print(fEpoch{epoch1}结束, 平均损失:{running_loss/len(train_loader):.4f})# 测试correcttotal0withtorch.no_grad():forimages,labelsintest_loader:images,labelsimages.to(device),labels.to(device)outputsmodel(images)_,predictedtorch.max(outputs.data,1)totallabels.size(0)correct(predictedlabels).sum().item()print(f 测试集准确率:{100*correct/total:.2f}%)# 可视化classes[飞机,汽车,鸟,猫,鹿,狗,青蛙,马,船,卡车]dataiteriter(test_loader)images,labelsnext(dataiter)images,labelsimages.to(device),labels.to(device)outputsmodel(images)_,predictedtorch.max(outputs,1)images,labels,predictedimages.cpu().numpy(),labels.cpu().numpy(),predicted.cpu().numpy()plt.figure(figsize(12,6))foriinrange(5):plt.subplot(1,5,i1)imgimages[i].transpose((1,2,0))imgimg*np.array([0.2023,0.1994,0.2010])np.array([0.4914,0.4822,0.4465])imgnp.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 每张图的尺寸是多少通道数呢padding1的作用是什么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)保存训练好的模型