PyTorch Sampler完全指南:从原理到实战,解决类别不均衡与分布式训练
2026/9/23 12:39:14 网站建设 项目流程

1. 为什么每个PyTorch新手都会在Sampler上栽跟头

1.1 一次"数据顺序错乱"事故的排查全过程

前阵子帮一个朋友调试训练脚本,现象非常诡异:同一个模型、同一份数据,在A机器上跑得好好的,换到B机器上loss曲线就开始抖动,验证集指标也忽高忽低。代码逐行比对了两遍,模型结构、优化器、学习率调度、数据增强逻辑全都一样,最后才发现问题出在DataLoader的sampler参数上——他在B机器上给DataLoader传了一个自定义的sampler,却没有关掉默认的shuffle逻辑,两个采样逻辑叠加在了一起,产生了预料之外的索引顺序。

这种问题在PyTorch的Sampler使用中太常见了。我见过不少同学在DataLoader里看到参数,只知道shuffle=True能打乱数据,却完全没有意识到背后真正干活的是Sampler。一旦需要处理类别不均衡、做难例挖掘、搞分布式训练,就会发现默认的shuffle远远不够,必须对采样过程有完全的控制权。

所以这篇博客想做的就是从底层机制到实战踩坑,把Sampler这件事彻底讲透。不管你是在做CV分类、NLP序列标注,还是在做多卡分布式训练,只要你在用PyTorch的DataLoader加载数据,理解Sampler就能帮你解决三类核心问题:数据顺序怎么控制、样本怎么按权重采样、多卡之间怎么切分数据。

1.2 先搞清楚DataLoader加载数据的三个角色

很多人的误区是把Dataset和数据的"取用顺序"绑在一起看。实际上,PyTorch把数据加载这件事拆成了三个独立的责任方:

  • Dataset:只知道"我有多少样本"和"给我索引i,我返回第i个样本"。它不关心你按什么顺序取。
  • Sampler:只负责"产生索引序列"。它决定了DataLoader每次去Dataset里取哪些样本、以什么顺序取。
  • DataLoader:拿着Sampler吐出来的索引序列,逐个交给Dataset取样本,再组装成batch,顺便做多进程预取和collate。

换句话说,Sampler是"取数策略"的决策者,DataLoader是"取数动作"的执行者。你把Sampler理解成一个迭代器:每次__next__它就给一个索引,DataLoader就把这个索引交给Dataset去取数。shuffle=True这个参数的本质,只是PyTorch在内部帮你自动创建了一个RandomSampler,并没有任何"魔法"。

2. Sampler的运行机制:从索引到batch的幕后链路

2.1 两个必须实现的接口:iter__和__len

要想弄懂Sampler,最直接的办法就是看它的抽象定义。PyTorch的torch.utils.data.Sampler是一个很薄的基础类,核心只有两个方法:

class Sampler: def __init__(self, data_source): self.data_source = data_source def __iter__(self): raise NotImplementedError def __len__(self): raise NotImplementedError

任何Sampler都必须实现__iter__和__len__。__iter__返回一个可迭代对象,每次迭代抛出一个整数索引,DataLoader的迭代循环里就是不停地从这个迭代器里取索引;__len__返回这个Sampler总共会产生多少个索引,这个值最终决定了一个epoch内DataLoader会"看到"多少样本。

这里有个初学者容易忽略的点:Sampler返回的索引数量不一定等于Dataset的长度。典型的例子是WeightedRandomSampler,你给它传了weights和num_samples,它产生的索引数量是由num_samples决定的,而不是数据集的样本数。这带来一个很重要的推论:一个epoch的步数不是由Dataset决定的,而是由Sampler决定的。很多人想控制每个epoch训练多少步,改了半天DataLoader参数没效果,其实应该直接改Sampler。

2.2 内置Sampler的适用场景与实现逻辑

PyTorch内置了六个常用Sampler,逐个拆开看它们的定位:

Sampler类型产生的索引序列典型场景
SequentialSampler0, 1, 2, ..., n-1 顺序不变验证集评估、测试集推理
RandomSampler随机打乱的索引常规训练(shuffle=True)
SubsetRandomSampler从指定子集中随机取手动划分训练/验证集
WeightedRandomSampler按权重概率放回采样类别不均衡、样本重要度不同
BatchSampler把上述Sampler的索引打包需要自定义batch大小或batch内结构
DistributedSampler按rank切分数据单机多卡/多机多卡训练

