当前位置:首页 > 技术 > 正文内容

PyTorch 卷积神经网络实现 MNIST 手写数字识别

访客 技术 2026年9月3日 1

环境配置

在开始项目之前,建议创建一个独立的 Conda 虚拟环境并安装必要的深度学习库:

conda create -n torch_cv python=3.10
conda activate torch_cv
conda install pytorch torchvision torchaudio pytorch-cuda=12.1 -c pytorch -c nvidia
pip install matplotlib tqdm pillow

原始数据处理

MNIST 数据集通常以二进制格式存储。为了更直观地处理数据并模拟真实业务场景,我们先将二进制文件解析为图片格式(PNG),并生成对应的标注文件。

import os
import struct
import numpy as np
from array import array
from PIL import Image

class MNISTRawLoader:
    def __init__(self, train_set, test_set):
        self.paths = {
            'train': train_set,
            'test': test_set
        }

    def _parse(self, img_path, lbl_path):
        with open(lbl_path, 'rb') as f:
            _, size = struct.unpack(">II", f.read(8))
            labels = array("B", f.read())
        
        with open(img_path, 'rb') as f:
            _, size, rows, cols = struct.unpack(">IIII", f.read(16))
            raw_data = array("B", f.read())
        
        images = []
        for i in range(size):
            img = np.array(raw_data[i * rows * cols : (i + 1) * rows * cols]).reshape(28, 28)
            images.append(img)
        return images, labels

    def export_assets(self, target_dir, subset='train'):
        img_p, lbl_p = self.paths[subset]
        images, labels = self._parse(img_p, lbl_p)
        
        output_path = os.path.join(target_dir, subset)
        os.makedirs(output_path, exist_ok=True)
        
        manifest = []
        for idx, (img_data, val) in enumerate(zip(images, labels)):
            file_name = f"{subset}_{idx:05d}_{val}.png"
            Image.fromarray(img_data).save(os.path.join(output_path, file_name))
            manifest.append(f"{subset}/{file_name}\t{val}")
            
        with open(os.path.join(target_dir, f"{subset}_labels.txt"), "w") as f:
            f.write("\n".join(manifest))

# 示例调用 (假设文件已下载至 ./data)
# loader = MNISTRawLoader(
#     train_set=('./data/train-images-idx3-ubyte', './data/train-labels-idx1-ubyte'),
#     test_set=('./data/t10k-images-idx3-ubyte', './data/t10k-labels-idx1-ubyte')
# )
# loader.export_assets('./mnist_data', 'train')
# loader.export_assets('./mnist_data', 'test')

自定义 Dataset 与数据预处理

在 PyTorch 中,通过继承 Dataset 类可以灵活地读取自定义格式的数据。在训练前,我们需要计算训练集的均值(Mean)和标准差(Std)用于归一化,这有助于加速模型收敛并避免梯度问题。

import torch
from torch.utils.data import Dataset, DataLoader, random_split
from torchvision import transforms

class DigitsDataset(Dataset):
    def __init__(self, root, label_file, transform=None):
        self.root = root
        self.transform = transform
        with open(label_file, 'r') as f:
            self.items = [line.strip().split('\t') for line in f.readlines()]

    def __len__(self):
        return len(self.items)

    def __getitem__(self, index):
        rel_path, label = self.items[index]
        full_path = os.path.join(self.root, rel_path)
        img = Image.open(full_path).convert('L')
        
        if self.transform:
            img = self.transform(img)
        return img, int(label)

# 计算统计量
temp_ds = DigitsDataset('./mnist_data', './mnist_data/train_labels.txt', transform=transforms.ToTensor())
loader = DataLoader(temp_ds, batch_size=1024)

def compute_stats(loader):
    data_sum, data_sq_sum, num_batches = 0, 0, 0
    for data, _ in loader:
        data_sum += torch.mean(data)
        data_sq_sum += torch.mean(data**2)
        num_batches += 1
    mean = data_sum / num_batches
    std = (data_sq_sum / num_batches - mean**2)**0.5
    return mean.item(), std.item()

