PyTorch Sampler完全指南:从数据加载到分布式训练的采样器详解
2026/9/24 23:56:35 网站建设 项目流程

先聊点实际场景。很多人刚开始用PyTorch时,都是Dataset接DataLoader,shuffle=True一开,数据顺序乱了,模型训起来了,基本就没人管线下面有个叫Sampler的东西。但等你真正去调分布式训练、处理极度不均衡的数据集、或者想控制每个batch该喂哪些样本时,Sampler就变成了绕不开的关键角色。

这篇文章就专门把PyTorch里的Sampler讲透,讲清楚它到底在数据流水线的哪个环节干活,内置的几种Sampler分别适合什么场景,以及怎么自己写一个符合业务需求的采样器。不管你是刚入门的小白,还是已经在跑模型但总被数据加载问题折磨的老手,这篇文章都值得看一眼。

1. 先搞清楚Sampler在PyTorch数据流水线里的位置

想要理解Sampler,先得把PyTorch载入数据的完整流程拆开看一遍。很多文档会告诉你Dataset负责取样本、DataLoader负责批量加载,但其实中间还藏着一层Sampler,它才是决定“每个batch拿哪些样本”的真正指挥官。

1.1 Dataset、DataLoader、Sampler的三方分工

一句话总结三者的关系:Dataset管的是“一个样本长什么样”,Sampler管的是“按什么顺序拿样本”,DataLoader管的是“怎么协调上面两层,最终拼出模型能吃的batch”。

展开来说,Dataset定义了两个核心能力:一是数据集里有多少样本,也就是__len__;二是给定一个索引值,能返回对应的数据样本和标签,也就是__getitem__。当你写出dataset[3]的时候,拿到的是第4个样本,这个动作和Sampler没有任何关系。

Sampler则负责生成“索引序列”。它做的事情本质上就是:我告诉你下一个该访问哪个下标。比如数据集有100个样本,SequentialSampler返回的就是0, 1, 2, ..., 99这串顺序索引;RandomSampler返回的则是0到99的一个随机排列。它只负责产生索引,不负责取数据。

DataLoader在上层做组装工作:它拿到Sampler给的索引,传给Dataset.__getitem__,把取出来的样本收集到一起,再用collate_fn拼成一个batch张量。所以你可以理解成,Dataset是仓库管理员,Sampler是拣货清单,DataLoader是打包员。

这里有个很常见的误解:以为shuffle=True是DataLoader在打乱数据。实际上,DataLoader看到shuffle=True后,会默默创建一个RandomSampler实例,真正打乱索引次序的是这个Sampler。也就是说,你平时每跑一步训练就在人家里干活,只是自己没察觉。

1.2 为什么单独拆出Sampler这一层

可能有朋友会问:直接把“打乱”写进DataLoader不行吗?为什么非要单独拆一个Sampler接口?这个设计其实是经过深思熟虑的,因为“如何决定采样顺序”本身就是一种值得灵活定制的策略。

最直接的理由是解耦。数据读取分成了“有什么数据”和“按什么顺序读”,这两件事经常需要独立变化。同样是100张图片,训练时想随机打乱,验证时想顺序读取,如果打乱逻辑耦合在DataLoader里,你就得把DataLoader也换掉。拆出Sampler后,换个Sampler实例就能切换策略。

第二个理由是满足复杂采样需求。常用的随机打乱只是最简单的策略,真实业务里还有大量场景需要精细控制采样方式。比如类别不平衡时,要给少数类样本更高的被抽中概率,就需要WeightedRandomSampler;要手动指定只用数据集里某一部分索引(比如训练集和验证集划分),就用SubsetRandomSampler;要做多轮采样、有放回采样,甚至自定义一套难样本挖掘逻辑,这些统统可以通过实现自定义Sampler来完成。

第三个理由是分布式训练。用DistributedDataParallel跑多机多卡时,每个进程要拿到不同的数据分片,否则梯度同步就乱了。这需要DistributedSampler按照ranknum_replicas对数据集做分片,DataLoader自身没有能力做到这一点。

还有一个容易忽略但很重要的点:BatchSampler让“凑batch”这个动作也变成了可配置的逻辑。你可以把一批索引合成一组,控制最后一组不足一个batch时要不要丢弃,甚至动态改变每个batch的大小。这套机制如果揉在DataLoader内部,就没有这么灵活了。

