先说个真实的体感:我见过太多新手拿到 NLP 项目,第一件事就是写 for i in range(len(texts)): 一条条把样本喂给模型。代码跑起来也能跑,但一到 batch 训练、数据量上到几万条就原形毕露——要么内存爆掉,要么 GPU 利用率只有十几。问题基本都出在没搞懂 PyTorch 数据管线这两兄弟:Dataset 和 DataLoader。
这篇博文会从底层机制讲起,配合一个完整的 NLP 文本分类实战案例,把这两者的设计逻辑、参数细节、常见坑位全部拆开揉碎。无论你是刚装好 PyTorch 准备跑第一个模型,还是已经写了不少训练脚本但总觉得数据加载这块“差点意思”,这篇文章都是按着你能直接上手的标准来写的。
1. 为什么说 Dataset 和 DataLoader 是训练循环的“地基”
1.1 新手最容易踩的误区:拿着原生数组硬怼
很多人的第一个 PyTorch 模型长这样:加载数据用 pd.read_csv,然后转成 numpy 数组,再 torch.tensor() 一下,全部塞进内存,接着 model(data) 开训。
数据量小的时候(几百条),这种方法完全没问题。但一旦进入真实项目,问题就全冒出来了:
- 内存失控:几万条新闻文本,每条几百到几千字,全部转成 tensor 存内存,显存和内存一起报警。
- 没有 shuffle:自己写
np.random.shuffle(index),写了半天还容易出错,而且每个 epoch 都要重新 shuffle,代码越来越乱。 - batch 逻辑全靠手写:
for i in range(0, len(data), batch_size)这种切片操作看着简单,但遇到 NLP 里最常见的“一个 batch 内文本长度不一致”就抓瞎,还要自己写 padding 逻辑。 - 没有多进程加速:数据预处理(比如分词、转 ID)全在主进程里同步执行,GPU 在等 CPU,训练速度惨不忍睹。
Dataset 和 DataLoader 就是专门解决这几类问题的。它们不只是一个简单的封装工具,而是 PyTorch 整个数据加载体系的根基。
1.2 Dataset 是仓库,DataLoader 是调度员
我习惯用一个类比来理解这两者的分工:Dataset 是仓库,DataLoader 是调度员。
Dataset 负责搞清楚“仓库里有多少件货(__len__)”,以及“给定货架编号,把对应那件货取出来(__getitem__)”。它不关心货怎么运、按什么顺序运、一次运多少,那是调度员的事。
DataLoader 负责调度:按什么顺序取货(shuffle、sampler)、一次取多少件打包(batch_size)、怎么把散装货物打包成适合运输的规格(collate_fn)、用几个搬运工同时干活(num_workers)、要不要用专线运输提高效率(pin_memory)。
这个分工非常重要。它意味着你可以随时换仓库(比如从 CSV 换成数据库、换成 JSONL、换成 HuggingFace 数据集),而调度逻辑完全不用动;反过来,你也可以在同一个 Dataset 上换不同的调度策略(比如训练时 shuffle、验证时不 shuffle),而取样逻辑完全不用动。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 手写 Dataset 前必须搞懂的底层逻辑
2.1 getitem 与 len:两个方法撑起整个数据管线
PyTorch 的 torch.utils.data.Dataset 是一个抽象类,它本身几乎不实现任何东西,全靠子类实现两个核心方法:
__len__(self):返回数据集样本总数。DataLoader 靠它知道一共有多少样本,从而计算一个 epoch 有多少个 batch,也靠它生成默认的采样索引。__getitem__(self, index):根据 index 返回第 index 个样本。这是整个数据管线里最核心的代码,PyTorch 的多进程加速就是靠不断的并发调用这个方法来完成的。
一个最朴素的 Dataset 长这样:
python复制from torch.utils.data import Dataset
class MyDataset(Dataset):
def __init__(self, texts, labels):
self.texts = texts
self.labels = labels
def __len__(self):
return len(self.texts)
def __getitem__(self, idx):
return self.texts[idx], self.labels[idx]
就这么简单。你可能会问:这也太没技术含量了,跟直接拿列表索引有什么区别?
区别在于,__getitem__ 的返回值可以是任意 Python 对象——字符串、字典、变长列表、嵌套结构,都行。DataLoader 会把这些返回值收集起来,通过 collate_fn 统一加工成 tensor。这就给了你极大的自由:原始数据往往不是规整的矩阵,而是在 __getitem__ 里做清洗、分词、转 ID,最后返回到统一的、可被装进 batch 的格式。
2.2 Map式还是 Iterable式:选错的代价
PyTorch 支持两种 Dataset 接口风格:Map-style(映射式) 和 Iterable-style(可迭代式)。
上面写的那个 MyDataset 就是 Map-style。核心特征是你可以通过 dataset[i] 随机访问任意样本。DataLoader 默认通过 Sampler 生成索引序列,然后逐个传给 __getitem__。因为支持随机访问,所以 shuffle 非常好实现——本质就是打乱索引顺序。
Iterable-style 是实现了 __iter__ 的 Dataset,类似于 Python 的生成器。它适合数据没法随机访问的场景,比如从网络流、数据库游标、TensorFlow Record 文件里顺序读取数据。
选错的代价是什么?如果你有 10 万条文本文件路径,用 Map-style 完全没问题——__getitem__ 里按需读取文件即可。但如果你面对的是一个 5TB 的流式日志文件,无法索引到具体某一行偏移量,或者数据源本身只支持顺序读取,那 Map-style 就玩不转了。此时强行用 Map-style,要么把整个数据集读进内存,要么就得写一堆复杂的缓存逻辑。
我的建议是:95% 的场景选 Map-style。不是因为 Iterable 不好,而是它有几个麻烦:
- 它不支持
len(),所以 DataLoader 不知道总样本数,部分功能(如进度条)会受限。 - shuffle 非常麻烦,需要自己实现 shuffle buffer,远不如 Map-style 的索引打乱简单。
- 多进程 worker 下,每个 worker 都会拿到同一个迭代器,需要自己处理样本分配逻辑,不然同一个 batch 会被多个 worker 重复读取。
只有当你确实遇到“无法随机访问”的数据源时,才考虑 Iterable-style。
2.3 预处理放在哪里才合理:init 与 getitem 的分工
新手写 Dataset 最常见的错误,是把所有预处理全部堆在 __init__ 里。我见过有人把分词、去停用词、构建词表、文本转 ID 全在 __init__ 做完,结果初始化一次要等几分钟,而且内存占用爆炸。
正确分工原则是:__init__ 只负责存储“轻量级”的原始引用,重活放在 __getitem__ 里按需执行。
举个例子,假设你的数据是 5 万条新闻文本,存在 CSV 里。合理做法:
python复制class NewsDataset(Dataset):
def __init__(self, csv_path, vocab):
# 这里只存文件路径,不读全部数据进内存
self.data = pd.read_csv(csv_path) # 如果你的 CSV 不是特别大,读进来也没事
self.vocab = vocab
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
text = self.data.iloc[idx]['content']
label = self.data.iloc[idx]['label']
# 分词、转 ID 在 getitem 里做,每个样本按需处理
tokens = jieba.lcut(text)
ids = [self.vocab.get(t, self.vocab['<unk>']) for t in tokens]
return ids, label
这样做的原因有两个:
- 内存友好:如果每条新闻原文有几 KB,5 万条也就几百 MB,读进来问题不大;但如果你有 100 万条,或者每条是几 MB 的文档,全读进内存就真的要命了。
__getitem__里按需读取可以大幅降低内存峰值。 - 多进程并行:DataLoader 开
num_workers > 0后,每个 worker 是独立进程,它们各自执行自己的__getitem__。重活在 worker 进程里并行执行,能充分利用多核 CPU。
但也不是所有东西都要放在 __getitem__ 里。凡是所有样本共享、且计算一次就能复用的,比如词表(vocab)、预训练 embedding 矩阵、停用词集合,都应该在 __init__ 或者全局只构建一次。
还有一个中间地带:像文本清洗(去掉 HTML 标签、特殊符号)这种操作,如果原始数据本身不大,也可以选择在预处理脚本里提前做一次,存成干净的文件,然后 Dataset 只负责加载干净数据。我的经验是:能提前做的预处理尽量提前做,__getitem__ 里只保留那些依赖样本个体特征的转换——比如文本长度不一导致的动态 padding 就是无法提前做的,只能在 batch 层面处理。
3. DataLoader 参数逐个拆解:别再只调 batch_size
很多人在实际使用中只关心 batch_size 和 shuffle,其他参数基本不管。但 DataLoader 的参数每一个都有明确的场景和代价,理解它们能让你的训练效率上一个大台阶。
3.1 sampler 与 shuffle:数据顺序由谁说了算
shuffle=True 的底层其实是换了一个采样器。DataLoader 的采样流程是:Sampler 生成索引序列 → 按索引调用 __getitem__ → 按 batch_size 打包。
默认情况下,shuffle=False 用 SequentialSampler,按 0、1、2、3... 的顺序依次取;shuffle=True 用 RandomSampler,每次迭代前对所有索引做一次随机打乱。
但有时候,默认采样器不够用。最典型的场景是样本不均衡的文本分类:假设你的语料里“科技”类有 4 万条,“体育”类只有 2000 条,直接用 RandomSampler,每个 batch 里体育样本凤毛麟角,模型训练时梯度被科技类主导。
这时候可以用 WeightedRandomSampler,给每个样本分配一个权重,让抽样时概率与权重成正比,从而让少数类被抽到的频率更高:
python复制from torch.utils.data import WeightedRandomSampler
labels = dataset.data['label'].values
class_counts = np.bincount(labels)
weights = 1.0 / class_counts[labels] # 每一类样本的权重 = 1/该类的总数
sampler = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)
dataloader = DataLoader(dataset, batch_size=32, sampler=sampler)
注意,sampler 和 shuffle 参数是互斥的,传了 sampler 就不能传 shuffle=True,因为 sampler 本身已经定义了采样顺序。另外还有一个 batch_sampler 参数,它直接决定每个 batch 由哪些索引组成,比 sampler 更底层。BatchSampler 把 sampler 产生的索引按 batch_size 切分成一个个小批次。一般情况下你不用直接碰它,但如果你需要实现特殊的 batch 组成逻辑(比如把相似长度的样本分到同一 batch,即 bucket batching),就需要自定义 batch_sampler。
3.2 collate_fn:NLP 里最省不掉的定制逻辑
collate_fn 可能是所有 DataLoader 参数里,对 NLP 任务最重要的一个。
默认情况下,DataLoader 的 collate_fn 假设 __getitem__ 返回的每个样本是一个(或一组)tensor,它把这些 tensor 打包成一个新的 tensor 作为 batch 的第一个维度。也就是说,默认行为是“沿 batch 维度堆叠”。这要求所有样本的第一维大小一致。
但 NLP 数据天然不满足这个假设。句子有长有短,分词后 token 数各不相同,强行堆叠会报维度不匹配的错误。这时候就需要 collate_fn 出场:把变长的样本序列调整成同一个 batch 内等长的 tensor。
最简单的做法是固定最大长度(比如每条截断到 128 tokens),padding 到 128。但更好的做法是动态 padding:只 padding 到当前 batch 内最长句子的长度,而不是整个数据集的最大长度。这样既避免了过长的填充浪费计算,又保证了 batch 内维度一致。
关于这部分,我会在下一章的实战里给出完整的 collate_fn 实现。
3.3 num_workers、pin_memory 与 prefetch_factor:性能三件套
这三个参数是 DataLoader 性能调优的核心。
num_workers:数据加载的子进程数。num_workers=0 表示在主进程里同步加载数据,此时如果 __getitem__ 里有耗时的分词操作,GPU 就只能干等着。num_workers=4 表示启动 4 个子进程并行加载数据,能显著缩短数据准备时间。但注意,不是越大越好:
- 进程数超过 CPU 核心数不仅不会提速,反而会因为进程切换耗尽 CPU。
- Windows 上
num_workers > 0要求训练代码放在if __name__ == '__main__':的守卫里,否则会疯狂报错。 - 每个 worker 会复制一份 Dataset 对象(通过 pickle 序列化传给孩子进程),如果 Dataset 里存了大量数据,复制开销和内存占用都会很大。
pin_memory=True:这个参数的作用是把数据加载到内存时,锁定在页锁定内存(pinned memory)里。GPU 从页锁定内存拷贝数据到显存的速度,比从普通内存拷贝快很多。简单说,如果你的数据最终要交给 GPU 训练,pin_memory=True 基本是白捡的性能提升。代价是占用的内存不能被操作系统换出,所以内存紧张时反而可能拖慢系统。
prefetch_factor:控制每个 worker 预先加载多少个 batch 的数据。默认值是 2,即每个 worker 提前准备 2 个 batch。当你的 __getitem__ 耗时较长,而 collate_fn 相对较快时,适当调大 prefetch_factor(比如 4 或 8)可以让数据管线更平滑。但要注意,prefetch_factor 值越大,每个 worker 占用的内存也越高。
我个人常用的配置是:num_workers=4、pin_memory=True、prefetch_factor=4。在 CPU 负载可接受范围内,这个配置基本能把数据加载时间压到训练时间的一个零头。
4. NLP 实战:文本分类从原始样本到可训练数据
这一章我们用一份模拟的新闻数据集,走一遍完整的 NLP 数据管线:从 CSV 原始文本,到可训练的 batch 数据。完整代码可以跑通,你可以在此基础上替换成自己的数据集。
4.1 数据准备与分词:别把预处理全塞进 Dataset
假设我们有一份 news.csv,两列:content 是新闻正文,label 是类别(整数编码,从 0 到 3)。
第一步是先做全局预处理。为什么不让 Dataset 自己处理?因为分词、去停用词这些操作的结果,在 Dataset 里会反复执行——每次调用 __getitem__,同一个样本的同一段文本都会被重新分词一次,浪费计算资源。
正确做法是:提前跑一次预处理,把文本转成词 ID 序列,存成新的文件。 Dataset 加载时的逻辑就变成“读取词 ID 序列 + 转成 tensor”,非常轻量。
python复制import pandas as pd
import jieba
df = pd.read_csv('news.csv')
# 简单清洗:去空格、去换行
df['content'] = df['content'].str.replace(r'\s+', '', regex=True)
# 分词(这里用 jieba,实际项目中也可以用更快的 pkuseg)
df['tokens'] = df['content'].apply(lambda x: jieba.lcut(x))
# 构建词表
from collections import Counter
word_counter = Counter()
for tokens in df['tokens']:
word_counter.update(tokens)
# 按频率排序,保留出现次数 >= 2 的词
vocab = {'<pad>': 0, '<unk>': 1}
for word, freq in word_counter.most_common():
if freq < 2:
break
vocab[word] = len(vocab)
# 文本转词 ID
df['ids'] = df['tokens'].apply(lambda tokens: [vocab.get(t, vocab['<unk>']) for t in tokens])
# 保存预处理后的结果,后续 Dataset 直接加载
df[['ids', 'label']].to_pickle('news_processed.pkl')
注意这里我用 to_pickle 保存,因为列表列在 CSV 里存会很麻烦。实际项目中你也可以用 parquet、jsonl 等格式。
4.2 构建词表与样本编码
上面代码里,词表构建有个小细节值得展开:
<pad>固定为 0,专门用于填充。填充 token 的 id 越靠前越方便,因为很多模型初始化时会把 padding id 对应的 embedding 置零。<unk>固定为 1,凡是未登录词都映射到这个 id,避免维度爆炸。- 最少出现次数
freq < 2就截断,这是 NLP 里的常见操作:只出现一次的单词基本都是噪声,保留它只会增长词表大小、增加过拟合风险,却几乎不带来任何有效信息。
4.3 自定义 Dataset 与 Dynamic Padding 的 collate_fn
现在加载预处理后的文件,写 Dataset 和 DataLoader:
python复制import torch
from torch.utils.data import Dataset, DataLoader
class TextDataset(Dataset):
def __init__(self, ids_list, labels):
self.ids_list = ids_list
self.labels = labels
def __len__(self):
return len(self.labels)
def __getitem__(self, idx):
# 返回一个元组(词ID序列,标签)
# 注意:这里不转 tensor,因为序列长度不一,collate_fn 里统一处理
return torch.tensor(self.ids_list[idx], dtype=torch.long), torch.tensor(self.labels[idx], dtype=torch.long)
然后是最关键的 collate_fn:
python复制def collate_fn(batch):
"""
batch: 一个包含 batch_size 个元组的列表,每个元组是 (ids, label)
"""
ids_list = [item[0] for item in batch]
labels = torch.stack([item[1] for item in batch])
# 动态 padding:当前 batch 内最大长度
max_len = max(len(ids) for ids in ids_list)
padded_ids = torch.zeros(len(ids_list), max_len, dtype=torch.long)
attention_mask = torch.zeros(len(ids_list), max_len, dtype=torch.long)
for i, ids in enumerate(ids_list):
length = len(ids)
padded_ids[i, :length] = ids # 前面放真实 token
attention_mask[i, :length] = 1 # 有效位置标记为 1
return padded_ids, attention_mask, labels
为什么 collate_fn 里要用 torch.zeros 统一创建 token 的 0 号位来池化?因为 <pad> 的 id 就是 0,直接用初始化好的 zero tensor 天然完成了 padding。这是把 <pad> 放在 index 0 的一个额外好处。
attention_mask 是给模型看的:哪些位置是真实 token(1),哪些位置是 padding(0)。Transformer 架构的模型在计算注意力时,会把 padding 位置 mask 掉,不参与注意力计算。
4.4 跑通一个完整训练循环
数据管线搭好后,训练循环就非常清爽了:
python复制import torch.nn as nn
from torch.utils.data import DataLoader
# 加载预处理数据
import pandas as pd
df = pd.read_pickle('news_processed.pkl')
dataset = TextDataset(df['ids'].values, df['label'].values)
train_loader = DataLoader(
dataset,
batch_size=32,
shuffle=True,
collate_fn=collate_fn,
num_workers=2,
pin_memory=True,
drop_last=True
)
# 一个极简的模型(真实 NLP 任务请换 Embedding + LSTM/Transformer)
class TinyTextModel(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.fc = nn.Linear(embed_dim, num_classes)
def forward(self, input_ids, attention_mask):
# input_ids: (batch, seq_len)
embedded = self.embedding(input_ids) # (batch, seq_len, embed_dim)
masked = embedded * attention_mask.unsqueeze(-1)
pooled = masked.mean(dim=1) # 对真实 token 取平均(padding 位置为 0,不影响均值)
return self.fc(pooled)
model = TinyTextModel(vocab_size=len(vocab), embed_dim=128, num_classes=4)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
for epoch in range(3):
total_loss = 0
for batch_idx, (input_ids, attention_mask, labels) in enumerate(train_loader):
optimizer.zero_grad()
logits = model(input_ids, attention_mask)
loss = loss_fn(logits, labels)
loss.backward()
optimizer.step()
total_loss += loss.item()
if batch_idx % 100 == 0:
print(f'Epoch {epoch}, Batch {batch_idx}, Loss: {loss.item():.4f}')
print(f'Epoch {epoch}, Avg Loss: {total_loss / len(train_loader):.4f}')
这里 drop_last=True 的用意:如果最后一个 batch 的样本数不足 batch_size,会拖慢训练且影响梯度稳定性,所以直接丢掉。如果你的模型要求严格等长的 batch(比如某些涉及 batch norm 的实现),这个参数也很重要。
5. 踩坑记录与性能调优实录
5.1 把全部数据读进 init 导致的内存爆炸
我见过一个真实案例:某 NLP 项目的 Dataset 写成了这样:
python复制def __init__(self, csv_path):
df = pd.read_csv(csv_path)
self.data = []
for _, row in df.iterrows():
tokens = jieba.lcut(row['content'])
ids = [vocab.get(t, 1) for t in tokens]
self.data.append((ids, row['label']))
数据量 30 万条时,这段代码在 __init__ 阶段就耗了将近 20 分钟,内存占用直逼 16GB。而且一旦 num_workers > 0,每个子进程都会复制一份这个 Dataset,内存直接翻好几倍。
解决办法:__init__ 里只存文件路径或 DataFrame 引用,分词和 ID 转换放到 __getitem__ 里。如果觉得分词太慢,用预处理脚本先把结果缓存成 pickle 或 parquet。
5.2 Worker 数量与 Windows 平台的诡异报错
在 Windows 上跑 DataLoader(..., num_workers=4),如果代码没有放在 if __name__ == '__main__': 的保护里,会直接报 RuntimeError: An attempt has been made to start a new process before the current process has finished its bootstrapping phase。
这不是 PyTorch 的 bug,而是 Windows 下多进程实现方式(spawn)与 Linux 不同(fork)导致的。解决办法很简单:所有包含 DataLoader 实例化和训练循环的代码,都放到 main() 函数里,并在脚本底部调用 if __name__ == '__main__': main()。
如果你在 Jupyter Notebook 里跑,多进程还容易遇到各种“spawn 进程重新执行 notebook 代码”的诡异问题。我个人的建议是:Notebook 里做调试和数据探索时,用 num_workers=0;正式跑训练脚本时,再配多个 worker。 省下的时间远比那点加速重要。
5.3 样本不均衡时的 sampler 选择
如果你用默认的 shuffle=True 训练一个样本严重不均衡的分类任务,可能出现一个 batch 里全是“科技”类、下一个 batch 里才偶尔有几条“体育”类的情况。模型会频繁被多数类更新,少数类几乎学不到。
这时候有两个选择:
- 用
WeightedRandomSampler,如上文所示,每个样本的采样概率与类别频率成反比。但要注意,这相当于人为改变了训练分布,训练集和验证集的类别分布会不一致,需要观察验证指标是否真的变好了。 - 用
Oversampler的思路,在__getitem__里对少数类样本做文本增强(同义词替换、回译),增大少数类样本量。这是更“数据驱动”的做法,效果通常更稳。
实操上,我一般先跑一版 WeightedRandomSampler 看模型能否收敛,如果收敛没问题,再考虑数据增强提升泛化。
5.4 实测对比:不同配置下的数据加载耗时
我用一份 5 万条、平均长度 120 词的中文新闻数据做了一次简单对比测试,环境是 8 核 CPU + RTX 4090,训练 1 个 epoch(30 个 batch 之后停止计时):
| 配置 | 平均每个 batch 数据加载耗时 | 说明 |
|---|---|---|
| num_workers=0, 无 pin_memory | 约 120 ms | 数据加载完全阻塞训练 |
| num_workers=4, 无 pin_memory | 约 35 ms | 多进程并行显著提速 |
| num_workers=4, pin_memory=True | 约 12 ms | 页锁定内存加速 GPU 拷贝 |
| num_workers=4, pin_memory=True, prefetch_factor=8 | 约 9 ms | 预取进一步掩盖数据加载延迟 |
注意,这个测试里 __getitem__ 只做简单的 ID 序列读取和转 tensor,没有做实时分词。如果 __getitem__ 里有实时分词,num_workers 的加速效果会更加明显。
最终建议配置:8 核及以上 CPU,num_workers=4,pin_memory=True,prefetch_factor=4~8。如果你的数据加载依然是瓶颈,把 num_workers 调到 6 或 8 试一下,但不要超过物理核心数。
最后再分享一个我在实际项目中悟出来的经验:不要一上来就追求数据加载的极致性能。 先把 Dataset 和 DataLoader 跑通,确认模型能正常训练,再逐步调参优化。很多时候你以为的“数据加载太慢”,其实是 __getitem__ 里埋了太多的重复计算(比如对同一个样本反复分词)。把预处理往前面挪一挪,比调 num_workers 效果大得多。数据管线这东西,架构对了比什么都重要。