每个Sampler的内部逻辑其实非常简单,我逐个说下它们适用但不为人知的细节。

SequentialSampler不说了,就是range(len(data_source))。RandomSampler的实现里有一个特别容易被忽略的参数:generator。PyTorch很多随机操作都支持传入一个torch.Generator对象,控制随机数种子。如果你在多个机器之间做可复现实验,或者做强化学习想固定环境随机性,一定要显式传入generator,否则它内部用的是全局默认的随机数生成器,不同进程之间没法独立控制。

SubsetRandomSampler接收一个indices列表,然后对这个列表做随机打乱。它常被用来做数据集划分,但有个坑:如果两个不同的任务要共享同一个数据集的随机拆分,就必须保证传进去的indices顺序一致,否则两个任务拿到的训练/验证子集就完全不一样了,实验结果也就没法对齐。

WeightedRandomSampler的机制更有意思。它接收weights和一个num_samples,内部逻辑是每次从0到n-1中按权重概率抽取一个样本,抽取方式默认是"有放回"的,也就是说同一条样本在一个epoch里可能被抽中多次,也可能一次都抽不中。如果你希望"每个样本最多被抽一次",就需要把replacement设为False,但这么做的前提是num_samples不能超过数据集长度,否则没那么多不重复的样本可抽。

BatchSampler和前面几个不是同一层的东西。前面那些产生的索引是一维的,DataLoader默认把它们按batch_size切成一个个小组;BatchSampler则是把这个"切组"的过程提前到Sampler层。它接收一个内部的base_sampler(比如RandomSampler)和一个batch_size,每次迭代返回一个索引列表,这个列表就是最终一个batch对应的所有索引。它的价值在于你可以完全自定义"哪些索引进同一个batch"——比如你想做自定义的batch内负样本采样,或者让一个batch内只包含同类别样本,就必须通过BatchSampler控制。

2.3 DataLoader拿到索引之后的处理顺序

理清了Sampler的职责,再看DataLoader的完整数据处理流水线:Sampler产出一个batch的索引列表,DataLoader拿着这些索引逐个去Dataset里调用__getitem__取出原始样本,再交给collate_fn把多个样本整理成一个batch张量。如果有num_workers>0,这个取样本的过程会在多个子进程中并行进行。

理解了这条链路,你就明白为什么不同的采样策略会直接影响训练效果。如果Sampler产出的索引有偏,那么模型每个epoch看到的样本分布就有偏,batch内梯度更新的方向和方差也会跟着变。很多在数据增强上做了半天功夫、模型还是收敛不稳的情况,根源其实是采样的随机性不够或者赋予了某些类别过高的采样概率。

3. 自定义Sampler的完整实战:从需求到落地的全过程

3.1 一个真实需求:类别不均衡的均衡采样

现在假设我们有一个分类任务,三个类别的样本数量分别是1000、100、10,类别差距非常大。直接用RandomSampler训练,模型会严重倾向于预测多数类。常见的解决方案是做类别均衡采样:让每个类别在一个epoch内被抽到的机会大致相等。

最简单的做法是用WeightedRandomSampler。先计算出每个样本的权重,让少数类的权重高、多数类的权重低:

import torch from torch.utils.data import DataLoader, WeightedRandomSampler # labels 是数据集中所有样本的类别标签,形状为 (N,) labels = dataset.targets # 或者从dataset里取出来 class_counts = torch.bincount(torch.tensor(labels)) class_weights = 1.0 / class_counts.float() sample_weights = class_weights[torch.tensor(labels)] sampler = WeightedRandomSampler( weights=sample_weights, num_samples=len(sample_weights), replacement=True, generator=torch.Generator().manual_seed(42) ) dataloader = DataLoader(dataset, batch_size=32, sampler=sampler)

这里num_samples设成len(sample_weights),意思就是每个epoch的总采样次数和数据集的样本总数相同,但因为是有放回抽取,实际上样本的利用率是重复采样大于1、低频样本多次出现。

3.2 继承Sampler实现"每类固定数量"的精确采样

