import torch
import torch.nn as nn
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import time

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))
])
train_set = datasets.MNIST('./data', train=True,  download=False, transform=transform)
test_set  = datasets.MNIST('./data', train=False, download=False, transform=transform)
train_loader = DataLoader(train_set, batch_size=128, shuffle=True)
test_loader  = DataLoader(test_set,  batch_size=256, shuffle=False)

device = 'cuda' if torch.cuda.is_available() else 'cpu'
print("device:", device)

# ===== CNN 模型 =====
class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.features = nn.Sequential(
            # 输入 [B,1,28,28]
            nn.Conv2d(1, 32, kernel_size=3, padding=1),   # ->[B,32,28,28] 32个卷积核提特征
            nn.ReLU(),
            nn.MaxPool2d(2),                              # ->[B,32,14,14] 降采样
            nn.Conv2d(32, 64, kernel_size=3, padding=1),  # ->[B,64,14,14]
            nn.ReLU(),
            nn.MaxPool2d(2),                              # ->[B,64,7,7]
        )
        self.classifier = nn.Sequential(
            nn.Flatten(),                # ->[B, 64*7*7=3136]
            nn.Linear(64*7*7, 128),
            nn.ReLU(),
            nn.Dropout(0.25),            # 防过拟合
            nn.Linear(128, 10)
        )
    def forward(self, x):
        x = self.features(x)
        return self.classifier(x)

model = CNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

# 参数量对比
n_params = sum(p.numel() for p in model.parameters())
print(f"CNN 参数量: {n_params:,}")

def evaluate():
    model.eval()
    correct = total = 0
    with torch.no_grad():
        for x, y in test_loader:
            x, y = x.to(device), y.to(device)
            pred = model(x).argmax(1)
            correct += (pred == y).sum().item()
            total += y.size(0)
    return correct / total

for epoch in range(5):
    model.train()
    t0 = time.time()
    for x, y in train_loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()
        loss = criterion(model(x), y)
        loss.backward()
        optimizer.step()
    acc = evaluate()
    print(f"Epoch {epoch+1} | loss={loss.item():.4f} | test acc={acc*100:.2f}% | {time.time()-t0:.1f}s")

print(f"\nCNN 最终测试准确率: {evaluate()*100:.2f}%")
