入门实践工程四:基于 PyTorch 的手写数字识别(MNIST 图像分类)|附:环境依赖环境及工程源码

🔥 别再死磕枯燥理论了!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 日志

工程源码

main.py

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()
相关推荐
A15362551 小时前
2026 电商物流系统推荐:多渠道仓配履约数字化选型指南
大数据·数据库·人工智能
SEO_juper1 小时前
2026年Schema自动注入实战:用Python批量给1000个页面加上JSON-LD,AI引用率实测提升38%
人工智能·python·json·seo·独立站
雪的季节1 小时前
不安装YOLO只安装 PyTorch,加载已有yolo数据,从无到有创建模型训练数据并加载使用(重要)
人工智能·pytorch·yolo
超智算科技1 小时前
超智算领衔协办“NVIDIA创业企业展示·西安站”,全栈AI服务赋能行业生态协同发展
大数据·网络·人工智能·物联网·百度
云端漫步19872 小时前
HarmonyOS NEXT AI 智能生活助手:统一 AIService 封装
人工智能·华为·生活·harmonyos
风途科技~2 小时前
FMCW 调频连续波雷达|非接触式雷达水位计精准把控液位变化
大数据·人工智能
程序员AI工坊2 小时前
Agent 开发:ReAct 循环与工具调用实战——从单次调用到自主 Agent
人工智能·后端·python·langchain·agent·react
小刘快学习2 小时前
品牌舆情监测的接入思路:从采集到预警的链路
人工智能
Ai-_Man2 小时前
豆包智能体内容批量导出:哪些该存、存成什么、怎么存v
人工智能·ai·小程序·word