针对《动手学深度学习》第 3.5 节 “图像分类数据集 (Fashion-MNIST)” 末尾的练习题,我为你整理了详细的逻辑分析与解答代码。

  1. 减少 batch_size(如设置为 1)对读取性能的影响
    解析:
  • 性能下降:当 batch_size=1 时,Python 必须为每一个单独的图像启动一次数据加载循环并进行预处理。
  • I/O 瓶颈:由于无法利用现代 CPU 的并行处理能力(SIMD)和磁盘的批量读取特性,系统开销(Overhead)将远大于实际处理数据的时间。
  • 实验结果:你会发现读取完整个数据集的时间会从几秒钟激增到几分钟。
  1. 数据迭代器的性能:数据读取速度 vs 模型训练速度
    解析:
  • 关键原则:数据读取速度必须快于模型训练速度(梯度更新速度),否则 GPU 就会处于空闲等待状态(Starvation),造成算力浪费。
  • 解决方案:
    • 增加 num_workers 以利用多进程并行读取。
    • 使用高速 SSD 存储数据集。
    • 在内存充足的情况下,将数据集预先加载到内存中。
  1. 查阅框架的在线 API 文档,还有哪些其他数据集可用?
    解析:
    以 PyTorch 为例,torchvision.datasets 提供了非常丰富的内置数据集:
  • 分类:CIFAR-10, CIFAR-100, ImageNet, CelebA, STL-10。
  • 检测/分割:COCO, VOC, Cityscapes。
  • 手写体:MNIST, EMNIST, KMNIST。
  1. 动手代码实验:调整 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')
  1. 图像变换 (Transforms) 的应用
    题目建议:针对图像分类,还有哪些常用的图像变换?
    除了本节提到的 ToTensor(将图像转为浮点张量并除以 255),常见的还有:
  • 水平翻转 (RandomHorizontalFlip):增加模型对方向的鲁棒性。
  • 随机裁剪 (RandomResizedCrop):让模型学会识别物体的局部特征。
  • 标准化 (Normalize):减去均值并除以方差,有助于模型更快收敛。
  • 颜色抖动 (ColorJitter):调整亮度、对比度和饱和度。
    下一步建议:
    数据准备工作完成后,最核心的环节就是 3.6 节的 Softmax 回归从零实现。在那一节,你将亲手编写一个循环,把今天加载的数据喂给模型,观察损失函数是如何下降的。
    你想让我先解释 交叉熵损失函数的代码实现,还是直接展示 完整的训练循环逻辑?
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