完整训练 MNIST CNN

🎉摘要:本教程提供PyTorch实现MNIST CNN训练循环的完整代码,涵盖网络定义、数据加载、训练函数及3个epoch的训练结果(最终准确率99.06%),适合深度学习初学者实战练习。

把下面代码保存为 chapter07_training_loop.py,运行后会真正开始训练。预计 3 个 epoch 约 2~3 分钟(CPU),具体时间要看你机器的性能。

代码如下:

"""
第 7 章:训练循环实战 —— 完整训练 MNIST CNN

运行方式:python chapter07_training_loop.py
"""
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader


# 网络定义(同第5章)
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):
    """训练一个 epoch"""
    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)

    avg_loss = total_loss / len(loader)
    accuracy = 100.0 * correct / total
    return avg_loss, accuracy


# 主程序
def main():
    print("=" * 50)
    print("  MNIST CNN 训练(第7章)")
    print("=" * 50)

    # 1. 数据准备
    print("\n[1/4] 加载数据...")
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    train_dataset = datasets.MNIST("./data", train=True, download=True, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
    print(f"  训练集: {len(train_dataset)} 张图, {len(train_loader)} 个 batch")

    # 2. 创建模型、损失函数、优化器
    print("\n[2/4] 创建模型...")
    model = MNIST_CNN()
    loss_fn = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    total_params = sum(p.numel() for p in model.parameters())
    print(f"  模型参数: {total_params:,}")

    # 3. 训练循环
    print("\n[3/4] 开始训练...")
    num_epochs = 3
    history = {"loss": [], "accuracy": []}

    for epoch in range(num_epochs):
        avg_loss, accuracy = train_one_epoch(model, train_loader, loss_fn, optimizer)
        history["loss"].append(avg_loss)
        history["accuracy"].append(accuracy)
        print(f"  Epoch {epoch+1}/{num_epochs} | Loss: {avg_loss:.4f} | Acc: {accuracy:.2f}%")

    # 4. 训练总结
    print("\n[4/4] 训练完成!")
    print("─" * 50)
    print(f"  初始 Loss: {history['loss'][0]:.4f}")
    print(f"  最终 Loss: {history['loss'][-1]:.4f}")
    print(f"  最终准确率: {history['accuracy'][-1]:.2f}%")
    print(f"  Loss 下降: {history['loss'][0]:.4f} → {history['loss'][-1]:.4f}")
    print(f"  准确率提升: {history['accuracy'][0]:.2f}% → {history['accuracy'][-1]:.2f}%")

    # 简单文本绘图
    print("\n  Loss 下降趋势:")
    for i, loss in enumerate(history["loss"]):
        bar = "█" * int(loss * 20)
        print(f"  Epoch {i+1}: {bar} {loss:.4f}")


if __name__ == "__main__":
    main()

运行示例,输出如下:

==================================================
  MNIST CNN 训练(第7章)
==================================================

[1/4] 加载数据...
  训练集: 60000 张图, 938 个 batch

[2/4] 创建模型...
  模型参数: 421,642

[3/4] 开始训练...
  Epoch 1/3 | Loss: 0.1367 | Acc: 95.80%
  Epoch 2/3 | Loss: 0.0429 | Acc: 98.67%
  Epoch 3/3 | Loss: 0.0292 | Acc: 99.06%

[4/4] 训练完成!
──────────────────────────────────────────────────
  初始 Loss: 0.1367
  最终 Loss: 0.0292
  最终准确率: 99.06%
  Loss 下降: 0.1367 → 0.0292
  准确率提升: 95.80% → 99.06%

  Loss 下降趋势:
  Epoch 1: ██ 0.1367
  Epoch 2:  0.0429
  Epoch 3:  0.0292

  

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