从零训练手写数字识别模型
本篇用 MNIST 手写数字走完一遍 PyTorch 训练闭环:装环境、加载数据、搭小 CNN、训练与评估、存权重、推理和交付。读完应能自己跑通从数据到 .pt 的流程,并知道自定义数据、迁移学习和灾难性遗忘要注意什么。
AI 模型训练是用大量标注数据,反复迭代,自动调整模型内部权重参数,让模型预测越来越准的完整过程。
核心步骤(PyTorch 通用流程):
- 准备数据:整理成统一格式。完整做法会分成训练集、验证集、测试集三类:训练集改权重,验证集调超参、看是否过拟合,测试集只在最后评估一次。本篇 MNIST 示例只用官方的训练、测试划分,没有再切验证集。
- 搭建模型结构:用卷积、全连接等层搭出骨架。此时参数是随机数,还不会这个任务。
- 前向传播:数据送进模型得到预测,和真实标签比,算出损失(loss)。损失越大越不准。
- 反向传播 + 参数更新:
- 自动求导(Autograd):算出每个参数该往哪边改才能降低损失。
- 优化器(SGD、Adam):按梯度小幅改参数。
- 全部训练样本过完一轮叫一个 epoch,通常要重复多轮。
循环「前向 → 反向 → 更新」,直到损失不再明显下降、效果稳定,训练结束。
PyTorch
PyTorch 是 Meta(原 Facebook)开源的 Python 深度学习框架,主打动态图、易上手、科研友好,广泛用于搭建、训练、推理神经网络,覆盖视觉、语言、语音、强化学习等。
两大核心模块:
torch:张量计算(类似 NumPy,可在 CPU、CUDA GPU、Apple MPS 上加速)。torch.nn:层、损失函数、优化器等高层 API。
一切运算的载体是 Tensor(张量):带自动求导、可放在不同设备上的多维数组。
适用场景:科研实验(动态图改结构方便)、大模型与多模态、工业里的检测、分类、语音、推荐、以及语法贴近 Python 的入门。JAX 等框架在部分论文里也很常见,并不是只有 PyTorch。
安装工具库
# CPU 版;若有 NVIDIA 显卡,按官网 https://pytorch.org 选对应 CUDA 的安装命令
pip install torch torchvision matplotlib
import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader
import matplotlib
import matplotlib.pyplot as plt
torch
深度学习的核心是大量矩阵运算(加法、乘法、求导)。Python 本身做这些很慢,所以 PyTorch 用 C++ 写底层引擎,再封装成 Python 包 torch,用 Python 语法调用底层算力。
torch 既是包名,也是核心库,围着 torch.Tensor 转,支持动态计算图,跑在 CPU、CUDA、MPS 上。
torch.Tensor:张量(Tensor)在数学中是多维数组,只存数字。torch.Tensor是 PyTorch 最基本的数据结构,是承载数据、参与运算、支持自动微分、可跨设备(CPU、GPU)计算的基础单元,是整个 PyTorch 深度学习框架的底层基石。torch.nn:神经网络层(Linear、Conv2d、LSTM 等)。torch.optim:优化器(SGD、Adam 等)。torch.nn.functional:激活函数(relu、softmax 等)。torch.utils.data:数据加载工具(Dataset、DataLoader)。
torchvision
PyTorch 的计算机视觉(CV)专属工具箱,提供数据集、预训练模型、图像变换。
做图像相关的 AI 任务时,有大量重复性工作:
- 下载公开数据集(MNIST、CIFAR-10、ImageNet)。
- 对图像做预处理(裁剪、缩放、归一化)。
- 用预训练好的模型(ResNet、VGG)做迁移学习。
torchvision 将这些都封装:
torchvision.datasets:内置常用数据集。MNIST、CIFAR-10 可一行下载;ImageNet 体量大,需要自行准备,不能当成「一行自动下完」。torchvision.transforms:图像预处理工具箱(裁剪、缩放、归一化)。torchvision.models:预训练好的经典模型(可直接用)。torchvision.io:读写图像、视频文件。torchvision.ops:计算机视觉专用算子(如 NMS、RoIAlign)。
matplotlib
Python 最流行的数据可视化库,用来画图、看数据分布、监控训练过程。
深度学习需要大量可视化:
- 训练过程中,画损失曲线(loss curve),看模型有没有在收敛。
- 查看输入数据长什么样(手写数字图片、彩色猫狗图)。
- 展示模型预测结果,对比预测 vs 真实标签。
- 画混淆矩阵,看模型在哪些类别上犯错。
核心用法:
plt.figure():创建一个画布(窗口)。plt.subplot()或 plt.subplots():在画布上划分多个子图区域。ax.imshow()或 ax.plot():在某个子图上画图。plt.show():渲染并显示整个画布。
加载数据集(MNIST)
MNIST(Modified National Institute of Standards and Technology)是深度学习领域最经典的手写数字灰度图像数据集,专门用来识别 0~9 手写数字,常作为深度学习入门 HelloWorld 数据集。
数据集总共 70000 张图片,其中训练集 60000 张,测试集 10000 张。每张图像都是单通道灰度图(黑白),28 × 28 像素,就是普通人手写的单个阿拉伯数字。
# 数据预处理:转成张量并缩放到 [0, 1],再标准化到 [-1.0, 1.0]
to_tensor = transforms.ToTensor()
normalize = transforms.Normalize((0.5,), (0.5,))
transform = transforms.Compose([
to_tensor,
normalize
])
train_dataset = torchvision.datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
test_dataset = torchvision.datasets.MNIST(
root='./data',
train=False,
download=True,
transform=transform
)
# 分批加载,每批 64 张
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)
print(f"训练集大小:{len(train_dataset)} 张")
print(f"测试集大小:{len(test_dataset)} 张")
└── data
│ └── MNIST
│ │ └── raw
│ │ │ └── t10k-images-idx3-ubyte
│ │ │ └── t10k-images-idx3-ubyte.gz
│ │ │ └── t10k-labels-idx1-ubyte
│ │ │ └── t10k-labels-idx1-ubyte.gz
│ │ │ └── train-images-idx3-ubyte
│ │ │ └── train-images-idx3-ubyte.gz
│ │ │ └── train-labels-idx1-ubyte
│ │ │ └── train-labels-idx1-ubyte.gz
transforms.Compose
创建了一个数据预处理流水线(Pipeline)。当图片经过这条流水线时,会依次执行两个操作:
ToTensor():把图片(PIL Image 或 NumPy 数组)转成 PyTorch 张量(Tensor),并把像素值从[0, 255]缩放到[0.0, 1.0]。Normalize((0.5,), (0.5,)):按 $(x-0.5)/0.5$ 做标准化,把[0.0, 1.0]映射到[-1.0, 1.0]。元组只有一个数,因为 MNIST 是单通道。
Compose 就像一个传送带,图片从一端进去,依次经过每个工位,从另一端出来时就已经处理好了。
import torchvision.transforms as transforms
import numpy as np
to_tensor = transforms.ToTensor()
normalize = transforms.Normalize((0.5,), (0.5,))
transform = transforms.Compose([
to_tensor,
normalize
])
# 创建一个 8x8 的二维数组,数值范围 0~255
# 这里我们造一个“半黑半白”的图案:左半全黑(0),右半全白(255)
manual_image = np.array([
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
[0, 0, 0, 0, 255, 255, 255, 255],
], dtype=np.uint8)
tensor_img = to_tensor(manual_image)
norm_img1 = normalize(tensor_img)
norm_img2 = transform(manual_image)
print(norm_img1)
print(norm_img2)
torchvision.datasets.MNIST
检查本地是否有 MNIST 数据,如果没有,从互联网自动下载,并绑定预处理规则(transform=transform)——但此时数据还没有真正加载到内存,当访问某个具体的数据项时,才会真正加载。
import torchvision
import torchvision.transforms as transforms
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
train_dataset = torchvision.datasets.MNIST(
root='./data',
train=True,
download=True, # 如果已存在,不会重复下载
transform=transform
)
print(f"数据量:{len(train_dataset)} 张")
image, label = train_dataset[0]
print(f"加载完成!")
print(f"图片张量 shape:{image.shape}")
print(f"标签:{label}")
print(f"像素值范围:[{image.min():.2f}, {image.max():.2f}]")
注意: MNIST 数据集里的文件并不是一张张独立的图片文件,其实是把成千上万张图片和标签,分别打包成了独立的 二进制数据文件。
DataLoader
torch.utils.data.DataLoader 是 PyTorch 用来构建数据迭代器的工具,接收 Dataset,自动批量打包、打乱数据,在训练循环中源源不断产出批量张量。
三大核心作用:
- 分批打包:模型训练时,通常不会一次只喂一张图(太慢),也不会一次把 6 万张全塞进去(内存不够)。指定
batch_size=64就会自动帮每次取出 64 张图片和对应的 64 个标签,组合成一批供模型处理。 - 数据打乱:训练时,希望模型看到的是随机打乱的数据,而不是所有 0 排在一起、所有 1 排在一起。
shuffle=True会在每次遍历数据前,自动帮你把 6 万个样本的顺序打乱。 - 并行加载:数据读取往往是训练过程中的一个瓶颈。
num_workers参数可以开启多个子进程,让电脑在 GPU 计算当前批次的同时,提前把下一批数据从硬盘加载到内存里,大幅提升 GPU 利用率。
自己按步长切片组 batch 既慢又容易出错,下面只是对照 DataLoader 在做什么,训练里不要这么写:
for i in range(0, len(dataset), 64):
batch_images = []
看一眼数据长什么样
DataLoader 返回一个 DataLoader 类实例,本质是一个可迭代对象(Iterable)。
# utils/plot.py
import matplotlib
import matplotlib.pyplot as plt
# 中文字体:macOS 用下列字体;Windows 可改成 ['Microsoft YaHei', 'SimHei']
matplotlib.rcParams['font.sans-serif'] = ['Arial Unicode MS', 'Heiti TC', 'PingFang SC']
matplotlib.rcParams['axes.unicode_minus'] = False
def plot_images(train_loader):
images, labels = next(iter(train_loader))
fig, axes = plt.subplots(2, 4, figsize=(10, 6))
for i, ax in enumerate(axes.flat):
ax.imshow(images[i][0], cmap='gray')
ax.set_title(f"标签:{labels[i].item()}")
ax.axis('off')
plt.tight_layout()
plt.show()
return fig, axes
next(iter()):把 DataLoader 可迭代对象 转换成迭代器,取出第一个 batch 的数据。subplots:plt.subplots(行数, 列数)创建一张画布,2 行 4 列一共 8 个子图,用来摆放 8 张图片。fig:整张画布对象。axes:二维数组[[ax0,ax1,ax2,ax3],[ax4,ax5,ax6,ax7]],存放每一个小绘图区域。figsize=(10,6):画布宽 10 英寸、高 6 英寸。
enumerate(axes.flat):axes.flat:把二维的 axes 数组展平成一维迭代器。原本 2 行 4 列:[[0,1,2,3],[4,5,6,7]],flat 之后顺序遍历:0,1,2,3,4,5,6,7。enumerate:同时拿到索引i和子图对象ax,i 取值:0 ~ 7,正好对应前 8 张图片。
tight_layout:自动调整各个子图间距,防止标题、图片互相重叠。show:弹出窗口渲染、显示整张画布。
神经网络模型
搭建一个非常简单的 CNN(卷积神经网络),比全连接网络效果好很多。
# models/cnn.py
import torch
import torch.nn as nn
class SimpleCNN(nn.Module):
def __init__(self):
super(SimpleCNN, self).__init__()
# 卷积层 1:输入 1 通道(灰度图),输出 32 通道
self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)
# 卷积层 2:输入 32 通道,输出 64 通道
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
# 池化层:压缩图片尺寸
self.pool = nn.MaxPool2d(2, 2)
# 全连接层:把 64 张 7×7 的特征图展平,映射到 10 个数字
self.fc1 = nn.Linear(64 * 7 * 7, 128)
self.fc2 = nn.Linear(128, 10)
# Dropout 防止过拟合
self.dropout = nn.Dropout(0.25)
def forward(self, x):
# 卷积 → 激活 → 池化
x = self.pool(torch.relu(self.conv1(x)))
x = self.pool(torch.relu(self.conv2(x)))
# 展平(把二维特征图拉成一维向量)
x = x.view(-1, 64 * 7 * 7)
# 全连接层
x = torch.relu(self.fc1(x))
x = self.dropout(x)
x = self.fc2(x)
return x
class SimpleCNN(nn.Module)::定义一个名为SimpleCNN的类,括号内是父类。nn.Module是 PyTorch 所有神经网络的基类,必须继承它,才能获得 PyTorch 的训练、推理能力。def __init__(self)::Python 类的构造函数,在创建对象时执行(model = SimpleCNN()时会执行)。super(SimpleCNN, self).__init__():调用父类nn.Module的构造函数,初始化 PyTorch 框架需要的基础属性。必须写这一行,否则模型无法正常工作。nn.Conv2d():2D 卷积层,专门处理图像数据的层,提取局部特征(边缘、纹理)。- 第 1 个参数:输入通道数。MNIST 是灰度图,只有 1 个通道,彩色图是 3(RGB)。
- 第 2 个参数:输出通道数(也叫卷积核数量)。用 32 个不同的过滤器去提取特征,每个过滤器专注提取不同模式。
- kernel_size=3:卷积核大小 3×3,每次只看图像的一小块(3×3 像素区域),提取局部特征。
- padding=1:填充 1 圈 0,保持图像尺寸不变(28×28 → 28×28),防止边缘信息丢失过快。
nn.MaxPool2d(2, 2):最大池化层,在 2×2 区域内取最大值,压缩图像尺寸。即,把 28×28 变成 14×14,再变成 7×7。压缩尺寸的同时保留最显著的特征(取最大值相当于"保留每个区域最强的激活信号"),同时大幅减少参数量。- 第 1 个参数:池化窗口大小,比如: 2×2。
- 第 2 个参数:步长 2,每次滑动 2 步。
nn.Linear(64 * 7 * 7, 128):全连接层,每个输入神经元连接到每个输出神经元。- 输入维度:64 7 7,上一层输出的特征图是 64 通道 × 7×7 = 3136 个像素。
- 输出维度:压缩到 128 个神经元,作为"高层特征"。
nn.Linear(128, 10):输出 10 个数,对应 0~9 的logits(未归一化的分对数)。还不是概率;交给CrossEntropyLoss时不要先做 Softmax,损失内部会做。要概率时在推理里再softmax。Dropout(0.25):防止过拟合。每次训练时随机丢弃 25% 的神经元,迫使模型不依赖任何单个特征,让网络学习更鲁棒的模式。只在训练时生效,推理时自动关闭。过拟合是模型把训练集里的细节、噪声、无关特征死记硬背下来,而没有学到通用规律。def forward(self, x)::定义前向传播函数,数据从输入层流向输出层的路径。x 是输入张量,形状是[batch_size, 1, 28, 28]。注意:在 PyTorch 中,只需要定义forward,不需要定义backward(反向传播)。PyTorch 的自动求导机制(autograd)会根据前向计算图自动推导反向传播。pool(torch.relu(self.conv1(x))):- conv1(x):卷积,提取局部特征,输出 32 通道特征图。
- torch.relu():激活函数,把负数变成 0(增加非线性)。
- pool():池化,压缩尺寸 28×28 → 14×14。
view(-1, 64 * 7 * 7):改变张量形状。-1 表示自动推断该维度大小(保持 batch_size 不变),64 7 7 把 64 通道 × 7×7 的特征图展平成一维向量。SimpleCNN().to(device):创建一个模型实例(在 CPU 内存里),把模型的所有参数(权重、偏置)搬到 device 指定的地方。如果device = "cuda",模型参数会被搬到显卡的显存里;如果device = "cpu",则留在电脑的内存里。
损失函数和优化器
损失函数(Loss Function、代价函数):用来衡量模型预测结果和真实标签之间的差距,输出数值越大预测越不准,数值越小预测越贴近真实答案。
两种最常用损失:
- 交叉熵损失
CrossEntropyLoss:多分类标配(不限 MNIST)。输入是 logits 和整数类别标签,内部做 LogSoftmax + NLLLoss。 - 均方误差 MSELoss(回归任务):用于预测连续数字(房价、温度),计算预测值与真值差值的平方平均。
优化器(Optimizer):拿到损失算出的梯度后,负责更新网络权重参数的工具。损失只算出 “差多少”,优化器负责 “怎么改参数减小误差”。
常见优化器:
- SGD:随机梯度下降,基础款。
- Adam、AdamW:工业界最常用,自适应学习率,收敛更快。
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
def train_one_epoch(model, train_loader, criterion, optimizer, device):
model.train()
running_loss = 0.0
correct = 0
total = 0
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
_, predicts = torch.max(outputs, 1)
total += labels.size(0)
correct += (predicts == labels).sum().item()
epoch_loss = running_loss / len(train_loader)
epoch_acc = 100 * correct / total
return epoch_loss, epoch_acc
nn.CrossEntropyLoss():损失函数。optim.Adam(model.parameters(), lr=0.001):优化器。model.parameters()是一个生成器(Generator),会遍历模型中所有的nn.Parameter对象;lr是学习率,控制每次参数更新的步伐大小。model.train():将模型切换到 训练模式,模型默认是训练模式,但显式声明是良好习惯。训练模式启用Dropout层(训练时随机丢弃神经元),启用BatchNorm层的统计更新。optimizer.zero_grad():将模型中所有参数的梯度清零。model(images):前向传播,把图片数据传入模型,得到预测结果。criterion(outputs, labels):计算损失值,衡量模型预测和真实标签的差距。loss.backward():反向传播,计算损失对每个参数的梯度。optimizer.step():根据当前梯度,更新模型的所有参数。loss.item():当前 batch 的损失值(Python 浮点数)。torch.max(outputs, 1):在类别维上取最大值,得到预测类别。logits 上 argmax 与先 Softmax 再 argmax 结果相同。labels.size(0):当前 batch 的样本数。predicts == labels:得到布尔张量,sum()统计预测正确的个数。
评估模型
训练完后,用从未见过的 10,000 张测试图来验证模型真正的泛化能力。
def evaluate(model, test_loader, criterion, device):
model.eval()
running_loss = 0.0
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)
loss = criterion(outputs, labels)
running_loss += loss.item()
_, predicts = torch.max(outputs, 1)
total += labels.size(0)
correct += (predicts == labels).sum().item()
avg_loss = running_loss / len(test_loader)
acc = 100 * correct / total
return avg_loss, acc
model.eval():关闭 Dropout、BatchNorm 的行为变化。torch.no_grad():关闭梯度计算,节省显存和加速推理。
保存模型
.pt、.pth 都是常见后缀,保存的是 state_dict()(各层参数张量),不包含网络结构。两种后缀格式没有区别,都是 Python 序列化;给 C++、其他语言用,应导出 TorchScript 或 ONNX,而不是指望换个后缀就能跨语言。
torch.save(model.state_dict(), 'mnist_model.pth')
torch.save(model.state_dict(), 'mnist_model.pt')
加载时建议写 map_location,并在较新的 PyTorch 里加 weights_only=True(只接受张量,避免任意反序列化):
model.load_state_dict(torch.load('mnist_model.pt', map_location=device, weights_only=True))
使用模型
加载模型并对单张图片做预测。
# deploy/apply.py
import sys
import os
import torch
import torchvision.transforms as transforms
from PIL import Image
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from models.cnn import SimpleCNN
# Apple silicon 可再判断 torch.backends.mps.is_available(),用 "mps"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SimpleCNN().to(device)
model.load_state_dict(torch.load('mnist_model.pt', map_location=device, weights_only=True))
model.eval()
transform = transforms.Compose([
transforms.Grayscale(num_output_channels=1),
transforms.Resize((28, 28)),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
def predict_image(image):
if isinstance(image, str):
image = Image.open(image)
elif hasattr(image, 'read'):
image = Image.open(image)
img_tensor = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
output = model(img_tensor)
probabilities = torch.softmax(output, dim=1)
confidence, predicted = torch.max(probabilities, 1)
return predicted.item(), confidence.item()
让预测脚本接受命令行参数,方便和其他工具集成:
import sys
from apply import predict_image
if len(sys.argv) > 1:
image_path = sys.argv[1]
pred, conf = predict_image(image_path)
print(f"{pred} ({conf:.2%})")
else:
print("用法:python predict.py <图片路径>")
python -B deploy/execute.py assets/images/img_002.jpeg
交付模型
打包成可执行文件
生成 dist/execute.exe(Windows)或可执行文件(Mac、Linux),使用者不需要装 Python 就能用。
pip install pyinstaller
pyinstaller --onefile deploy/execute.py
# 冲突,忽略 PyQt5
pyinstaller --onefile --exclude PyQt5 deploy/execute.py
# 生成 .app(Mac),更通用
pyinstaller --onefile --windowed --exclude PyQt5 deploy/execute.py
部署成 Web API
# deploy/server.py
from flask import Flask, request, jsonify
from flask_cors import CORS
from apply import predict_image
app = Flask(__name__)
CORS(app)
@app.route('/predict', methods=['POST'])
def predict():
if 'image' not in request.files:
return jsonify({'error': '请上传图片文件'}), 400
file = request.files['image']
if file.filename == '':
return jsonify({'error': '未选择文件'}), 400
try:
pred, conf = predict_image(file.stream)
return jsonify({
'success': True,
'prediction': pred,
'confidence': round(conf, 4),
'confidence_percent': f"{conf:.2%}"
})
except Exception as e:
import traceback
traceback.print_exc()
return jsonify({'error': str(e)}), 500
print("服务启动:http://localhost:5000")
print("curl -X POST -F 'image=@test.png' http://localhost:5000/predict")
app.run(host='0.0.0.0', port=5000, debug=True)
转成 ONNX 格式
ONNX 是通用的模型格式,可以在 C++、Java、C# 等其他语言中加载。
pip install onnx onnxruntime onnxscript
import torch
from models.cnn import SimpleCNN
model = SimpleCNN()
model.load_state_dict(torch.load('mnist_model.pt', map_location='cpu', weights_only=True))
model.eval()
dummy_input = torch.randn(1, 1, 28, 28)
torch.onnx.export(model, dummy_input, 'mnist_model.onnx')
# 较新的 PyTorch 也可用 dynamo 导出路径,见官网 ONNX 文档
自定义训练数据集
训练自己的图片数据集:
- 整理图片文件:把所有图片放在一个文件夹里,图片格式不限(JPG、PNG 都可以),文件名随意。
- 创建标签文件(CSV):创建一个 CSV 文件,人工标注每张图片对应的数字标签。
- 写一个自定义 Dataset 类。
- 用 DataLoader 加载,开始训练。
data/
├── images/
│ ├── img_001.png
│ ├── img_002.png
│ └── ...
└── labels.csv
filename,label
img_001.jpeg,0
img_002.jpeg,1
import torch
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import pandas as pd
import torchvision.transforms as transforms
from train import train_one_epoch
from utils.custom_curves import plot_curves
class MyDataset(Dataset):
def __init__(self, image_dir, label_file, transform=None):
self.image_dir = image_dir
self.transform = transform
self.df = pd.read_csv(label_file) # CSV 列:filename, label
def __len__(self):
return len(self.df)
def __getitem__(self, idx):
row = self.df.iloc[idx]
image_path = f"{self.image_dir}/{row['filename']}"
image = Image.open(image_path)
image_tensor = self.transform(image)
label = int(row['label'])
return image_tensor, label
def custom_train(model, ai_model_path, device, target_labels, epochs=20, lr=0.001):
model.load_state_dict(torch.load(ai_model_path, map_location=device, weights_only=True))
criterion = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
transform = transforms.Compose([
transforms.Grayscale(),
transforms.Resize((28, 28)),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
dataset = MyDataset(
image_dir='assets/images',
label_file='assets/labels.csv',
transform=transform
)
loader = DataLoader(dataset, batch_size=64, shuffle=True)
train_losses, train_accs = [], []
for epoch in range(epochs):
loss, acc = train_one_epoch(model, loader, criterion, optimizer, device)
train_losses.append(loss)
train_accs.append(acc)
print(f"第 {epoch+1} 轮:损失 {loss:.4f},准确率 {acc:.2f}%")
# plot_curves(train_losses, train_accs)
__init__:创建 MyDataset 对象时 Python 自动调用。__len__:PyTorch 内部机制自动调用,len(dataset)或DataLoader内部。__getitem__:PyTorch 内部机制自动调用,DataLoader内部(或手动dataset[idx])。target_labels:函数签名里留了这个参数,下面示例没有用到;真正按标签过滤时,要在读 CSV 之后自己筛。
注意:如果 labels.csv 里只有 0 和 1,模型在这些新数据上多轮训练后,会逐渐忘掉原来的 2~9,只剩 0 和 1。这叫灾难性遗忘。
简单解决方案:把 MNIST 旧数据和新数据混在一起训练,模型同时看到新旧数据,不会遗忘 2-9。更优解决方案,参考下方的『知识蒸馏』。
训练曲线可视化
训练曲线就是把每一轮的损失和准确率画成折线图,直观看到模型是否在进步、有没有过拟合。
损失和准确率曲线的趋势:一起下降、上升 = 正常;验证线反弹 = 过拟合;下降太慢 = 欠拟合;剧烈震荡 = 学习率太大。
# utils/custom_curves.py
import matplotlib
import matplotlib.pyplot as plt
# 中文字体:macOS 用下列字体;Windows 可改成 ['Microsoft YaHei', 'SimHei']
matplotlib.rcParams['font.sans-serif'] = ['Arial Unicode MS', 'Heiti TC', 'PingFang SC']
matplotlib.rcParams['axes.unicode_minus'] = False
def plot_curves(train_losses, train_accs):
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
epochs = range(1, len(train_losses) + 1)
# 左图:损失曲线
ax1.plot(epochs, train_losses, 'b-', label='训练损失')
ax1.set_xlabel('轮次')
ax1.set_ylabel('损失')
ax1.set_title('损失曲线')
ax1.legend()
ax1.grid(True)
# 右图:准确率曲线
ax2.plot(epochs, train_accs, 'b-', label='训练准确率')
ax2.set_xlabel('轮次')
ax2.set_ylabel('准确率 (%)')
ax2.set_title('准确率曲线')
ax2.legend()
ax2.grid(True)
plt.tight_layout()
plt.savefig('training_curves.png', dpi=150) # 保存图片
plt.show()
迁移学习
迁移学习(Transfer Learning) 是一种机器学习方法,利用在大规模数据上预训练好的模型作为起点,保留其已经学到的通用特征提取能力,只针对新任务进行微调, 从而用更少的数据和更短的时间达到更好的效果。
加载预训练权重:
# custom_train.py
model = SimpleCNN().to(device)
model.load_state_dict(torch.load('mnist_model.pt', map_location=device, weights_only=True))
迁移学习以加载 .pt 权重作为起点,在新数据上继续训练。如果只有少量数据,比如只有几张为 0 和 1 的图片,为避免遗忘旧知识,核心策略是冻结所有层,只训练最后一层(fc2)——这样 2-9 的能力完整保留,新 0、1 的图片能被"记住",但其他 0、1 的图片并不能正确识别了。
# custom_train.py
for param in model.parameters():
param.requires_grad = False
for param in model.fc2.parameters():
param.requires_grad = True
optimizer = torch.optim.Adam(model.fc2.parameters(), lr=lr)
数据量是迁移学习的天花板,任何参数调整都无法替代足够多的样本。
部分源码
.
└── assets
│ └── html
│ │ └── ai-digital-recognize.html
│ └── images
│ │ └── img_001.jpeg
│ │ └── img_002.jpeg
│ └── labels.csv
└── deploy
│ └── apply.py
│ └── execute.py
│ └── server.py
└── models
│ └── cnn.py
└── utils
│ └── custom_curves.py
│ └── imageToData.py
│ └── plot.py
└── main.py
└── train.py
└── custom_train.py
主文件
# main.py
import torch
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader
from utils.plot import plot_images
from models.cnn import SimpleCNN
from train import train
from custom_train import custom_train
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"正在使用:{device}")
to_tensor = transforms.ToTensor()
normalize = transforms.Normalize((0.5,), (0.5,))
transform = transforms.Compose([
to_tensor,
normalize
])
train_dataset = torchvision.datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
test_dataset = torchvision.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=64, shuffle=False)
# plot_images(train_loader)
model = SimpleCNN().to(device)
train(model, train_loader, test_loader, device, epochs=5, lr=0.001)
torch.save(model.state_dict(), 'mnist_model.pt')
# custom_train(model, 'mnist_model.pt', device, target_labels=[0,1], epochs=20, lr=0.001)
# torch.save(model.state_dict(), 'mnist_model.pt')
模型训练
# train.py
import torch
import torch.nn as nn
import torch.optim as optim
def train_one_epoch(model, train_loader, criterion, optimizer, device):
model.train()
running_loss = 0.0
correct = 0
total = 0
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
_, predicts = torch.max(outputs, 1)
total += labels.size(0)
correct += (predicts == labels).sum().item()
epoch_loss = running_loss / len(train_loader)
epoch_acc = 100 * correct / total
return epoch_loss, epoch_acc
def evaluate(model, test_loader, criterion, device):
model.eval()
running_loss = 0.0
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)
loss = criterion(outputs, labels)
running_loss += loss.item()
_, predicts = torch.max(outputs, 1)
total += labels.size(0)
correct += (predicts == labels).sum().item()
avg_loss = running_loss / len(test_loader)
acc = 100 * correct / total
return avg_loss, acc
def train(model, train_loader, test_loader, device, epochs=5, lr=0.001):
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=lr)
for epoch in range(epochs):
train_loss, train_acc = train_one_epoch(
model, train_loader, criterion, optimizer, device
)
test_loss, test_acc = evaluate(
model, test_loader, criterion, device
)
print(f"第 {epoch+1}/{epochs} 次训练:训练损失-{train_loss:.4f} 训练正确率-{train_acc:.2f}% 评测损失-{test_loss:.4f} 评测正确率-{test_acc:.2f}%")
自定义训练
# custom_train.py
import torch
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import pandas as pd
import torchvision.transforms as transforms
from train import train_one_epoch
from utils.custom_curves import plot_curves
class MyDataset(Dataset):
def __init__(self, image_dir, label_file, transform=None):
self.image_dir = image_dir
self.transform = transform
self.df = pd.read_csv(label_file) # CSV 列:filename, label
def __len__(self):
return len(self.df)
def __getitem__(self, idx):
row = self.df.iloc[idx]
image_path = f"{self.image_dir}/{row['filename']}"
image = Image.open(image_path)
image_tensor = self.transform(image)
label = int(row['label'])
return image_tensor, label
def custom_train(model, ai_model_path, device, target_labels, epochs=20, lr=0.001):
model.load_state_dict(torch.load(ai_model_path, map_location=device, weights_only=True))
criterion = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
transform = transforms.Compose([
transforms.Grayscale(),
transforms.Resize((28, 28)),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
dataset = MyDataset(
image_dir='assets/images',
label_file='assets/labels.csv',
transform=transform
)
loader = DataLoader(dataset, batch_size=64, shuffle=True)
train_losses, train_accs = [], []
for epoch in range(epochs):
loss, acc = train_one_epoch(model, loader, criterion, optimizer, device)
train_losses.append(loss)
train_accs.append(acc)
print(f"第 {epoch+1} 轮:损失 {loss:.4f},准确率 {acc:.2f}%")
# plot_curves(train_losses, train_accs)
同一套训练循环(前向、损失、反向、optimizer.step)换数据与网络就能做更复杂的分类;大模型也是这条回路,只是数据换成文本、结构换成 Transformer。
| 数字识别 | 复杂图像识别(动物、花草) | 大语言模型(GPT、Llama) | |
|---|---|---|---|
| 数据 | 6 万张 28×28 灰度图 | 几万+张彩色大图(224×224+) | TB 级互联网文本 |
| 模型结构 | 自己设计的 2 层小 CNN | 复用预训练模型(ResNet、EfficientNet) | Transformer(多头注意力机制) |
| 训练方式 | 从零开始训练 | 预训练模型 + 微调 | 两阶段:预训练 + 指令微调 |
| 代码差异 | 标准训练循环 | 几乎一样 | 训练循环代码相同,DataLoader 换成文本数据即可 |
训练人员的核心工作是根据任务选对损失、优化器、网络结构,准备高质量数据。这些决定效果上限;权重由优化器自动改,不必手调每个系数。卷积层、损失、优化器 PyTorch 都已实现。改算法、发明新结构是研究人员的事,例如用 Transformer 替代 RNN、设计新损失或更省算力的优化器。
相关问题
每次重新训练出来的权重完全一样吗?
不一样。 即使训练数据和代码完全相同,每次重新训练得到的权重都不相同。
原因有三点:
- 参数随机初始化:创建模型时,PyTorch 会对所有权重进行随机初始化(默认使用 Kaiming Uniform 初始化)。每次运行程序,随机种子不同,初始值就不同。
- 数据打乱:每次 epoch 遍历数据时,样本的顺序被随机打乱。模型看到的数据顺序不同,参数更新的路径就不同。
- 优化器的随机性:比如 dropout 随机丢弃神经元。所以即使初始值相同,更新路径也可能产生微小差异。
知识蒸馏
知识蒸馏的核心思想是让学生模型(Student Model)在训练时,不仅学习有限的真实标签(硬标签),还要模仿教师模型(Teacher Model)对数据的输出分布(软标签),从而保留教师模型对特征空间的原有认知结构。
训练过程中,学生模型同时优化两个目标:
- 硬标签损失:确保模型拟合当前任务的正确答案,维持对训练数据的拟合能力。
- 软标签损失:让学生模型的输出分布逼近教师模型的输出分布,保留教师模型对数据相似性的先验知识,起到正则化作用。
最终损失函数定义为:
Loss = α · Hard_Loss + (1−α) · Soft_Loss
通过软标签的引导,学生模型在适应新任务的同时,受到教师模型原有知识结构的约束,从而缓解因少量数据导致的过拟合和灾难性遗忘。
要复现某一次训练,需同时固定 torch.manual_seed、numpy、Python 的随机种子,并把 DataLoader 的 generator、cudnn 确定性等一并设好;只设一个种子往往不够。
总结
本篇用 MNIST 走完 PyTorch 闭环:数据经 ToTensor + Normalize 进 DataLoader,小 CNN 输出 logits,CrossEntropyLoss + Adam 做「前向 → 反向 → 更新」,再用测试集评估,把 state_dict 存成 .pt。之后可以单张推理、打成可执行文件、挂 Flask 或导出 ONNX。换自己的图时要防灾难性遗忘;迁移学习靠冻层或混旧数据,数据量仍是上限。