做深度学习训练的这些年,我越来越觉得一句话是对的:模型决定上限,数据管线决定你能不能摸到上限。 很多人把精力全砸在改网络结构、调超参上,结果训练一跑起来就卡在数据加载上,GPU 吃不满、Loss 乱跳、跑到一半 OOM,排查半天才发现问题出在 Dataset 和 DataLoader 这两个最不起眼的环节。DAY38 这篇,我把自己在实际项目里用 PyTorch Dataset 类和 DataLoader 类的完整思路、踩坑记录、调优经验一次性整理出来,新手可以照着抄,老手也能看看有没有自己忽略的细节。
这篇内容适合正处于 PyTorch 入门到进阶过渡期的朋友,也适合那些已经跑通 MNIST、但一接触真实业务数据就手足无措的人。我会先从这两个类为什么非要拆开讲起,再逐个拆解实现细节,最后给出一套可直接复用的数据管线代码,以及在多进程加载、随机采样、内存占用这些场景下的实战排查记录。
1. 数据管线的整体设计思路
1.1 为什么 PyTorch 要把数据加载拆成 Dataset 和 DataLoader 两个类
我最早接触 PyTorch 的时候也觉得奇怪,TensorFlow 那边一个 tf.data.Dataset 好像就把事情干完了,为什么 PyTorch 非要拆成两个类,徒增理解成本。后来在项目里被数据问题折磨过几轮,才明白这个拆分是刻意为之的。
核心原因可以用一句话概括:Dataset 负责"数据是什么",DataLoader 负责"数据怎么喂"。 这是一个非常经典的分层思想,类似把"业务逻辑"和"调度逻辑"分开。Dataset 类只需要回答三个问题:这个数据集有多大、怎么取第 i 个样本、这个样本长什么样。它完全不关心你要不要打乱顺序、一次取几个、用几个进程去取、取完要不要做数据增强。这些事全部交给 DataLoader 处理。
这样的好处在实际项目中非常明显。比如你换了数据集,从图片分类换成文本分类,只需要重写 Dataset;你的训练策略从全量梯度下降换成 mini-batch SGD,只需要改 DataLoader 的参数,不用动数据读取的代码。如果这两个职责揉在一个类里,每次改动都要小心翼翼,尤其在数据格式复杂、预处理链条长的时候,耦合带来的维护成本会被无限放大。
还有一个容易被忽略的点:Dataset 天然支持"随时索引任意样本"的语义。你写 dataset[3] 就能拿到第 4 个样本。这个特性在做数据集可视化、错误样本核查、分层采样时太重要了。DataLoader 则是把 dataset 包装成一个可迭代对象,你 for batch in loader 就能拿到一个 batch 的数据。一个解决"随机访问"问题,一个解决"顺序迭代"问题,两个类配合起来,数据管线的关注点被切得清清楚楚。
1.2 两个类各自主攻的痛点
Dataset 类主攻的是数据来源的多样性。真实项目里的数据从来不可能是干净整齐的,可能是几千张尺寸不一的图片散落在好几个文件夹里,可能是一个巨大的 CSV 文件按行分块读取,可能是数据库里查出来的记录,也可能是内存里已经加载好的 numpy 数组。Dataset 就是把这些五花八门的来源统一成"给定索引、返回样本"的接口,让你的训练代码永远不需要关心底层数据到底存在哪里。
DataLoader 类主攻的是训练效率与随机性的保障。深度学习训练有两个硬性需求:一是数据要分 batch 送入,二是每个 epoch 的样本顺序要随机打乱。如果自己在训练循环里写这些逻辑,代码会越写越乱,而且很容易踩多进程、内存复制的坑。DataLoader 把 batching、shuffling、并行加载、内存固定这些事全部封装好,还提供了 sampler、collate_fn 这样的扩展点,让你不需要重造轮子。
我自己的体会是,理解这两个类的分工之后,再去看 PyTorch 官方文档的 DataLoader 参数列表,就不会觉得头大了。参数虽多,但本质都是在问"数据怎么喂"这个问题——一次喂多少(batch_size)、按什么顺序喂(shuffle/sampler)、怎么并行喂(num_workers)、喂之前怎么把样本拼成一个 batch(collate_fn)。把问题归类之后,选参数就有了方向。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. Dataset 类核心细节与实现要点
2.1 必须实现的三个方法:len 与 getitem 的底层约定
PyTorch 的 torch.utils.data.Dataset 是一个抽象基类,官方要求子类必须重写 __len__ 和 __getitem__ 两个方法。有些教程还会提到 __init__,但严格来说 __init__ 不是必须的,只是通常情况下你总得在初始化时把数据路径、标签、预处理方式存下来。
__len__ 方法必须返回数据集的样本总数。这个值会被 DataLoader 用来计算每个 epoch 有多少个 batch,也会被 torch.utils.data.random_split 用来做数据集划分。如果返回的值和实际 __getitem__ 能取到的最大索引不一致,训练时大概率会报 IndexError,这类问题在自定义数据集里非常常见。
__getitem__ 方法是整个数据管线的核心,它接收一个整数索引 idx,返回一个样本。关于返回值,我想多说几句:训练时你返回的一般是 (data, label) 这样的元组,DataLoader 会自动把多个样本的 data 堆成 batch、label 堆成 batch。但如果你做的是推理任务、或者样本本身结构比较复杂(比如目标检测的图片加多个 bbox),返回值可以是任意结构,只要你自己能处理就行。此时就要靠 collate_fn 来告诉 DataLoader 怎么把这些复杂结构拼起来,这个我在后面的章节详细讲。
一个我在实际项目里常用的写法是,在 __getitem__ 里只做"取原始样本 + 必要的格式转换",把耗时的数据增强、归一化等操作放到外部处理(通过 DataLoader 传入 transform)。这么设计的理由是:__getitem__ 会被多进程并发调用,如果在这里面做太多 CPU 密集型预处理,会直接拖慢数据加载速度。当然,如果你做的是离线预处理(比如把数据提前处理好存成内存映射格式),那放进 __init__ 或单独写预处理脚本会更合理。
2.2 从零实现一个图片分类 Dataset 的完整示例
我拿实际项目中最常见的图片分类场景来演示。假设你的数据目录结构是 data/train/cat/xxx.jpg、data/train/dog/yyy.jpg 这种按类别分文件夹的布局,那么可以这样实现:
python复制import os
from PIL import Image
from torch.utils.data import Dataset
from torchvision import transforms
class ImageFolderDataset(Dataset):
def __init__(self, root_dir, transform=None):
self.samples = [] # 每个元素是 (图片路径, 类别索引)
self.classes = sorted(os.listdir(root_dir))
self.class_to_idx = {cls: i for i, cls in enumerate(self.classes)}
self.transform = transform
for cls in self.classes:
cls_dir = os.path.join(root_dir, cls)
if not os.path.isdir(cls_dir):
continue
for fname in os.listdir(cls_dir):
if fname.lower().endswith(('.jpg', '.jpeg', '.png')):
self.samples.append(
(os.path.join(cls_dir, fname), self.class_to_idx[cls])
)
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
img_path, label = self.samples[idx]
image = Image.open(img_path).convert('RGB')
if self.transform is not None:
image = self.transform(image)
return image, label
这个实现有几个细节值得注意。第一,我在 __init__ 里把所有文件路径和标签一次扫描完成,而不是每次 __getitem__ 才去遍历目录,因为遍历目录是磁盘 IO 操作,频繁执行会严重拖慢节奏。第二,用 os.listdir 前先做排序,保证类别索引稳定,否则每次初始化数据集类别顺序都可能变化,训练时标签就乱了。第三,Image.open 之后加 .convert('RGB') 是为了统一通道数,避免遇到灰度图或带透明通道的 PNG 导致后续张量维度不一致。
使用的时候配合 torchvision 的 transforms 做预处理即可:
python复制transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
dataset = ImageFolderDataset(root_dir='data/train', transform=transform)
2.3 处理文本、表格等非图像数据时要注意什么
图像只是数据中最简单的一种形态。我处理过不少文本分类和表格类数据的项目,Dataset 的实现思路会有一些不同,这里挑关键的点来说。
文本数据最大的特点是样本长度不固定。如果你在 __getitem__ 里返回的是变长的 token 序列,直接丢给 DataLoader 会报错,因为它默认尝试把所有样本堆成一个张量,长度不一致就堆不了。解决方式有两种:一是自己在 __getitem__ 里做 padding,让所有样本长度一致,缺点是 batch 内有效 token 占比低、浪费算力;二是返回原始 (token_ids, attention_mask, label) 这样的字典结构,在 collate_fn 里再按当前 batch 的最大长度动态 padding,这是目前最主流的做法。
表格类数据则要小心特征类型混杂的问题。数值型特征、类别型特征、缺失值同时存在,你在 __getitem__ 里返回的样本可能是由 numpy 数组、整数、浮点数组成的 tuple,DataLoader 的默认 collate_fn 会把它们转成不同的 tensor 类型。我的经验是,尽量在 Dataset 内部就把类型统一好,比如类别特征转成 long 类型的张量,数值特征转成 float32 的张量,缺失值提前填好,避免在 batch 拼接阶段才暴露出类型冲突。
还有一类比较特殊的情况是样本之间有依赖关系,比如序列预测任务里第 t 个样本依赖第 t-1 个样本的输出。这种任务不适合用普通的 Dataset + DataLoader 组合,因为 DataLoader 的采样和 shuffle 会破坏时序。通常的做法是构造一个"窗口化"的 Dataset,让 __getitem__ 返回长度为 window_size 的一段序列,窗口之间的依赖被封装在样本内部,这样既保留了随机采样的灵活性,又不破坏时序逻辑。
3. DataLoader 类核心参数解析
3.1 参数速查与推荐配置
DataLoader 的参数我在项目里基本都试过一遍,这里整理成一张表,方便大家查阅:
| 参数 | 作用 | 我的常用配置 | 备注 |
|---|---|---|---|
| batch_size | 每个 batch 的样本数 | 16~128 | 受 GPU 显存和模型大小影响 |
| shuffle | 每个 epoch 是否打乱数据 | 训练 True,验证 False | 验证集不打乱且 batch 大小可调大 |
| num_workers | 并行加载数据的进程数 | CPU 核数的一半到全部 | 不是越大越好,见 3.2 |
| drop_last | 最后一个不足 batch_size 的 batch 是否丢弃 | 训练 True,验证 False | 防止 batch norm 在极小 batch 上出问题 |
| collate_fn | 自定义样本拼接逻辑 | 由数据形态决定 | 变长数据、目标检测场景必用 |
| sampler | 自定义采样策略 | 类别不平衡时用 WeightedRandomSampler | 与 shuffle 互斥 |
| pin_memory | 是否锁页内存 | GPU 训练 True | 能小幅提升 CPU 到 GPU 的拷贝速度 |
| prefetch_factor | 每个 worker 预取的样本数 | 2~4 | num_workers > 0 时才有效 |
关于 batch_size 的选择,很多人有个误区,觉得 batch 越大训练越快。实际上 batch 太大有两个问题:一是显存不够,强行用梯度累积又增加复杂度;二是在某些任务上大 batch 会导致模型泛化能力下降,需要相应调大学习率。我的经验是从 32 或 64 起步,观察显存占用和 Loss 曲线再做调整。
3.2 num_workers 调参的血泪经验
num_workers 是我见过被误解最深的参数。很多新手以为这个值越大数据加载越快,一上来就设成 32、64,结果程序直接卡死或者内存爆炸。
首先要理解它的机制。num_workers=0 表示数据加载在主进程里同步执行,训练时 GPU 要等 CPU 读完数据,效率极低;num_workers>0 时 PyTorch 会派生多个 worker 进程,它们各自从 Dataset 里取数并维护自己的队列,数据加载和模型训练可以一定程度上并行。但 worker 进程越多,内存开销越大,因为每个 worker 都会复制一份 Dataset 对象。如果你的 Dataset 比较大(比如把所有图片一次性读进内存),设 8 个 worker 就意味着内存占用变成大约 8 份。
我的调参经验是:先看机器有多少个物理核,num_workers 一般设置为物理核数的一半左右,不要超过 CPU 核数。然后在训练时观察 GPU 利用率,如果 GPU 利用率经常掉到 80% 以下且 CPU 没有跑满,适当增加 num_workers;如果发现内存占用异常飙升,说明 worker 开多了。还有一个容易踩的坑是 Windows 系统下多进程数据加载容易报错,需要在主脚本里加 if __name__ == '__main__': 保护,否则 worker 会递归执行主模块导致死循环。
3.3 collate_fn 是解决复杂数据拼接的万能钥匙
我接触的大多数初学者都用默认的 collate_fn,直到遇到变长数据或者多标签数据才发现搞不定。默认的 collate_fn 做的事情很简单:把 batch 里的每个样本(假设是 (data, label) 这样的 tuple)分别堆叠成张量。要求是每个样本的 data 形状完全一致、label 形状完全一致。
真实业务数据经常不满足这个要求,这时候就要自己写 collate_fn。我用一个文本分类的例子来说明。假设每个样本是 (input_ids, attention_mask, label),其中 input_ids 长度可变,那么可以这样处理:
python复制import torch
from torch.nn.utils.rnn import pad_sequence
def collate_fn(batch):
input_ids = [item['input_ids'] for item in batch]
attention_mask = [item['attention_mask'] for item in batch]
labels = torch.tensor([item['label'] for item in batch])
input_ids = pad_sequence(input_ids, batch_first=True, padding_value=0)
attention_mask = pad_sequence(attention_mask, batch_first=True, padding_value=0)
return {
'input_ids': input_ids,
'attention_mask': attention_mask,
'labels': labels
}
pad_sequence 会在 batch 内部按最长序列补齐,padding_value=0 对应 tokenizer 里的 pad_token_id。这样做的优点是每个 batch 的 padding 长度不同,短 batch 不会浪费太多计算。注意此时 __getitem__ 返回的数据必须是长度为 1 的张量序列,不能把不同长度的 python list 直接放进 batch,否则 torch.tensor 转换时会报维度不一致。
还有一个实际经验:collate_fn 里尽量不要做太重的预处理,因为它是在主进程执行的(准确说是数据加载进程),如果里面做了图片解码、文本编码这类耗时操作,会抵消多进程加载的优势。重活应该放在 Dataset 的 __getitem__ 里,让 worker 进程去分担。
4. 实操:从零构建一个完整可复用的数据管线
4.1 场景定义与方案选型
光讲概念容易飘,我拿一个自己近期做过的图像多标签分类项目来串一遍完整流程。这个项目的数据是医疗影像,每张图片可能同时属于多个类别(多标签),图片存储在磁盘上,标注存在一个 CSV 文件里,样本总量大概 20 万张。这个场景有几个关键约束:数据量大不能一次性全读进内存、标签是多标签需要特殊编码、训练时需要数据增强、还要保证每个 epoch 的随机性。
基于这些约束,我做了三个决策。第一,Dataset 里只存图片路径和标签索引,不提前读图片内容,让内存占用保持低位。第二,__getitem__ 里用 PIL 读图 + torchvision 的 transforms 做在线增强,增强操作直接作用在返回样本上,不需要单独事先生成增强数据。第三,DataLoader 开启多进程加载,collate_fn 保持默认(因为图片经过 Resize 后尺寸统一、标签是定长 one-hot 向量,默认拼接逻辑够用)。
4.2 完整代码实现与逐步解说
Dataset 部分我按多标签场景实现如下:
python复制import pandas as pd
from PIL import Image
from torch.utils.data import Dataset
from torchvision import transforms
class MultiLabelImageDataset(Dataset):
def __init__(self, csv_path, img_dir, transform=None):
self.df = pd.read_csv(csv_path)
self.img_dir = img_dir
self.transform = transform
# 假设 CSV 里有 image_id 列和若干标签列,标签值 0/1
self.label_cols = [c for c in self.df.columns if c not in ['image_id']]
def __len__(self):
return len(self.df)
def __getitem__(self, idx):
row = self.df.iloc[idx]
img_path = f"{self.img_dir}/{row['image_id']}"
image = Image.open(img_path).convert('RGB')
if self.transform is not None:
image = self.transform(image)
label = self.df.iloc[idx][self.label_cols].values.astype('float32')
return image, torch.from_numpy(label)
这里有几个我踩过坑后补上的细节。label_cols 在 __init__ 里提前算好,避免每个样本都重新算一次列名列表。label 转成 float32 是因为多标签分类通常用 BCEWithLogitsLoss,它的目标张量要求浮点类型,而不是整数类型,如果你用整数 one-hot 会报类型错误。图片读进来后转 RGB 是为了统一通道数,医疗影像中偶发灰度图的情况比较多,不加这行就可能在后续模型前向时报维度不匹配。
train/val 数据集的构建和 DataLoader 的配置如下:
python复制from torch.utils.data import DataLoader, random_split
transform_train = transforms.Compose([
transforms.Resize((256, 256)),
transforms.RandomCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
transform_val = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
full_dataset = MultiLabelImageDataset(
csv_path='data/train_labels.csv',
img_dir='data/images',
transform=transform_train
)
train_size = int(0.9 * len(full_dataset))
val_size = len(full_dataset) - train_size
train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])
random_split 返回的子集是一个包装视图,它会记住原始 Dataset 的映射关系,所以这里有个小坑:如果你用 random_split 切分后再分别对两个子集应用不同的 transform,直接改子集的 transform 属性是没用的,因为子集内部的 __getitem__ 调用的还是原始 Dataset 的 transform。正确做法是把 transform 的切换放在原始 Dataset 里,或者干脆不依赖 random_split,用 Subset + 自定义实现。我在项目里倾向于用索引列表做切分,然后把两个数据集分别构建,代码更直观:
python复制from torch.utils.data import Subset
import numpy as np
indices = np.arange(len(full_dataset))
rng = np.random.default_rng(42)
rng.shuffle(indices)
train_idx = indices[:train_size]
val_idx = indices[train_size:]
train_dataset = Subset(MultiLabelImageDataset(..., transform=transform_train), train_idx)
val_dataset = Subset(MultiLabelImageDataset(..., transform=transform_val), val_idx)
4.3 数据增强与归一化的执行顺序问题
数据增强的编排顺序看似小事,实际影响很大。我在项目里的基本原则是:几何变换(翻转、旋转、裁剪)在前,像素变换(颜色抖动、归一化)在后,ToTensor 放在它们之间。原因是 PIL 图像和 numpy 数组的操作接口不同,很多 torchvision 的 transform 只接受 PIL Image 类型;ToTensor 会把 HWC 的 PIL 图转成 CHW 的 float 张量,并且把像素值从 [0, 255] 缩放到 [0, 1],所以它的位置决定了后续操作是面向图像还是面向张量。
举一个因为顺序不对导致 Bug 的真实例子。之前有个同事把 Normalize 写在了 ToTensor 之前,结果 Normalize 对 PIL 图像直接报了 TypeError,因为 PIL Image 不支持张量减法。还有一次把 RandomCrop 放在 Resize 之前,导致每个 epoch 裁剪区域完全一致,数据增强形同虚设,模型训练到后面 Loss 怎么都降不下去。这些问题的排查方法很简单,打印一个 batch 的数据形状和数值范围,一眼就能发现问题。
另外,验证集和测试集不要使用随机增强,只保留 Resize、ToTensor、Normalize 这类确定性的预处理。否则验证集的评价指标每次运行都会因为随机性而波动,模型调参时很难判断是改动生效了还是随机种子造成的。
5. 常见问题与排查技巧实录
5.1 问题速查表
我在各种群里看到的数据管线问题,八成以上都集中在下面这些场景,整理成表格方便大家对照:
| 问题现象 | 可能原因 | 解决方案 |
|---|---|---|
| 训练进程卡死无响应 | num_workers 过大,或 Windows 缺少 main 保护 | 减小 num_workers;添加 if __name__ == '__main__': |
报错 IndexError: index out of range |
Dataset __len__ 与实际样本数不一致 |
检查文件列表长度与最大索引的关系 |
报错 TypeError: default_collate |
batch 内样本形状或类型不一致 | 在 collate_fn 里动态处理,或统一预处理 |
报错 RuntimeError: CUDA out of memory |
batch_size 过大或数据加载内存溢出 | 减小 batch_size;开启 pin_memory 但控制 worker 数 |
| GPU 利用率持续低于 50% | 数据加载成为瓶颈 | 增大 num_workers;减少 __getitem__ 里的耗时操作;用 prefetch_factor |
| 每个 epoch 的训练结果差异巨大 | 验证集使用了随机数据增强 | 验证/测试集只用确定性 transform |
| 样本类别极端不平衡 | 默认 sampler 均匀采样 | 使用 WeightedRandomSampler,按类别反比加权 |
其中最后一个问题我多说两句。类别不平衡在真实业务里非常普遍,用 WeightedRandomSampler 是最简单的缓解手段。它的原理是给每个样本赋予一个采样权重,权重高的样本被抽到的概率更大。权重的计算方式我做了一步简化版:
python复制from torch.utils.data import WeightedRandomSampler
import torch
labels = train_dataset.dataset.df[train_dataset.dataset.label_cols].values
class_counts = labels.sum(axis=0)
weights_per_class = 1.0 / class_counts
sample_weights = labels @ weights_per_class
sampler = WeightedRandomSampler(sample_weights, num_samples=len(sample_weights), replacement=True)
注意 WeightedRandomSampler 和 shuffle=True 是互斥的,因为 sampler 本身就是一种采样策略,两者同时设置会报错。另外设置 sampler 后 DataLoader 内部的索引生成逻辑完全交给 sampler,如果 num_samples 设置不当,可能导致一个 epoch 内的样本数和预期不一致。
5.2 性能优化与内存控制的进阶心得
数据管线的性能优化,我总结了一个优先级:先看 GPU 利用率,再定位瓶颈,最后针对性优化,不要一上来就盲目调参。常用的观察方法是 nvidia-smi 查看 GPU 利用率,如果模型很小而 GPU 利用率仍然低,基本可以断定是数据加载跟不上。
常见优化手段按收益从高到低排列:把 __getitem__ 里的重复运算提出来(比如提前解析好路径列表、预计算标签张量);图片解码换成更快的后端或者预先缩小到合适尺寸再存盘;num_workers 配合 prefetch_factor 一起调整;使用 pin_memory=True 减少 CPU 到 GPU 的拷贝开销。还有一个容易被忽略的点:如果你的数据增强完全确定且可以离线完成,建议离线预处理一次存成内存映射格式,训练时直接读 mmap,能把数据加载时间从秒级降到毫秒级。
内存控制方面,最激进的做法是在 Dataset 里做"懒加载":__init__ 只记录元信息,__getitem__ 才真正读文件。这在 20 万张图片的场景下是必须的,否则 8 个 worker 每个都持有完整的数据集副本,内存直接爆掉。如果数据集已经小到能全部放内存,反而建议一次性读进来用共享内存传递,比每个 worker 各自读磁盘快得多。具体什么时候切换策略,取决于你的内存容量和数据总量,我一般以"数据总量不超过物理内存的 30%"作为参考线。
5.3 两个容易忽视的细节问题
最后说两个我最近踩过、文档里又不怎么强调的细节。
第一个是 DataLoader 的默认 collate_fn 对字典类型的样本处理。当 __getitem__ 返回 dict 时,默认 collate 会对 dict 里每个 key 分别做拼接,所以要求每个 value 都是可拼接的张量或标量。一旦某个 key 的 value 是一个可变长度的 python list 或字符串,就会报错。如果你不想为了一个字段专门写 collate_fn,可以在 __getitem__ 里把所有字段都转成定长张量,虽然不是最优方案但能快速跑通。
第二个是随机种子的设置问题。很多人发现同样一个 seed,每次训练的数据顺序还是不一样。原因是 DataLoader 的 worker 进程有自己的随机状态,它们不会继承主进程设置的 numpy/pytorch 随机种子,所以即使你设置了 torch.manual_seed(42),多进程加载下的 shuffle 顺序也无法完全复现。如果想要严格的实验可复现性,需要把每个 worker 的随机种子也固定下来,可以通过给 DataLoader 传 generator=torch.Generator().manual_seed(42) 来实现一部分。但要注意,分布式训练时这个问题的复杂度会再上一个台阶,那种场景下建议把数据切分的逻辑从 DataLoader 抽出来,自己在训练流程里控制。
我在实际项目里逐渐养成的一个习惯是,写完 Dataset 之后先不急着接 DataLoader,直接对一个样本做可视化、对标签做统计,确认数据正确了再接训练循环。看似多花了十分钟,但能省下大量排查"模型不收敛到底是不是数据问题"的时间。数据管线是深度学习里最不性感、却最值得下功夫的部分,把 Dataset 和 DataLoader 吃透,你后面做任何复杂数据处理都会顺手很多。