2. 内置Sampler逐个过一遍,搞懂各自适合什么场景

PyTorch的torch.utils.data里自带了好几种Sampler,从最基础的顺序采样到处理不均衡数据的加权采样都有。我建议你把这些内置采样器当成“答案库”,大部分场景不需要自己造轮子,先选一个合适的内置实现能省很多事。

2.1 SequentialSampler:最朴素的顺序读取

SequentialSampler的逻辑简单到几乎不用解释:对长度为N的数据集,它会依次返回0, 1, 2, ..., N-1。这是设置DataLoader(shuffle=False)时默认使用的采样器。

它有两个典型使用场景:一个是验证集、测试集评估阶段,模型每次看到的样本顺序应该是固定的,这样多次评估结果可直接对比;另一个是某些对时间顺序敏感的任务,比如时间序列预测、视频帧序列,打乱样本会破坏时序关系,必须顺序读取。

需要提醒的是,SequentialSampler一般不需要你手动创建,DataLoader默认行为已经帮你处理好了。你只有在需要拿到索引序列本身做后续操作时,才可能手动实例化它,比如:

from torch.utils.data import SequentialSampler sampler = SequentialSampler(range(10)) print(list(sampler)) # [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]

2.2 RandomSampler:默认随机打乱的实现者

RandomSamplershuffle=True时DataLoader使用的默认采样器。它会在每次迭代时生成一个新的随机排列,从而保证每个epoch里数据顺序都不一样。

它有两个关键参数,replacementnum_samples。当replacement=False时,它会对0到N-1做一次随机排列,每个样本正好出现一次;当replacement=True时,它允许同一个索引被重复抽到,相当于有放回抽样,此时可以通过num_samples控制总共抽多少次。

手动使用方式:

from torch.utils.data import RandomSampler sampler = RandomSampler( range(10), replacement=False, num_samples=10 )

replacement=True的一个实际应用是做简单的过采样。比如少数类样本只有50条,你想让它在每个epoch里多出现几次,就把num_samples设大一些,比如100或200。这样模型在同一个epoch里能多次看到这些少数类样本,比单纯复制数据更自然。

2.3 SubsetRandomSampler:手动切分数据集的利器

如果你已经有一整个数据集,想按索引划分训练集和验证集,可以用SubsetRandomSampler。它接收一个indices列表,每次迭代时从这个列表里随机抽样,实现“只在这些索引范围内打乱”。

因为数据集不需要真正切分成两个对象,只是采样范围不同,所以这种方法非常轻量,也避免了用torch.utils.data.Subset时可能出现的索引混乱问题。如果两个采样器用的是不相交的索引,训练集和验证集之间就不会有数据泄漏。

基本用法:

from torch.utils.data import SubsetRandomSampler n_samples = 1000 indices = list(range(n_samples)) train_idx = indices[:800] val_idx = indices[800:] train_sampler = SubsetRandomSampler(train_idx) val_sampler = SubsetRandomSampler(val_idx)

需要注意,indices列表的顺序会影响初始状态,但每次迭代时它都会在内部重新打乱,所以最终返回的批次顺序是随机的。另外,这种划分方式通常用于单机场景,分布式训练时还需要配合DistributedSampler做进一步分片。

2.4 WeightedRandomSampler:对付样本不均衡的杀手锏

这个采样器在处理类别严重失衡的任务时非常好用。它接收三个参数:weightsnum_samplesreplacement。其中weights是一个长度为N的列表,每个元素代表该位置样本的权重;num_samples表示每次迭代采多少样本;replacement表示是否有放回。

它的原理也简单:每个样本被抽中的概率正比于它的权重。比如样本A权重是3,样本B权重是1,那么A被抽中的概率就是B的3倍。这比手动复制样本再shuffle要高效得多,也不会增加数据集大小。

我拿一个二分类正负样本严重失衡的场景来说。假设负样本有5000条,正样本只有100条,直接在原始数据上训练,模型很容易被负样本带偏。这时可以这样做:

