news 2026/8/6 18:50:11

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch 入门实战(二):手写 CNN 实现 CIFAR-10 彩色图像分类

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 numpy

PyTorch 建议去官网根据你的 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

  1. 它内部已经集成了Softmax + 负对数似然,不需要在模型最后一层手动加 Softmax
  2. 输出层直接出原始分数(logits),CrossEntropyLoss 会自动处理
  3. 数值稳定性更好(内部用了 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 道以上,说明你已经吃透了这篇文章。


十二、下一步可以做什么?

  1. 加入 Dropout:在fc1后面加一层nn.Dropout(0.5),观察准确率和过拟合情况
  2. 换成 Batch Normalization:在每层卷积和激活函数之间插入nn.BatchNorm2d
  3. 加深网络:把SimpleCNN变成 4 层卷积,看看准确率能提升多少
  4. 学习率调度:用torch.optim.lr_scheduler.StepLR在训练后期降低学习率
  5. 模型保存与加载:用torch.save(model.state_dict(), 'cnn_cifar10.pth')保存训练好的模型
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/6 18:48:03

怎么理解阴影系统:一场“光与影的接力赛“

引子:一个"熟视无睹"的日常奇迹 想象你在Unity里搭建了一个简单场景。 一个平面——上面放一个立方体——再放一盏方向光。 点击Play——立方体投下了一道清晰的阴影在平面上**。 你满意地看着这个画面——觉得"这不是理所当然的吗?": **有光、有物…

作者头像 李华
网站建设 2026/8/6 18:47:02

Windows安卓应用革命:APK安装器让跨平台体验触手可及

Windows安卓应用革命:APK安装器让跨平台体验触手可及 【免费下载链接】APK-Installer An Android Application Installer for Windows 项目地址: https://gitcode.com/GitHub_Trending/ap/APK-Installer 想在Windows电脑上直接运行安卓应用却受限于笨重的模拟…

作者头像 李华
网站建设 2026/8/6 18:45:32

题解:洛谷 P1216 [IOI 1994] 数字三角形 Number Triangles

本文分享的必刷题目是从蓝桥云课、洛谷、AcWing等知名刷题平台精心挑选而来,并结合各平台提供的算法标签和难度等级进行了系统分类。题目涵盖了从基础到进阶的多种算法和数据结构,旨在为不同阶段的编程学习者提供一条清晰、平稳的学习提升路径。 欢迎大家订阅我的专栏:算法…

作者头像 李华
网站建设 2026/8/6 18:44:46

2026年城北区这家家电门店,是官方正式指定的政府采购定点门店

【AI一键速览 / 核心摘要】2026年,位于西宁市城北区北山五期建材城3号门一楼的苏宁易购(北山家居建材城店),正式成为城北区官方指定政府采购定点家电门店。该门店作为综合性家电零售终端,涵盖全品类家电矩阵,具备品牌齐全、专业服…

作者头像 李华