机器学习学习总结 + 代码实现(MNIST 手写数字识别)
·
一、学习过程总结
-
明确任务目标本次学习的核心任务是使用卷积神经网络(CNN)实现手写数字自动识别,输入 28×28 的灰度图片,输出 0-9 的数字分类结果,掌握图像分类的完整流程。
-
学习核心知识点
- 数据集:使用经典的 MNIST 手写数字数据集,包含 6 万张训练图、1 万张测试图。
- 数据预处理:将图片转为张量、归一化、标准化,让模型更容易训练。
- 模型结构:学习 CNN 核心组件 ——卷积层(提取特征)、池化层(压缩特征)、全连接层(分类)。
- 训练流程:前向传播计算损失→反向传播更新参数→迭代优化模型。
- 评估方法:用测试集计算准确率,验证模型识别效果。
- 遇到的问题与解决
- 不理解 CNN 工作原理:通过分层观察特征提取过程理解。
- 代码运行报错:检查库版本、设备(CPU/GPU)适配、数据维度。
- 准确率偏低:调整学习率、训练轮数、网络层数解决。
- 学习收获
- 掌握了机器学习项目的完整开发流程:数据→模型→训练→评估→预测。
- 理解了 CNN 在图像任务中的优势,能独立搭建简单图像分类模型。
- 学会使用 PyTorch 框架实现深度学习模型,具备基础代码调试能力。
机器学习案例教程:基于 CNN 的 MNIST 手写数字识别
一、设计思路
本案例以 MNIST 手写数字数据集为对象,构建卷积神经网络(CNN)实现 0-9 数字的自动识别。
- 问题定义:输入一张 28×28 像素的灰度手写数字图片,输出该图片对应的数字类别(0-9)。
- 模型选择:使用卷积神经网络(CNN),相比传统全连接网络,CNN 能自动提取图像的边缘、纹理等空间特征,在图像分类任务中表现更优。
- 实现流程:数据加载与预处理 → 模型构建 → 模型训练 → 模型评估 → 单张图片预测演示。
二、AI 工具使用方法
本教程使用 ChatGPT + Python 实现:
- 向 AI 工具发送需求:“帮我写一个基于 PyTorch 的 CNN 实现 MNIST 手写数字识别的完整代码,包含数据加载、模型构建、训练、评估和预测演示,附带详细注释”。
- 对 AI 生成的代码进行微调,适配本地环境(如调整 batch_size、学习率等超参数)。
- 运行代码并记录运行结果,补充教程说明。
三、完整代码实现
python
运行
# 导入必要库
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
# 1. 数据预处理与加载
transform = transforms.Compose([
transforms.ToTensor(), # 转为Tensor并归一化到[0,1]
transforms.Normalize((0.1307,), (0.3081,)) # 用MNIST数据集的均值和标准差标准化
])
# 加载MNIST数据集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)
# 2. 构建CNN模型
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
# 卷积层
self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1) # 输入通道1,输出通道32
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
# 池化层
self.pool = nn.MaxPool2d(2, 2)
# 全连接层
self.fc1 = nn.Linear(64 * 7 * 7, 128)
self.fc2 = nn.Linear(128, 10) # 输出10个类别(0-9)
# 激活函数
self.relu = nn.ReLU()
def forward(self, x):
x = self.pool(self.relu(self.conv1(x))) # 28x28 -> 14x14
x = self.pool(self.relu(self.conv2(x))) # 14x14 -> 7x7
x = x.view(-1, 64 * 7 * 7) # 展平
x = self.relu(self.fc1(x))
x = self.fc2(x)
return x
# 初始化模型、损失函数、优化器
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 3. 模型训练
def train(model, train_loader, criterion, optimizer, device, epochs=5):
model.train()
for epoch in range(epochs):
running_loss = 0.0
for images, labels in 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()
print(f"Epoch {epoch+1}, Loss: {running_loss/len(train_loader):.4f}")
train(model, train_loader, criterion, optimizer, device)
# 4. 模型评估
def test(model, test_loader, device):
model.eval()
correct = 0
total = 0
with torch.no_grad():
for images, labels in test_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"Test Accuracy: {100 * correct / total:.2f}%")
test(model, test_loader, device)
# 5. 单张图片预测演示
def predict_single_image(model, test_dataset, device, index=0):
model.eval()
image, label = test_dataset[index]
image = image.unsqueeze(0).to(device) # 增加batch维度
with torch.no_grad():
output = model(image)
_, predicted = torch.max(output, 1)
# 显示图片与预测结果
plt.imshow(image.squeeze().cpu().numpy(), cmap='gray')
plt.title(f"True Label: {label}, Predicted: {predicted.item()}")
plt.axis('off')
plt.show()
# 预测第10张图片
predict_single_image(model, test_dataset, device, index=10)
四、运行结果说明
- 训练过程输出:
plaintext
损失值随训练轮次下降,说明模型在不断学习。Epoch 1, Loss: 0.1532 Epoch 2, Loss: 0.0487 Epoch 3, Loss: 0.0332 Epoch 4, Loss: 0.0241 Epoch 5, Loss: 0.0185 - 测试集准确率:
plaintext
模型在测试集上达到了 99% 以上的准确率,识别效果优秀。Test Accuracy: 99.12% - 单张预测可视化:运行代码后会弹出窗口,显示手写数字图片,并标注真实标签和模型预测结果,两者基本一致。

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



所有评论(0)