保存、加载、推理、断点续训

🎉摘要:基于PyTorch实现MNIST手写数字识别CNN,完整演示训练、保存checkpoint、加载断点续训、推理演示及导出样本。运行5个epoch后实时显示训练曲线,并绘制混淆矩阵和错误样本,适合深度学习初学者学习模型保存与推理。

将下面代码保存到 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()

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号