WeightedRandomSampler的缺点是权重比例全靠试,你没法精确控制"每个batch里少数类至少占几个"。如果你的业务场景对batch内的类别构成有硬性要求,就得自定义Sampler。

下面这段代码就是我实际项目中用过的方案,核心思路是:在每个batch内,多类样本随机抽k1个,少数类样本随机抽k2个,保证batch内类别比例固定:

import torch from torch.utils.data import Sampler class FixedClassSampler(Sampler): def __init__(self, labels, samples_per_class, batch_size, drop_last=True): self.labels = labels self.samples_per_class = samples_per_class self.batch_size = batch_size self.drop_last = drop_last self.class_to_indices = {} for idx, label in enumerate(labels): self.class_to_indices.setdefault(label, []).append(idx) def __iter__(self): num_classes = len(self.class_to_indices) per_class_batch = {c: self.samples_per_class[c] for c in self.class_to_indices} for cls, indices in self.class_to_indices.items(): random.shuffle(indices) pos = {c: 0 for c in self.class_to_indices} batches = [] # 每个batch:从每类中取固定数量 while True: batch_indices = [] for c in self.class_to_indices: start = pos[c] end = start + per_class_batch[c] if end > len(self.class_to_indices[c]): break batch_indices.extend(self.class_to_indices[c][start:end]) pos[c] = end else: if len(batch_indices) == self.batch_size: batches.append(batch_indices) continue break # 如果batch太大,可以再shuffle一次 for batch in batches: random.shuffle(batch) return iter(batches) def __len__(self): # 大致估算batch数量 total = 0 for c, indices in self.class_to_indices.items(): total += len(indices) // self.samples_per_class[c] return total // self.batch_size if self.drop_last else total

这个实现比WeightedRandomSampler更可控,你清楚知道每个batch里的类别分布。实际使用中还要注意:如果某类样本太少,可能撑不到生成足够的batch,最好提前检查一下每个类别的样本数量是否满足要求。

3.3 自定义Sampler与shuffle、drop_last的边界

很多人问过一个问题:自定义Sampler之后,DataLoader的shuffle参数还能用吗?

答案是:Sampler和shuffle是互斥的。只要传入了sampler,DataLoader会强制忽略shuffle参数,然后它内部再也不会创建RandomSampler。源码里的判断是if sampler is not None: self.sampler = sampler。同理,如果传入sampler还同时设置batch_sampler,两者也只能二选一。

这里有个实践经验:如果自定义Sampler返回的索引数不等于数据集的样本数,那么drop_last的行为也会受到影响。drop_last的作用是在batch切分时丢弃最后一个不足batch_size的batch。当你用BatchSampler时,这个逻辑要自己管理,DataLoader不会再处理。

4. 分布式训练中的Sampler:DistributedSampler的特殊之处

4.1 单机多卡为什么要单独处理数据切分

当你从单卡切换到多卡训练时,Sampler的重要性会急剧上升。多卡训练的本质是数据并行:每张卡只负责数据的一个子集,每张卡单独算梯度,然后做梯度同步。那"数据怎么切分"就成了关键:如果两张卡拿到的数据完全相同,梯度算了两遍,等于白算;如果切分不均衡,有的卡数据多,有的卡数据少,整体训练时间会被最慢的那张卡卡住。

DistributedSampler就是来解决这个切分问题的。它会把所有样本按卡数(world_size)均匀切分,保证每张卡拿到的样本子集互不重叠。切分逻辑默认是近似均匀的:把索引按顺序划分为多个块,每张卡拿一块。

4.2 shuffle、seed对齐和epoch的关系

DistributedSampler最容易被忽略的一点是它的shuffle逻辑和随机种子。它不像RandomSampler那样由你传入generator,而是自己根据epoch和rank来生成确定性随机序列。具体规律是:它内部用self.epoch + self.seed作为随机数种子来生成打乱的顺序,所以每个epoch开始前必须调用set_epoch方法,否则每个epoch的采样结果完全一样。

来看标准用法:

sampler = torch.utils.data.distributed.DistributedSampler( dataset, num_replicas=world_size, rank=rank, shuffle=True, seed=your_seed ) for epoch in range(num_epochs): sampler.set_epoch(epoch) for batch in dataloader: train_step(batch)