from torch.utils.data import WeightedRandomSampler # 假设dataset.samples是一个[(img, label), ...]列表 class_counts = [5000, 100] sample_weights = [1.0 / class_counts[label] for _, label in dataset.samples] sampler = WeightedRandomSampler( sample_weights, num_samples=len(dataset.samples), replacement=True )

这里给少数类样本分配了更大的权重。严格来说,weights并不需要归一化,因为采样的核心是权重比例,归一化与否不影响相对概率。

有一个特别容易踩的坑:weights的长度必须和dataset长度完全一致,并且权重要对应到具体样本,而不是类别。很多人会误以为传入每个类别的权重就行,结果运行时报错或者样本分布完全不对。如果数据集很大,一次性建出长度的权重列表可能内存开销不小,但这是方案本身的特点,数据加载阶段多占点内存通常是可以接受的。

2.5 BatchSampler:把单个索引拼成一批数据

前面提到的几个Sampler都在输出“单个索引”,而BatchSampler是负责把单个索引组合成“一批索引”的采样器。它在内部包装另一个Sampler,然后按照batch_size把索引切分成List[List[int]]。

下面这段代码能直观看到它的效果:

from torch.utils.data import BatchSampler, SequentialSampler sampler = SequentialSampler(range(10)) batch_sampler = BatchSampler(sampler, batch_size=3, drop_last=False) for batch_indices in batch_sampler: print(batch_indices) # [0, 1, 2] # [3, 4, 5] # [6, 7, 8] # [9]

drop_last参数控制的是:当最后一组不够一个batch时,是保留还是丢弃。取值为True时,[9]这一组会被丢掉。训练阶段一般会设置为True,因为很多损失函数和BatchNorm对最后一个小batch可能有负面影响;验证阶段则经常设置为False,保证所有样本都被评估到。

需要注意,如果你在DataLoader里显式传入了batch_sampler,那么batch_sizeshufflesamplerdrop_last这几个参数都不能再传,它们彼此冲突。这个约束在源码里是直接断言检查的,报错信息也很明确。

2.6 DistributedSampler:分布式训练里不可缺席的成员

DistributedDataParallel做多卡训练时,DistributedSampler基本是标配。它做的事情有两个:第一,根据num_replicas(总共几个进程)和rank(当前进程编号)把数据集切分成互不重叠的分片;第二,在每个分片内部做shuffle。

如果没有它,多个进程会在每个epoch里重复取到同样的数据,梯度同步时就等于每个卡都看到一样的样本,分布式训练的效果会大打折扣。

这里有个特别重要的细节:每轮epoch开始前,必须调用set_epoch(epoch)方法,否则shuffle的随机顺序在每个epoch都是相同的。因为DistributedSampler内部通过一个epoch属性来生成不同的随机种子,不更新的话随机性就退化了。很多人在分布式训练里调了半天不收敛,最后发现问题就出在这。

基本用法:

from torch.utils.data import DistributedSampler sampler = DistributedSampler( dataset, num_replicas=world_size, rank=rank, shuffle=True ) for epoch in range(max_epochs): sampler.set_epoch(epoch) for batch in dataloader: train_step(batch)

如果是单机多卡,world_size就是显卡数,rank就是当前卡的编号。DDP框架初始化后,这些值可以从dist.get_world_size()dist.get_rank()获取。

3. 手写一个自定义Sampler:先搞懂那套约定

内置Sampler虽然覆盖了大部分常见需求,但总有一些业务场景,比如主动学习、课程学习、难样本挖掘,需要你按自己的想法定制采样顺序。一旦理解自定义Sampler的约定,你就能像搭积木一样控制数据流。

3.1 Sampler基类和必须实现的两个方法

PyTorch里的Sampler是一个抽象基类,它只规定了两件事:__iter____len__。前者必须返回一个可迭代对象,每次迭代出的值就是要传给Dataset的下标;后者返回整个采样过程会产生多少个索引(注意不是数据集长度)。

值得注意的是,官方对该基类的实现约等于一个空壳:

class Sampler(Generic[T_co]): def __init__(self, data_source): self.data_source = data_source def __iter__(self) -> Iterator[T_co]: raise NotImplementedError def __len__(self) -> int: raise NotImplementedError

