沐神-动手学习深度学习-3.5图像分类数据集 (Fashion-MNIST)
·
针对《动手学深度学习》第 3.5 节 “图像分类数据集 (Fashion-MNIST)” 末尾的练习题,我为你整理了详细的逻辑分析与解答代码。
- 减少 batch_size(如设置为 1)对读取性能的影响
解析:
- 性能下降:当 batch_size=1 时,Python 必须为每一个单独的图像启动一次数据加载循环并进行预处理。
- I/O 瓶颈:由于无法利用现代 CPU 的并行处理能力(SIMD)和磁盘的批量读取特性,系统开销(Overhead)将远大于实际处理数据的时间。
- 实验结果:你会发现读取完整个数据集的时间会从几秒钟激增到几分钟。
- 数据迭代器的性能:数据读取速度 vs 模型训练速度
解析:
- 关键原则:数据读取速度必须快于模型训练速度(梯度更新速度),否则 GPU 就会处于空闲等待状态(Starvation),造成算力浪费。
- 解决方案:
- 增加 num_workers 以利用多进程并行读取。
- 使用高速 SSD 存储数据集。
- 在内存充足的情况下,将数据集预先加载到内存中。
- 查阅框架的在线 API 文档,还有哪些其他数据集可用?
解析:
以 PyTorch 为例,torchvision.datasets 提供了非常丰富的内置数据集:
- 分类:CIFAR-10, CIFAR-100, ImageNet, CelebA, STL-10。
- 检测/分割:COCO, VOC, Cityscapes。
- 手写体:MNIST, EMNIST, KMNIST。
- 动手代码实验:调整 batch_size 和 num_workers
你可以运行以下测试脚本来直接观察性能变化:
import time
import torch
import torchvision
from torch.utils import data
from torchvision import transforms
def get_dataloader_workers():
"""使用4个进程来读取数据"""
return 4
def load_data_fashion_mnist(batch_size, num_workers):
trans = transforms.ToTensor()
mnist_train = torchvision.datasets.FashionMNIST(
root="../data", train=True, transform=trans, download=True)
return data.DataLoader(mnist_train, batch_size, shuffle=True,
num_workers=num_workers)
测试不同配置下的速度
for b_size in [1, 64, 1024]:
for n_work in [0, 4]:
train_iter = load_data_fashion_mnist(b_size, n_work)
start = time.time()
for X, y in train_iter:
continue
print(f'Batch size: {b_size}, Workers: {n_work} -> Time: {time.time() - start:.2f} sec')
- 图像变换 (Transforms) 的应用
题目建议:针对图像分类,还有哪些常用的图像变换?
除了本节提到的 ToTensor(将图像转为浮点张量并除以 255),常见的还有:
- 水平翻转 (RandomHorizontalFlip):增加模型对方向的鲁棒性。
- 随机裁剪 (RandomResizedCrop):让模型学会识别物体的局部特征。
- 标准化 (Normalize):减去均值并除以方差,有助于模型更快收敛。
- 颜色抖动 (ColorJitter):调整亮度、对比度和饱和度。
下一步建议:
数据准备工作完成后,最核心的环节就是 3.6 节的 Softmax 回归从零实现。在那一节,你将亲手编写一个循环,把今天加载的数据喂给模型,观察损失函数是如何下降的。
你想让我先解释 交叉熵损失函数的代码实现,还是直接展示 完整的训练循环逻辑?
更多推荐

所有评论(0)