这里的seed如果不设置,在不同机器之间默认随机不一致,会导致不同机器拿到的数据子集不稳定。我见过有人用DistributedSampler训练,模型在单卡上测试正常,多卡上怎么都不收敛,排查很久才发现是seed不一致导致每卡的数据分布对不上。

4.3 一个epoch内"数据重复"还是"数据缺失"的排查

分布式训练里见过最多的问题有两个:一个是数据重复,一个是数据缺失。

数据重复的典型原因:把Dataset做了多次DataLoader实例化,每个DataLoader都用默认的RandomSampler,但每个进程没有设置独立的随机种子,导致多卡之间拿到的索引相同。数据缺失的典型原因是忘了设置shuffle=True,DistributedSampler会按顺序切分,如果数据集的类别标签是顺序排列的(前1万个都是类别0,后1万个都是类别1),那么rank0那张卡拿到的全是类别0,rank1那张卡拿到的全是类别1,模型根本不收敛。

这个问题的根因不难理解:DistributedSampler做的是"按位置切分"而不是"均匀采样"。你只有保证在切分之前数据顺序已经被打乱,每张卡拿到的子集在类别分布上才会接近全局分布。所以养成一个习惯:无论单卡还是多卡,数据集的样本顺序最好先全局随机打乱一次,再交给DataLoader做进一步随机,别把"数据恰好有序"当成一种可靠保障。

5. 我在实际项目中踩过的Sampler的坑

5.1 坑一:WeightedRandomSampler的replacement参数选错

做文本分类的时候,数据集类别不均衡,我用WeightedRandomSampler做均衡采样,起初replacement设成了False,结果一个epoch只能采几百个batch,而且少数类样本频繁没有出现在batch里。仔细读文档才发现:当replacement=False时,这个Sampler的行为是"不放回地按权重抽样本",一旦某个样本被抽中,它就从候选池里移除,权重大的样本会先被抽走,等到后期剩下的全是权重小的样本,抽样结果反而变得更加不均衡。

正确的做法取决于你想要的行为:

  • 想要"每个样本平均被看到的次数一致":replacement=True,num_samples通常设成总样本数。
  • 想要"每个样本最多被抽一次":replacement=False,此时权重只影响抽样的先后顺序。

这个坑的深层原因是很多人把"权重"理解成"这个样本一定会被抽到几次",实际不是的。权重定义的是相对概率,同一个样本在大量采样里会依据概率被抽中,但抽中次数是一个随机变量。如果你想让少类样本在一个epoch里100%被看到至少一次,最简单的办法还是像我后面讲的那样自定义Sampler,或者用多个Sampler组合。

5.2 坑二:传入sampler时忘了shuffle失效这回事

另一个高频bug:自定义Sampler + DataLoader(shuffle=True),从代码上看"既有采样器又想随机打乱",但实际上shuffle被静默忽略了。有一天我发现自己的验证集指标比训练集还低,排查后发现训练阶段的loader因为原始数据顺序刚好容易学习,而验证阶段换了loader导致评估时数据分布差异大。纠正之后才发现一个epoch内数据出现顺序根本没变——因为我传入了自定义sampler但没关掉DataLoader的shuffle,结果shuffle根本没生效,训练数据顺序完全由Sampler决定。

解决思路很简单:如果你写自定义Sampler,就默认把DataLoader的shuffle参数保持False,避免后人误以为shuffle还在生效。反过来,如果只是简单打乱数据,就用shuffle=True别去动Sampler,不要叠床架屋。

5.3 坑三:多进程worker下的采样状态复制问题

这是一个比较隐蔽的问题。DataLoader的num_workers>0时,主进程里的Sampler负责生成索引,然后这些索引会被分发到多个worker进程里,由worker进程调用Dataset的__getitem__去取数。这里的关键在于:Sampler是在主进程运行的,它的随机状态不会被复制到每个worker进程。

也就是说,如果你在自定义Sampler里创建了一个随机数生成器,并且用全局随机状态,每个epoch的采样结果理论上每个worker看到的是同一个Sampler状态,不会因为worker间随机性不同而重复采样。但如果你的Sampler惰性地缓存中间结果,比如使用了全局的python random模块,就可能因为主进程随机状态在每个epoch之前没有被重置,导致多个worker看到同一个索引序列。

