code/08-neural-nets/digit_classifier.py
85 行 · 3.3 KB程式碼和執行結果保留原樣(簡體中文),與實際執行時完全一致。
"""训练一个多层感知机识别手写数字。数据是 scikit-learn 自带的 8×8 手写数字,不需要下载。
python digit_classifier.py
准备:uv add torch scikit-learn
全部在 CPU 上运行,几秒钟就能训练完。固定了随机种子。
"""
import numpy as np
import torch
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from torch import nn
torch.manual_seed(0)
digits = load_digits()
print(f"数据:{digits.images.shape[0]} 张图,每张 {digits.images.shape[1]}×{digits.images.shape[2]} 像素,"
f"像素值 {digits.images.min():.0f}~{digits.images.max():.0f}")
def show(image):
"""用字符把一张 8×8 的图画出来,深浅用不同的字符表示。"""
shades = " .:-=+*#"
return "\n".join(" " + "".join(shades[min(int(v) // 2, 7)] * 2 for v in row) for row in image)
print(f"第一张图,标签是 {digits.target[0]}:\n{show(digits.images[0])}")
# 64 个像素拉成一行作为输入;除以 16 把像素值缩放到 0~1
X = torch.tensor(digits.data / 16.0, dtype=torch.float32)
y = torch.tensor(digits.target, dtype=torch.long)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.25, random_state=0)
print(f"\n训练集 {len(X_train)} 张,测试集 {len(X_test)} 张")
model = nn.Sequential(
nn.Linear(64, 64), # 64 个像素 → 64 个隐藏神经元
nn.ReLU(),
nn.Linear(64, 10), # → 10 个输出,分别对应数字 0~9
)
print(f"模型共 {sum(p.numel() for p in model.parameters())} 个参数")
loss_fn = nn.CrossEntropyLoss() # 分类问题用交叉熵
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
def accuracy(X, y):
with torch.no_grad():
return (model(X).argmax(dim=1) == y).float().mean().item()
print("\n训练:")
for epoch in range(1, 31):
# 每一轮把训练集打乱,每次取 64 张图更新一次参数
order = torch.randperm(len(X_train))
for i in range(0, len(X_train), 64):
idx = order[i:i + 64]
loss = loss_fn(model(X_train[idx]), y_train[idx])
optimizer.zero_grad()
loss.backward()
optimizer.step()
if epoch in (1, 2, 5, 10, 20, 30):
print(f" 第 {epoch:2d} 轮 最后一批的损失 {loss.item():.4f} "
f"训练集准确率 {accuracy(X_train, y_train):.1%} 测试集准确率 {accuracy(X_test, y_test):.1%}")
# 模型输出的是 10 个"分数",softmax 把它们变成概率
with torch.no_grad():
probs = torch.softmax(model(X_test[:1]), dim=1)[0]
print(f"\n测试集第一张图(标签 {y_test[0].item()})的预测概率:")
print(" " + " ".join(f"{d}:{p:.2f}" for d, p in enumerate(probs.tolist())))
# 看看错在哪里
with torch.no_grad():
pred = model(X_test).argmax(dim=1)
wrong = (pred != y_test).nonzero().flatten()
print(f"\n测试集 {len(X_test)} 张里错了 {len(wrong)} 张。混淆矩阵(行是真实数字,列是预测数字):")
matrix = np.zeros((10, 10), dtype=int)
for t, p in zip(y_test.tolist(), pred.tolist()):
matrix[t, p] += 1
print(" " + " ".join(f"{d:3d}" for d in range(10)))
for d in range(10):
print(f" {d} " + " ".join(f"{v:3d}" if v else " ." for v in matrix[d]))
i = wrong[0].item()
print(f"\n一张判错的图:真实是 {y_test[i].item()},模型认为是 {pred[i].item()}")
print(show((X_test[i] * 16).reshape(8, 8).numpy()))