实际写自定义Sampler时,通常会继承Sampler,在__init__里保存数据源引用,在__iter__里实现核心逻辑,在__len__里返回采样总数。__len__不是可有可无的,DataLoader在计算epoch总步数、显示进度条时会用到它。如果你懒得严格继承,也可以直接定义一个包含这两个方法的普通类,PyTorch更多是靠鸭子类型来识别。

还要注意类型和范围:迭代器返回的索引必须是Python的int类型,并且取值范围在0len(dataset)-1之间。如果返回numpy.int64或越界索引,轻则类型报错,重则取到不存在的样本或出现难以排查的乱序问题。

3.2 实战:实现一个“难样本优先”采样器

为了演示,我写一个基于“样本loss”的采样器。思路是先让模型跑一个epoch,记录每个样本的loss,然后把loss值转换成采样概率,loss越高,样本越容易被选中。这是难样本挖掘(Hard Example Mining)的一种简化版,模型能把更多注意力放在它当前最容易出错的样本上。

简化版代码如下:

import torch from torch.utils.data import Sampler class HardExampleSampler(Sampler): def __init__(self, data_source, num_samples=None): super().__init__(data_source) self.data_source = data_source self.num_samples = num_samples if num_samples is not None else len(data_source) def __iter__(self): losses = getattr(self.data_source, "losses", None) if losses is None: indices = torch.randperm(len(self.data_source)).tolist() else: losses = torch.as_tensor(losses, dtype=torch.float32) # 防止负值或极端值影响,先用温度系数缩放 logits = losses / max(losses.max().item(), 1e-8) probs = torch.softmax(logits, dim=0) sampled = torch.multinomial(probs, self.num_samples, replacement=True) indices = sampled.tolist() return iter(indices) def __len__(self): return self.num_samples

这段代码的关键在于,把采样逻辑和Dataset实现解耦。Dataset只需要维护一个losses属性,训练过程中不断更新它就行;采样器每次迭代都读取这个属性,重新计算采样概率。这样就能实现“越难样本,越常出现”的效果。

使用起来也很直接:

train_loader = DataLoader( dataset, batch_size=64, sampler=hard_sampler, num_workers=4 )

需要说明的是,真实项目里的难样本挖掘通常比这个复杂得多,往往还要配合课程学习、Focal Loss等策略一起用,但上面这段代码已经足够帮你理解自定义Sampler的套路。

3.3 自定义Sampler实操时的几个注意事项

写自定义Sampler的时候,最大的坑往往不在逻辑本身,而在和DataLoader其他参数的配合上。首先,一旦显式传了sampler,DataLoader的shuffle参数就必须是False,因为shuffle本身就会创建一个RandomSampler,两个采样器一起传会直接报冲突错误。

其次,如果你需要drop_last功能,不能用DataLoader提供的drop_last参数,因为那个参数和sampler参数互斥。正确做法是把自定义Sampler先用BatchSampler包一层,再把batch_sampler传给DataLoader,同时把batch_size设成None

另外还有一个在多进程数据加载场景特别容易踩的坑:Sampler必须有良好的序列化支持,因为DataLoader在num_workers > 0时会用多个worker进程,采样器对象需要能被pickle。换句话说,不要在__iter__里写lambda函数,也不要在__init__里保存无法序列化的对象(比如文件句柄、线程锁等)。真遇到过把torch.Generator放进去使用却忘了它能正常pickle的坑,建议用前先在本地打印一下torch.__version__相关行为。

4. 常见问题与排查技巧实录

这一节整理一下我在实际使用中碰到过的、社群提问里也高频出现的Sampler相关问题。这些问题单看文档很难一次讲明白,但搞清原因后基本都能快速定位。

4.1 参数冲突类报错,先看shuffle

最常见的错误是:

RuntimeError: sampler option is mutually exclusive with shuffle 或者 AssertionError: batch_size should be None when batch_sampler is provided

这类报错其实是保护机制在起作用,防止你给了多种采样策略让DataLoader无所适从。解决思路很简单:先确定你最终想用哪种方式控制索引。如果用一个自定义Sampler,就把shuffle置为False;如果用了batch_sampler,就把batch_size置为None。别想着同时传,让参数自洽比“绕过规则”安全得多。

4.2 WeightedRandomSampler的权重“对样本不对类别”

