🔥 别再死磕枯燥理论了!AI 时代,拿实战作品说话才是硬道理!(原创不易哈,希望可以帮到还有些许学习劲儿的同学们) 🔥
【进阶版还在创作中,耗费精力中......】
入门实践工程四:基于 PyTorch 的手写数字识别(MNIST 图像分类)|附:环境依赖环境及工程源码
简介
使用一个轻量卷积神经网络(CNN)在经典 MNIST 数据集上完成 0-9 手写数字识别,涵盖数据加载、模型定义、训练、测试与预测对照全流程,是理解深度学习「数据→模型→损失→优化」闭环的最佳起点。
工程详细介绍
核心思想: 卷积神经网络通过局部感受野、权值共享与下采样,自动从像素中学习「边缘→纹理→部件→类别」的层次特征,是深度学习视觉任务的基石,完整体现「数据→模型→损失→优化」训练闭环。
实现方法:
- 数据: MNIST(torchvision 自动下载,6 万训练 + 1 万测试,28×28 灰度图)。
- 模型: 两层卷积块(Conv+ReLU+MaxPool)把 28×28 压到 7×7×32,展平后两层全连接输出 10 类 logits。
- 流程: 交叉熵损失 + Adam 优化器训练 3 轮,每轮打印训练准确率;测试集评估;取一批打印预测与真实标签对照;保存权重供后续项目复用。
- 输出: 98%+ 测试准确率 + 预测对照 + mnist_cnn.pth(供入门实践工程 10 复用)。
目录结构
04_mnist_cnn/
├── main.py # 训练 + 测试 + 预测主程序
├── requirements.txt
└── data/ # 首次运行自动下载 MNIST 数据
正确安装
bash
pip install torch torchvision numpy
python main.py
- 首次运行自动下载 MNIST 数据集(约 10MB),需联网。
- 如需 GPU,按 PyTorch 官网选择对应 CUDA 版本安装 torch。
运行方式
bash
pip install -r requirements.txt
python main.py
预期结果
- 3 个 Epoch 后测试准确率约 98%+(CPU 几分钟内完成)
- 末尾打印前 8 张测试图片的预测对照
- 权重保存为
mnist_cnn.pth(供入门实践工程 10 复用)
扩展方向
- 替换为 ResNet18 + CIFAR-10,过渡到完整图像分类项目
- 增加数据增强、学习率调度、TensorBoard 日志
工程源码
python
"""
入门实践工程四:基于 PyTorch 的手写数字识别(MNIST 图像分类)
=====================================================
使用一个轻量卷积神经网络(CNN)在 MNIST 数据集上完成 0-9 手写数字识别,
涵盖:数据加载 -> 模型定义 -> 训练 -> 测试 -> 预测可视化 全流程。
运行:
python main.py
数据集会在首次运行时由 torchvision 自动下载到 ./data 目录。
"""
import os
import random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
# ---------------- 全局配置 ----------------
SEED = 42
BATCH_SIZE = 64
EPOCHS = 3 # 入门演示用,CPU 也能几分钟跑完
LR = 1e-3
DATA_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data")
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def set_seed(seed: int = 42):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
# ---------------- 模型定义 ----------------
class MnistCNN(nn.Module):
"""轻量 CNN:两个卷积块 + 全连接分类头。"""
def __init__(self, num_classes: int = 10):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
self.fc1 = nn.Linear(32 * 7 * 7, 128)
self.fc2 = nn.Linear(128, num_classes)
def forward(self, x):
x = F.relu(self.conv1(x)) # [B,16,28,28]
x = F.max_pool2d(x, 2) # [B,16,14,14]
x = F.relu(self.conv2(x)) # [B,32,14,14]
x = F.max_pool2d(x, 2) # [B,32,7,7]
x = x.flatten(1) # [B,32*7*7]
x = F.relu(self.fc1(x))
return self.fc2(x) # [B,10] logits
# ---------------- 数据加载 ----------------
def get_loaders():
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)), # MNIST 均值/标准差
])
train_set = datasets.MNIST(DATA_DIR, train=True, download=True, transform=transform)
test_set = datasets.MNIST(DATA_DIR, train=False, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)
test_loader = DataLoader(test_set, batch_size=BATCH_SIZE, shuffle=False)
return train_loader, test_loader
# ---------------- 训练 & 测试 ----------------
def train(model, loader, optimizer, epoch):
model.train()
total, correct, loss_sum = 0, 0, 0.0
for batch_idx, (x, y) in enumerate(loader):
x, y = x.to(DEVICE), y.to(DEVICE)
optimizer.zero_grad()
out = model(x)
loss = F.cross_entropy(out, y)
loss.backward()
optimizer.step()
loss_sum += loss.item() * x.size(0)
pred = out.argmax(1)
correct += (pred == y).sum().item()
total += x.size(0)
if batch_idx % 100 == 0:
print(f" [Epoch {epoch}] batch {batch_idx}/{len(loader)} loss={loss.item():.4f}")
print(f" [Epoch {epoch}] 训练准确率: {correct/total:.4f} 平均损失: {loss_sum/total:.4f}")
@torch.no_grad()
def evaluate(model, loader):
model.eval()
correct, total = 0, 0
for x, y in loader:
x, y = x.to(DEVICE), y.to(DEVICE)
out = model(x)
correct += (out.argmax(1) == y).sum().item()
total += x.size(0)
acc = correct / total
print(f"==> 测试集准确率: {acc:.4f} ({correct}/{total})")
return acc
# ---------------- 预测可视化 ----------------
@torch.no_grad()
def predict_and_show(model, loader, n: int = 8):
"""取一批数据,打印模型预测与真实标签的对照。"""
model.eval()
x, y = next(iter(loader))
x, y = x.to(DEVICE), y.to(DEVICE)
out = model(x[:n])
preds = out.argmax(1).cpu().tolist()
truths = y[:n].cpu().tolist()
print("\n前 %d 张图片预测对照:" % n)
for i in range(n):
flag = "✓" if preds[i] == truths[i] else "✗"
print(f" 第 {i+1} 张: 预测={preds[i]} 真实={truths[i]} {flag}")
def main():
set_seed(SEED)
print(f"设备: {DEVICE}")
train_loader, test_loader = get_loaders()
print(f"训练集样本数: {len(train_loader.dataset)} 测试集样本数: {len(test_loader.dataset)}")
model = MnistCNN().to(DEVICE)
optimizer = optim.Adam(model.parameters(), lr=LR)
for epoch in range(1, EPOCHS + 1):
print(f"\n--- Epoch {epoch}/{EPOCHS} ---")
train(model, train_loader, optimizer, epoch)
evaluate(model, test_loader)
predict_and_show(model, test_loader)
save_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mnist_cnn.pth")
torch.save(model.state_dict(), save_path)
print(f"\n模型权重已保存到: {save_path}")
if __name__ == "__main__":
main()