将下面代码保存到 chapter10_inference.py 文件,运行后会演示保存、加载、推理、断点续训等。
代码如下:
"""
第 10 章:保存、加载、推理、断点续训
运行方式:python chapter10_inference.py
"""
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import os
# ── 网络定义 ────────────────────────────────────────────────
class MNIST_CNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
self.pool1 = nn.MaxPool2d(2)
self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.pool2 = nn.MaxPool2d(2)
self.fc1 = nn.Linear(64 * 7 * 7, 128)
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
x = torch.relu(self.conv1(x))
x = self.pool1(x)
x = torch.relu(self.conv2(x))
x = self.pool2(x)
x = x.view(x.size(0), -1)
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
# ── 训练/评估 ────────────────────────────────────────────────
def train_one_epoch(model, loader, loss_fn, optimizer):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in loader:
optimizer.zero_grad()
outputs = model(images)
loss = loss_fn(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item()
correct += outputs.argmax(dim=1).eq(labels).sum().item()
total += labels.size(0)
return total_loss / len(loader), 100.0 * correct / total
@torch.no_grad()
def evaluate(model, loader, loss_fn):
model.eval()
total_loss, correct, total = 0.0, 0, 0
for images, labels in loader:
outputs = model(images)
loss = loss_fn(outputs, labels)
total_loss += loss.item()
correct += outputs.argmax(dim=1).eq(labels).sum().item()
total += labels.size(0)
return total_loss / len(loader), 100.0 * correct / total
def save_checkpoint(epoch, model, optimizer, loss, best_acc, history, path):
torch.save({
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'loss': loss,
'best_acc': best_acc,
'history': history,
}, path)
def load_checkpoint(path, model, optimizer):
if os.path.exists(path):
cp = torch.load(path, map_location='cpu')
model.load_state_dict(cp['model_state_dict'])
optimizer.load_state_dict(cp['optimizer_state_dict'])
print(f"Checkpoint 已加载:Epoch {cp['epoch']}, Best Acc {cp['best_acc']:.2f}%")
return cp['epoch'], cp['best_acc'], cp.get('history', {})
return 0, 0.0, {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}
@torch.no_grad()
def predict_and_show(model, loader):
"""演示推理:取测试集前5张图预测并打印结果"""
model.eval()
print("\n" + "=" * 50)
print(" 推理演示(测试集前5张图)")
print("=" * 50)
images, labels = next(iter(loader))
for i in range(min(5, len(images))):
img = images[i:i+1] # 保持 batch 维度
true_label = labels[i].item()
outputs = model(img)
probs = torch.softmax(outputs, dim=1)[0]
pred = outputs.argmax(dim=1).item()
correct_mark = "✅" if pred == true_label else "❌"
print(f"\n 第{i+1}张图 {correct_mark}")
print(f" 真实数字: {true_label}")
print(f" 预测数字: {pred}(置信度 {probs[pred].item()*100:.1f}%)")
print(f" Top 3 概率: ", end="")
top3 = probs.topk(3)
for j in range(3):
print(f"{top3.indices[j].item()}({top3.values[j].item()*100:.0f}%)", end=" ")
print()
# ── 导出推理用样本 ──────────────────────────────────────────
def export_samples(num_per_class=1, output_dir='./examples'):
os.makedirs(output_dir, exist_ok=True)
raw_dataset = datasets.MNIST("./data", train=False, download=True,
transform=transforms.ToTensor())
collected = {i: [] for i in range(10)}
from torchvision.utils import save_image
for img, label in raw_dataset:
label_int = int(label)
if len(collected[label_int]) < num_per_class:
collected[label_int].append(img)
if all(len(v) >= num_per_class for v in collected.values()):
break
for digit, imgs in collected.items():
for j, img in enumerate(imgs):
save_image(img, f"{output_dir}/digit_{digit}_sample{j}.png")
print(f"样例图片已导出到 {output_dir}/")
# ── 主程序 ──────────────────────────────────────────────────
def main():
print("=" * 50)
print(" MNIST CNN 完整流程(第10章)")
print("=" * 50)
# 数据
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
train_dataset = datasets.MNIST("./data", train=True, download=True, transform=transform)
test_dataset = datasets.MNIST("./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)
# 模型
model = MNIST_CNN()
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 尝试加载 checkpoint
checkpoint_path = "checkpoint.pth"
start_epoch, best_acc, history = load_checkpoint(checkpoint_path, model, optimizer)
# 训练
num_epochs = 5
print(f"\n从 Epoch {start_epoch + 1} 开始,目标 {num_epochs} 轮")
print(f"{'Epoch':<8}{'Train Loss':<14}{'Train Acc':<12}{'Test Loss':<14}{'Test Acc':<12}")
print("-" * 60)
for epoch in range(start_epoch, num_epochs):
train_loss, train_acc = train_one_epoch(model, train_loader, loss_fn, optimizer)
test_loss, test_acc = evaluate(model, test_loader, loss_fn)
for key, val in zip(
["train_loss", "train_acc", "test_loss", "test_acc"],
[train_loss, train_acc, test_loss, test_acc]
):
if key not in history:
history[key] = []
history[key].append(val)
# 保存最佳模型
if test_acc > best_acc:
best_acc = test_acc
torch.save(model.state_dict(), "best_model.pth")
save_mark = " ★"
else:
save_mark = ""
print(f"{epoch+1:<8}{train_loss:<14.4f}{train_acc:<12.2f}{test_loss:<14.4f}"
f"{test_acc:<12.2f}{save_mark}")
# 每个 epoch 保存 checkpoint
save_checkpoint(epoch + 1, model, optimizer, train_loss, best_acc, history, checkpoint_path)
print(f"\n训练完成!最佳测试准确率: {best_acc:.2f}%")
print(f"最佳模型已保存: best_model.pth")
print(f"Checkpoint 已保存: {checkpoint_path}")
# 演示推理
predict_and_show(model, test_loader)
# 导出推理样本
export_samples(num_per_class=1, output_dir='./examples')
print(f"\n可以用以下命令对样本做推理测试:")
print(f" python infer.py --image examples/digit_0_sample0.png")
if __name__ == "__main__":
main()