WeightedRandomSampler时,我见过最多的错误是:

sampler = WeightedRandomSampler( weights=[0.1, 0.9], # 只有两个数,但dataset有1000个样本 num_samples=1000, replacement=True )

运行后要么直接报错,说weights长度不匹配,要么即便不报错,采样结果也完全不符合预期。记住一句话:weights是一份“样本级别”的数组,它第几个元素就代表dataset里第几个样本的权重。正确写法是先按类别算好class weight,再把它映射到每个样本上。类别权重可以用sklearn.utils.class_weight.compute_class_weight这类现成工具计算,也可以手工用样本数倒数。

4.3 分布式训练时每个epoch的shuffle顺序完全一样

这是DistributedSampler最隐蔽的坑。它的实现逻辑是:用self.epoch作为随机种子的一部分生成打乱顺序。如果你在训练循环里没有在每个epoch开头调用set_epoch(epoch),那么每个epoch生成的随机排列其实一模一样。尤其在模型不太收敛的时候,很多人会怀疑是学习率或初始化的问题,查了一圈才发现是数据顺序根本没变。

我个人的习惯是,在分布式训练代码里,把sampler.set_epoch(epoch)放在最显眼的位置:

for epoch in range(epochs): train_sampler.set_epoch(epoch) for batch in train_loader: ...

这个动作虽然只有一行,但作用相当于给每个epoch重新“洗牌”,对收敛稳定性和最终精度都有实际影响。

4.4 自定义Sampler返回了非int类型的索引

自定义Sampler的第二个高频坑,是索引类型不合法。比如你用了numpy生成索引:

indices = np.random.permutation(len(self.data_source))

此时indicesnumpy.ndarray,元素类型是np.int64。虽然很多场景下PyTorch能隐式转换,但如果你交给Dataset.__getitem__后出现了奇怪的切片行为,或者和某个版本的DataLoader不兼容,还是会踩坑。稳妥的做法是显式转成Python int列表:

indices = np.random.permutation(len(self.data_source)).tolist()

这样返回的就是标准int列表,兼容性最好。

4.5 想调试Sampler,直接把它当可迭代对象用

如果你不确定某个Sampler到底生成了什么索引序列,最直接的办法就是把它包进list()里打印出来。Sampler本身就是可迭代对象,list(sampler)会触发一次完整的迭代,把生成的索引全部展示出来。这个小技巧在调试BatchSampler时特别有用,你能直观看到每个batch包含哪些索引、最后一个batch是不是不完整。

from torch.utils.data import BatchSampler, RandomSampler sampler = RandomSampler(range(10), replacement=True, num_samples=15) batch_sampler = BatchSampler(sampler, batch_size=4, drop_last=False) print(list(batch_sampler)) # 例如 [[0, 7, 2, 2], [9, 5, 3, 1], [7, 4, 8, 0], [6]]

打印出来之后,你既能直接验证replacement是否生效,也能检查drop_last的行为是否符合预期,比对着源码空想要快得多。

写在最后的经验之谈

Sampler这套机制,平时看着不起眼,但真正影响训练效果的地方还挺多。我在实际项目中最大的体会是,任何“数据侧调优”需求,最好先想一想能不能通过Sampler解决,而不是上来就改写Dataset或DataLoader。比如过采样少数类,用WeightedRandomSampler就比重建一个Dataset要轻量得多;分布式训练,不引入DistributedSampler甚至会有梯度同步问题。

另外还有一个细节想分享:如果你的数据集非常大,Sampler生成的索引序列只是一个轻量级CPU对象,它的开销相比真正从磁盘读图片、做数据增强来说几乎可以忽略。但如果你在自定义Sampler里做了太复杂的计算,比如每次迭代都要全量跑一次模型分数,那就要小心它成为训练瓶颈。这种情况下,把采样逻辑分阶段缓存起来,或者降低采样频率,效果会更好。

Sampler这个设计不算复杂,但它把“如何选数据”这件事从“如何读数据”里优雅地剥了出来。理解它是深入PyTorch数据加载体系的关键一步,也值得你在遇到数据相关问题时,第一时间想到去“动一下”训练管线里的这个灵活的旋钮。

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

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

立即咨询