在本篇文章中,我们将展示如何使用神经网络(ANN)对MNIST数据集进行手写数字识别。MNIST数据集包含28x28像素的手写数字图片(0到9),每张图片对应一个标签,表示图片中的数字是什么。本文将使用PyTorch库来实现一个简单的神经网络,通过多层感知器(MLP)模型对手写数字进行分类。

1. 加载MNIST数据集

        PyTorch 提供了非常方便的工具,可以直接加载常用的公开数据集。MNIST就是其中之一。我们可以通过 torchvision 库将数据下载并加载到程序中。

import torch
import torch.nn as nn
import torch.nn.functional as F          
from torch.utils.data import DataLoader  
from torchvision import datasets, transforms

import numpy as np
import pandas as pd
from sklearn.metrics import confusion_matrix  
import matplotlib.pyplot as plt
%matplotlib inline

        我们首先定义一个简单的 ToTensor 转换,用于将数据集中的图片转换为张量格式。

transform = transforms.ToTensor()

        接下来,加载训练集和测试集:

train_data = datasets.MNIST(root='../data', train=True, download=True, transform=transform)
test_data = datasets.MNIST(root='../data', train=False, download=True, transform=transform)

2. 数据加载与批处理

        数据集加载完成后,我们可以使用 DataLoader 来创建批量数据,这在训练大规模数据时非常重要。

torch.manual_seed(101)  # 为了结果一致性
train_loader = DataLoader(train_data, batch_size=100, shuffle=True)
test_loader  = DataLoader(test_data, batch_size=500, shuffle=False)

        使用 DataLoader 可以帮助我们按批次加载数据,确保每次训练过程中模型只需处理一定数量的数据,节省内存并提高训练效率。

3. 构建神经网络模型

        我们将构建一个多层感知器模型(MLP),其输入层为784个节点(28x28像素展开),输出层为10个节点(对应0到9的10个分类)。模型的隐藏层包含两层,分别为120个和84个节点。

class MultilayerPerceptron(nn.Module):
    def __init__(self, in_sz=784, out_sz=10, layers=[120,84]):
        super().__init__()
        self.fc1 = nn.Linear(in_sz,layers[0])
        self.fc2 = nn.Linear(layers[0],layers[1])
        self.fc3 = nn.Linear(layers[1],out_sz)
    
    def forward(self,X):
        X = F.relu(self.fc1(X))
        X = F.relu(self.fc2(X))
        X = self.fc3(X)
        return F.log_softmax(X, dim=1) 

4. 定义损失函数与优化器

        我们将使用交叉熵损失函数 CrossEntropyLoss,并采用 Adam 优化器来更新模型参数。

criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

5. 训练模型

        接下来,我们开始训练模型。训练过程中,我们将每个批次的训练数据输入模型,并计算损失。每经过一个 epoch 后,我们还将测试模型在测试集上的表现。

epochs = 10
train_losses = []
test_losses = []
train_correct = []
test_correct = []

for i in range(epochs):
    trn_corr = 0
    tst_corr = 0
    
    # 训练过程
    for b, (X_train, y_train) in enumerate(train_loader):
        b+=1
        
        # 输入模型并计算损失
        y_pred = model(X_train.view(100, -1))  
        loss = criterion(y_pred, y_train)

        # 统计正确分类的数量
        predicted = torch.max(y_pred.data, 1)[1]
        batch_corr = (predicted == y_train).sum()
        trn_corr += batch_corr
        
        # 更新参数
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
    
    # 测试过程
    with torch.no_grad():
        for b, (X_test, y_test) in enumerate(test_loader):
            y_val = model(X_test.view(500, -1))
            predicted = torch.max(y_val.data, 1)[1]
            tst_corr += (predicted == y_test).sum()

6. 评估模型性能

        我们可以通过绘制损失曲线和准确率曲线来评估模型的性能。

train_losses = [loss.item() for loss in train_losses]
test_losses  = [loss.item() for loss in test_losses]

plt.plot(train_losses, label='training loss')
plt.plot(test_losses, label='validation loss')
plt.title('Loss at the end of each epoch')
plt.legend();

        最后,我们计算在测试集上的准确率:

with torch.no_grad():
    correct = 0
    for X_test, y_test in test_load_all:
        y_val = model(X_test.view(len(X_test), -1))
        predicted = torch.max(y_val,1)[1]
        correct += (predicted == y_test).sum()
print(f'Test accuracy: {correct.item()}/{len(test_data)} = {correct.item()*100/(len(test_data)):7.3f}%')

7. 错误分析

        我们可以查看模型分类错误的样本,帮助我们进一步分析模型的不足之处。

misses = np.array([])
for i in range(len(predicted.view(-1))):
    if predicted[i] != y_test[i]:
        misses = np.append(misses,i).astype('int64')

# 打印错误分类的样本
for idx in misses[:10]:
    print(f'Index: {idx}, Label: {y_test[idx]}, Prediction: {predicted[idx]}')

结语

        在这篇文章中,我们使用 PyTorch 实现了一个简单的神经网络模型,并对MNIST手写数字进行了分类。通过多层感知器(MLP),我们可以达到约98%的分类准确率。这证明了即使在没有卷积层的情况下,简单的全连接层也能够较好地解决手写数字识别问题。

        下一步可以尝试进一步优化模型,例如加入卷积层、池化层,或者调整模型的超参数,以提高准确率。

如果你觉得这篇博文对你有帮助,请点赞、收藏、关注我,并且可以打赏支持我!

欢迎关注我的后续博文,我将分享更多关于人工智能、自然语言处理和计算机视觉的精彩内容。

谢谢大家的支持!

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