我的建议:自定义Sampler里使用torch.Generator,并且通过manual_seed固定;每个epoch开始前如果需要重置,就在训练循环里显式调用sampler.reset()(如果你实现了这个方法),不要让Sampler内部依赖全局随机函数。

5.4 坑四:在BatchSampler的__len__里算错epoch步数

最后再说一个关于"epoch步数"的坑。训练循环里经常需要知道一个epoch有多少个batch,用来算warmup步数或者log的频率。很多人直接写len(dataloader),但如果你自定义了BatchSampler,len(dataloader)返回的是BatchSampler的__len__。如果__len__写得不精确,就会出现"打印的epoch步数和实际跑到的步数不一致"。

这个问题在IterableDataset场景下尤其明显:流式数据不知道总量,__len__常常返回一个估算值。我建议所有自定义采样器在实现__len__时,用代码原样跑一遍索引生成逻辑来统计batch数,而不是单纯靠除法和取整估算,避免边界条件漏算。

6. 采样器的进阶玩法与选型思路

6.1 不均衡样本下的采样策略选择

结合前面的内容,当训练数据类别不均衡时,你面前其实有三个层级的选择:

  • 最低成本:直接用WeightedRandomSampler,把类别频率的倒数作为权重。适合快速实验、基线模型,代码改动最小。
  • 中等控制:自定义Sampler精确控制每个batch内的类别构成。适合固定batch结构、对比实验、需要保证每次迭代都能看到少类的场景。
  • 稳定提升:把"采样均衡"和"损失函数加权"结合使用。比如在采样上做的均衡度低一点,保留一点真实的类别分布,然后通过loss_weight把少类样本的梯度放大更多倍,这种组合往往比单用一种方法的效果更稳。

从实践角度来说,我倾向于在训练的前期用均衡采样让模型快速学到每个类别的特征,后期再逐步放松采样权重、让模型在接近真实分布的样本上微调。具体怎么退火,可以用一个简单的线性衰减实现:训练初期weights按频率的反比,训练结束时weights全部为1。把这个退火权重和WeightedRandomSampler配合,效果会比固定的均衡采样好很多。

6.2 难例挖掘与采样器的配合

如果你的任务涉及难例挖掘,可以考虑在Sampler层做文章:难例不是"随机抽样"抽出来的,而是根据模型上一个epoch的loss来决定的。一个经典的做法是每完成一个epoch,记录每个样本的loss值,然后下一个epoch里对这些loss值做加权采样,loss大的样本被抽中的概率高。

这个方案用自定义Sampler实现起来并不复杂:你在训练循环里维护一个sample_loss数组,epoch结束后把它传给Sampler的set_loss_weights方法,然后Sampler根据权重重新构建采样分布。关键点是这个权重更新是异步的,用上一个epoch的loss影响下一个epoch的分布,在训练过程中会导致数据分布略微滞后于模型状态,但整体收敛效果通常比完全均匀要好。

6.3 采样器选型的场景对照表

最后把选型逻辑整理成一个对照表,方便你按需取用:

使用场景推荐方案原因
常规分类/回归训练RandomSampler(shuffle=True)随机性足够,成本最低
验证集/测试集评估SequentialSampler保证推理结果可复现
类别不均衡(简单处理)WeightedRandomSampler一行代码解决问题
类别不均衡(精确控制)自定义Sampler控制batch内类别比例
手动划分训练/验证子集SubsetRandomSampler直接传indices最省事
自定义batch结构BatchSampler + 自定义base_sampler掌控batch内的索引组合
单机多卡/多机多卡DistributedSampler自动切分+shuffle+seed对齐
每个epoch步数固定自定义num_samples控制训练时长与数据循环轮数

根据我个人的项目经验,大部分Sampler相关的问题都不是"不会用",而是"不知道某个参数被静默忽略"或者"不清楚Sampler返回索引的数量与Dataset长度的关系"。写自定义Sampler时,最值得你多花时间的不是实现,而是把__len__算准确、把与shuffle的互斥关系处理好、把分布式场景的seed对齐解决掉。做到这三点,你的数据加载层就会非常稳定,后续调模型的时候也能少排查一大堆莫名其妙的玄学问题。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询