m, s = compute_stats(loader)
print(f"Dataset Mean: {m:.4f}, Std: {s:.4f}")

数据增强与加载

为了提高模型的泛化能力,我们在训练集中加入随机旋转和裁剪。

train_aug = transforms.Compose([
    transforms.RandomRotation(10),
    transforms.RandomResizedCrop(28, scale=(0.9, 1.1)),
    transforms.ToTensor(),
    transforms.Normalize((m,), (s,))
])

test_aug = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((m,), (s,))
])

full_train_ds = DigitsDataset('./mnist_data', './mnist_data/train_labels.txt', transform=train_aug)
test_ds = DigitsDataset('./mnist_data', './mnist_data/test_labels.txt', transform=test_aug)

# 划分验证集
train_size = int(0.9 * len(full_train_ds))
val_size = len(full_train_ds) - train_size
train_ds, val_ds = random_split(full_train_ds, [train_size, val_size])

train_loader = DataLoader(train_ds, batch_size=64, shuffle=True)
val_loader = DataLoader(val_ds, batch_size=64)
test_loader = DataLoader(test_ds, batch_size=64)

构建卷积神经网络 (CNN)

我们设计一个经典的卷积架构,包含三层卷积和两层全连接,用于提取手写数字的特征。

import torch.nn as nn

class RecognitionNet(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        self.feature_extractor = nn.Sequential(
            nn.Conv2d(1, 16, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(16, 32, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.AdaptiveAvgPool2d((4, 4))
        )
        self.head = nn.Sequential(
            nn.Flatten(),
            nn.Linear(64 * 4 * 4, 128),
            nn.ReLU(),
            nn.Dropout(0.2),
            nn.Linear(128, num_classes)
        )

    def forward(self, x):
        features = self.feature_extractor(x)
        logits = self.head(features)
        return logits

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = RecognitionNet().to(device)

模型训练与验证

使用 Adam 优化器和交叉熵损失函数进行模型优化。

import torch.optim as optim
from tqdm import tqdm

optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

def execute_epoch(model, loader, opt, crit, is_train=True):
    model.train() if is_train else model.eval()
    total_loss, correct = 0, 0
    
    with torch.set_grad_enabled(is_train):
        for imgs, lbls in loader:
            imgs, lbls = imgs.to(device), lbls.to(device)
            
            outputs = model(imgs)
            loss = crit(outputs, lbls)
            
            if is_train:
                opt.zero_grad()
                loss.backward()
                opt.step()
            
            total_loss += loss.item() * imgs.size(0)
            preds = outputs.argmax(dim=1)
            correct += (preds == lbls).sum().item()
            
    return total_loss / len(loader.dataset), correct / len(loader.dataset)

# 训练循环
epochs = 10
best_acc = 0

for epoch in range(epochs):
    t_loss, t_acc = execute_epoch(model, train_loader, optimizer, criterion)
    v_loss, v_acc = execute_epoch(model, val_loader, optimizer, criterion, is_train=False)
    
    if v_acc > best_acc:
        best_acc = v_acc
        torch.save(model.state_dict(), 'best_model.pth')
        
    print(f"Epoch {epoch+1:02d}: Train Acc {t_acc:.4f} | Val Acc {v_acc:.4f}")

模型在经过几个 Epoch 的训练后,通常在验证集上能达到 98% 以上的准确率。最终可以在测试集上载入最佳权重进行最后的评估。

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Dom\HTML_NO_DEFAULT_NS 的副作用:自动加闭合标签

在使用Dom\HTMLDocument时,Dom\HTML_NO_DEFAULT_NS 将禁止在解析过程中设置元素的命名空间, 此设置是为了与DOMDocument向后兼容而存在的。当使用它时,已知的一个副作用就是:自动加闭合标签例如 </img> 为什么会这样?当你使用:Dom\HTML_NO_DEFAULT_NS文档会变成 无命名空间模式,此时内部更接近 XML...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。