import torch
import torch.nn as nn
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np

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)

device = 'cpu'

class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.pool = nn.MaxPool2d(2)
        self.fc1 = nn.Linear(64*7*7, 128)
        self.drop = nn.Dropout(0.25)
        self.fc2 = nn.Linear(128, 10)
    def forward(self, x):
        x = self.pool(torch.relu(self.conv1(x)))
        x = self.pool(torch.relu(self.conv2(x)))
        x = x.flatten(1)
        x = self.drop(torch.relu(self.fc1(x)))
        return self.fc2(x)

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

# 训练2轮够看效果了
print("训练中...")
for epoch in range(2):
    model.train()
    for x, y in train_loader:
        optimizer.zero_grad()
        loss = criterion(model(x), y)
        loss.backward()
        optimizer.step()
    print(f"  epoch {epoch+1} done, loss={loss.item():.4f}")

# ===== 图1: 第一层32个卷积核 =====
kernels = model.conv1.weight.data.clone()   # [32,1,3,3]
fig, axes = plt.subplots(4, 8, figsize=(10, 5.5))
for i, ax in enumerate(axes.flat):
    k = kernels[i, 0].numpy()
    ax.imshow(k, cmap='RdBu_r')
    ax.set_title(f'#{i}', fontsize=7)
    ax.axis('off')
fig.suptitle('Layer-1: 32 learned 3x3 conv kernels (each detects an edge/pattern)', fontsize=12)
plt.tight_layout()
plt.savefig('cnn_kernels.png', dpi=120, bbox_inches='tight')
print("saved cnn_kernels.png")

# ===== 图2: 一张数字图 经过第一层卷积后的特征图 =====
model.eval()
img, label = test_set[0]   # 取一张测试图
with torch.no_grad():
    feat = torch.relu(model.conv1(img.unsqueeze(0)))   # [1,32,28,28]
feat = feat[0].numpy()

fig, axes = plt.subplots(4, 9, figsize=(13, 6))
# 第一格放原图
axes[0,0].imshow(img[0].numpy(), cmap='gray')
axes[0,0].set_title(f'INPUT (digit={label})', fontsize=9, color='red')
axes[0,0].axis('off')
# 剩下放32张特征图(放前35格)
idx = 0
for i, ax in enumerate(axes.flat):
    if i == 0:
        continue
    if idx < 32:
        ax.imshow(feat[idx], cmap='viridis')
        ax.set_title(f'feat#{idx}', fontsize=6)
        idx += 1
    ax.axis('off')
fig.suptitle('Same digit through 32 conv kernels -> 32 feature maps (different parts light up)', fontsize=12)
plt.tight_layout()
plt.savefig('cnn_featuremaps.png', dpi=120, bbox_inches='tight')
print("saved cnn_featuremaps.png")
