共计 2048 个字符,预计需要花费 6 分钟才能阅读完成。
问题背景
MNIST 数据集包含 60,000 张训练图像和 10,000 张测试图像,每张都是 28×28 像素的灰度手写数字(0-9)。在图像分类任务中,我们通常使用准确率作为主要评价指标。

全连接网络 (FNN) 在处理图像数据时会面临维度灾难问题。例如,对于 MNIST 的 28×28 图像,输入层就需要 784 个节点。如果第一隐藏层有 500 个节点,仅这一层就会产生 784*500=392,000 个权重参数。这种全连接方式不仅参数量爆炸,而且忽略了图像的局部空间结构信息。
技术方案对比
FNN 与 CNN 结构差异
- FNN:每个神经元与上一层的所有神经元全连接,适合处理向量化数据
- CNN:通过卷积核实现局部感受野,具有空间层次化特征提取能力
卷积操作的数学表达:
O[i,j] = ∑∑ (K[m,n] * I[i+m, j+n]) + b
其中 K 是卷积核,I 是输入图像,b 是偏置项。这种局部连接方式大幅减少了参数量。
PyTorch Lightning 实现框架
使用 PyTorch Lightning 可以规范训练流程,以下是一个模板类:
class LitModel(pl.LightningModule):
def __init__(self):
super().__init__()
self.example_input_array = torch.rand(1, 1, 28, 28)
def forward(self, x):
return self.net(x)
def training_step(self, batch, batch_idx):
x, y = batch
y_hat = self(x)
loss = F.cross_entropy(y_hat, y)
self.log('train_loss', loss)
return loss
核心代码实现
FNN 网络构建
class FNN(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Flatten(),
nn.Linear(784, 512),
nn.ReLU(),
nn.Dropout(0.2), # 防止过拟合
nn.Linear(512, 256),
nn.BatchNorm1d(256), # 批归一化
nn.ReLU(),
nn.Linear(256, 10)
)
CNN 网络构建
class CNN(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Conv2d(1, 32, 3, padding=1), # 保持空间维度
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1),
nn.BatchNorm2d(64), # 2D 批归一化
nn.ReLU(),
nn.MaxPool2d(2),
nn.Flatten(),
nn.Linear(64*7*7, 128),
nn.Dropout(0.5),
nn.Linear(128, 10)
)
训练优化技巧
- 梯度裁剪防止爆炸
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
- 早停机制实现
early_stop = EarlyStopping(
monitor='val_loss',
patience=5,
mode='min'
)
性能优化
训练速度对比
| 模型类型 | 每 epoch 时间(GPU) | 参数量 |
|---|---|---|
| FNN | 45s | 550K |
| CNN | 60s | 120K |
虽然 CNN 单次迭代稍慢,但达到相同准确率需要的 epoch 更少。
显存优化策略
- 使用混合精度训练
trainer = Trainer(precision=16) # 自动混合精度
- 梯度累积
trainer = Trainer(accumulate_grad_batches=4) # 相当于增大 batch size
避坑指南
- 张量格式问题
PyTorch 默认使用 NCHW 格式,但某些数据加载器可能输出 NHWC。转换方法:
x = x.permute(0, 3, 1, 2) # NHWC -> NCHW
- 学习率调整
批量大小增大 N 倍时,学习率也应增大√N 倍(线性缩放规则)。使用学习率预热:
lr_scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, start_factor=0.1, total_iters=5)
- 多 GPU 训练
注意 DataLoader 的 num_workers 设置:
dataloader = DataLoader(..., num_workers=min(4, os.cpu_count()//2))
完整代码与扩展
推荐进一步学习:
–《Deep Learning with PyTorch》书籍
– PyTorch 官方 Tutorials
– torch.profiler 性能分析工具
通过本次实践可以看到,CNN 在图像任务上具有明显优势。虽然实现稍复杂,但通过合理的网络设计和训练技巧,可以轻松达到 99%+ 的准确率。
正文完
发表至: 未分类
近一天内
