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

# ===== 0. 数据：MNIST 手写数字 (28x28灰度图, 0~9共10类) =====
transform = transforms.Compose([
    transforms.ToTensor(),                       # [0,255] -> [0,1], shape [1,28,28]
    transforms.Normalize((0.1307,), (0.3081,))   # 标准化(MNIST经验均值方差)
])
train_set = datasets.MNIST('./data', train=True,  download=True, transform=transform)
test_set  = datasets.MNIST('./data', train=False, download=True, 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)

# ===== 1. 模型：纯全连接 MLP (把28x28拉平成784维) =====
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Flatten(),            # [B,1,28,28] -> [B,784]
            nn.Linear(784, 256),
            nn.ReLU(),
            nn.Linear(256, 128),
            nn.ReLU(),
            nn.Linear(128, 10)       # 输出10维(对应0~9), 不加softmax(交给loss)
        )
    def forward(self, x):
        return self.net(x)

model = MLP().to(device)

# ===== 2. 损失 + 优化器 =====
criterion = nn.CrossEntropyLoss()   # 内部含softmax+log+NLL, 多分类标配
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

# ===== 3. 训练 (五步骨架, 多了batch循环) =====
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()
    for x, y in train_loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()        # 1
        out = model(x)               # 2 前向
        loss = criterion(out, y)     # 3 算loss
        loss.backward()              # 4 反向
        optimizer.step()             # 5 更新
    acc = evaluate()
    print(f"Epoch {epoch+1} | loss={loss.item():.4f} | test acc={acc*100:.2f}%")

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