Skip to content

实战第 1 步:数据准备

工程习惯:先让数据"流起来",再搭模型。本篇建好 DataLoader,并检查每个批次的形状——这些形状就是稍后模型输入层的规格书。

创建项目目录

powershell
# Windows PowerShell(macOS/Linux 终端把 md 换成 mkdir -p)
md mnist-pytorch; cd mnist-pytorch; md data, weights

编写数据管道

新建 train.py,先写数据部分:

python
# train.py —— 第 1 步:数据准备
import torch
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

BATCH_SIZE = 64

transform = transforms.Compose([
    transforms.ToTensor(),                      # PIL图 → [1,28,28] float32,像素除255
    transforms.Normalize((0.1307,), (0.3081,)),  # MNIST 全局均值/标准差
])

train_set = datasets.MNIST(root="./data", train=True,  download=True, transform=transform)
test_set  = datasets.MNIST(root="./data", train=False, download=True, transform=transform)

train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True,  drop_last=False)
test_loader  = DataLoader(test_set,  batch_size=256,        shuffle=False, drop_last=False)

if __name__ == "__main__":   # Windows + num_workers>0 时必须有这个保护块
    images, labels = next(iter(train_loader))   # 取第一个批次看看
    print("批次图片形状:", tuple(images.shape))  # 期望 (64, 1, 28, 28)
    print("批次标签形状:", tuple(labels.shape))  # 期望 (64,)
    print("标签类型:", labels.dtype)              # 期望 torch.int64
    print("像素范围:", float(images.min()), float(images.max()))  # 均值0方差1,约 -3 ~ 3

在激活了 PyTorch 环境的终端运行(首次会下载约 11MB 数据到 ./data):

powershell
python train.py

预期输出:

text
批次图片形状: (64, 1, 28, 28)
批次标签形状: (64,)
标签类型: torch.int64
像素范围: -1.9519999 2.7150

逐行说"为什么"

  • ToTensor() 是必经一步:Dataset 吐出的原始样本是 0~255 的灰度图,ToTensor 同时完成三件事——转 float、除 255 归到 [0,1]、把 H×W×C 转成 PyTorch 要求的 CHW。跳过它,模型吃到的就是未归一化大数值,训练寸步难行。
  • shuffle=True 只给训练集:测试集要固定顺序,保证每次评估结果一致、可对比。
  • 标签是 int64 刚好合 CrossEntropyLoss 的意:如果哪天你用了 one-hot 标签(float32),损失函数会立刻用报错提醒你。
  • if __name__ == "__main__": 保护块:Windows 下 DataLoader 多进程用 spawn 启动子进程,不加保护块子进程会反复导入主脚本造成卡死/报错。养成习惯:主逻辑永远放进保护块。

记住这三个形状

后续模型设计完全围绕它们展开:

text
输入批次:   [64, 1, 28, 28]   (batch, channel, H, W)
标签批次:   [64]              (batch,) int64
模型应输出: [64, 10]          (batch, 类别数) float32 的 logits

数据检查通过,就可以动手搭模型了 → 第 2 步:模型设计