import torch
import torch.nn as nn
from torchvision import datasets, transforms
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

# 取一张真实MNIST图
transform = transforms.Compose([transforms.ToTensor()])
test_set = datasets.MNIST('./data', train=False, download=False, transform=transform)
img, label = test_set[0]      # [1,28,28]
print(f"原图 label = {label}")

conv = nn.Conv2d(1, 32, kernel_size=3, padding=1)
pool = nn.MaxPool2d(2)

x0 = img.unsqueeze(0)         # [1,1,28,28] 加batch维
x1 = torch.relu(conv(x0))     # 卷积+relu
x2 = pool(x1)                 # 池化

# ===== 打印尺寸变化 =====
print("\n=== 尺寸变化（眼见为实）===")
print(f"输入图片      : {tuple(x0.shape)}  (batch=1, 通道=1, 28×28)")
print(f"卷积+ReLU 后  : {tuple(x1.shape)}  (32个核 -> 32通道, padding=1所以还是28×28)")
print(f"MaxPool(2) 后 : {tuple(x2.shape)}  (尺寸减半 28->14, 通道不变=32)")

# 验证卷积尺寸公式
H_in=28; pad=1; k=3; s=1
H_out = (H_in + 2*pad - k)//s + 1
print(f"\n卷积尺寸公式: (28 + 2×1 - 3)/1 + 1 = {H_out}  ✓ 对上了")
print(f"池化尺寸公式: 28 / 2 = 14  ✓")

# ===== 可视化：原图 -> 某个卷积通道 -> 池化后 =====
ch = 5  # 看第5个通道
fig, axes = plt.subplots(1, 3, figsize=(13, 4.5))

axes[0].imshow(x0[0,0].detach(), cmap='gray')
axes[0].set_title(f'1) INPUT\nshape (1,28,28)\ndigit={label}', fontsize=11)
axes[0].axis('off')

axes[1].imshow(x1[0,ch].detach(), cmap='viridis')
axes[1].set_title(f'2) after Conv+ReLU (channel #{ch})\nshape (32,28,28)\nsame size, 32 feature maps', fontsize=11)
axes[1].axis('off')

axes[2].imshow(x2[0,ch].detach(), cmap='viridis')
axes[2].set_title(f'3) after MaxPool(2) (channel #{ch})\nshape (32,14,14)\nhalf size, sharper', fontsize=11)
axes[2].axis('off')

fig.suptitle('Size flow:  (1,28,28) --Conv+pad1--> (32,28,28) --MaxPool2--> (32,14,14)', fontsize=13, y=1.02)
plt.tight_layout()
plt.savefig('cnn_size_flow.png', dpi=120, bbox_inches='tight')
print("\nsaved cnn_size_flow.png")
