
1. 为什么每个PyTorch新手都会在Sampler上栽跟头1.1 一次数据顺序错乱事故的排查全过程前阵子帮一个朋友调试训练脚本现象非常诡异同一个模型、同一份数据在A机器上跑得好好的换到B机器上loss曲线就开始抖动验证集指标也忽高忽低。代码逐行比对了两遍模型结构、优化器、学习率调度、数据增强逻辑全都一样最后才发现问题出在DataLoader的sampler参数上——他在B机器上给DataLoader传了一个自定义的sampler却没有关掉默认的shuffle逻辑两个采样逻辑叠加在了一起产生了预料之外的索引顺序。这种问题在PyTorch的Sampler使用中太常见了。我见过不少同学在DataLoader里看到参数只知道shuffleTrue能打乱数据却完全没有意识到背后真正干活的是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去取数。shuffleTrue这个参数的本质只是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随机打乱的索引常规训练shuffleTrueSubsetRandomSampler从指定子集中随机取手动划分训练/验证集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_workers0这个取样本的过程会在多个子进程中并行进行。理解了这条链路你就明白为什么不同的采样策略会直接影响训练效果。如果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( weightssample_weights, num_sampleslen(sample_weights), replacementTrue, generatortorch.Generator().manual_seed(42) ) dataloader DataLoader(dataset, batch_size32, samplersampler)这里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_lastTrue): 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是互斥的。只要传入了samplerDataLoader会强制忽略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. 分布式训练中的SamplerDistributedSampler的特殊之处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_replicasworld_size, rankrank, shuffleTrue, seedyour_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但每个进程没有设置独立的随机种子导致多卡之间拿到的索引相同。数据缺失的典型原因是忘了设置shuffleTrueDistributedSampler会按顺序切分如果数据集的类别标签是顺序排列的前1万个都是类别0后1万个都是类别1那么rank0那张卡拿到的全是类别0rank1那张卡拿到的全是类别1模型根本不收敛。这个问题的根因不难理解DistributedSampler做的是按位置切分而不是均匀采样。你只有保证在切分之前数据顺序已经被打乱每张卡拿到的子集在类别分布上才会接近全局分布。所以养成一个习惯无论单卡还是多卡数据集的样本顺序最好先全局随机打乱一次再交给DataLoader做进一步随机别把数据恰好有序当成一种可靠保障。5. 我在实际项目中踩过的Sampler的坑5.1 坑一WeightedRandomSampler的replacement参数选错做文本分类的时候数据集类别不均衡我用WeightedRandomSampler做均衡采样起初replacement设成了False结果一个epoch只能采几百个batch而且少数类样本频繁没有出现在batch里。仔细读文档才发现当replacementFalse时这个Sampler的行为是不放回地按权重抽样本一旦某个样本被抽中它就从候选池里移除权重大的样本会先被抽走等到后期剩下的全是权重小的样本抽样结果反而变得更加不均衡。正确的做法取决于你想要的行为想要每个样本平均被看到的次数一致replacementTruenum_samples通常设成总样本数。想要每个样本最多被抽一次replacementFalse此时权重只影响抽样的先后顺序。这个坑的深层原因是很多人把权重理解成这个样本一定会被抽到几次实际不是的。权重定义的是相对概率同一个样本在大量采样里会依据概率被抽中但抽中次数是一个随机变量。如果你想让少类样本在一个epoch里100%被看到至少一次最简单的办法还是像我后面讲的那样自定义Sampler或者用多个Sampler组合。5.2 坑二传入sampler时忘了shuffle失效这回事另一个高频bug自定义Sampler DataLoader(shuffleTrue)从代码上看既有采样器又想随机打乱但实际上shuffle被静默忽略了。有一天我发现自己的验证集指标比训练集还低排查后发现训练阶段的loader因为原始数据顺序刚好容易学习而验证阶段换了loader导致评估时数据分布差异大。纠正之后才发现一个epoch内数据出现顺序根本没变——因为我传入了自定义sampler但没关掉DataLoader的shuffle结果shuffle根本没生效训练数据顺序完全由Sampler决定。解决思路很简单如果你写自定义Sampler就默认把DataLoader的shuffle参数保持False避免后人误以为shuffle还在生效。反过来如果只是简单打乱数据就用shuffleTrue别去动Sampler不要叠床架屋。5.3 坑三多进程worker下的采样状态复制问题这是一个比较隐蔽的问题。DataLoader的num_workers0时主进程里的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)但如果你自定义了BatchSamplerlen(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 采样器选型的场景对照表最后把选型逻辑整理成一个对照表方便你按需取用使用场景推荐方案原因常规分类/回归训练RandomSamplershuffleTrue随机性足够成本最低验证集/测试集评估SequentialSampler保证推理结果可复现类别不均衡简单处理WeightedRandomSampler一行代码解决问题类别不均衡精确控制自定义Sampler控制batch内类别比例手动划分训练/验证子集SubsetRandomSampler直接传indices最省事自定义batch结构BatchSampler 自定义base_sampler掌控batch内的索引组合单机多卡/多机多卡DistributedSampler自动切分shuffleseed对齐每个epoch步数固定自定义num_samples控制训练时长与数据循环轮数根据我个人的项目经验大部分Sampler相关的问题都不是不会用而是不知道某个参数被静默忽略或者不清楚Sampler返回索引的数量与Dataset长度的关系。写自定义Sampler时最值得你多花时间的不是实现而是把__len__算准确、把与shuffle的互斥关系处理好、把分布式场景的seed对齐解决掉。做到这三点你的数据加载层就会非常稳定后续调模型的时候也能少排查一大堆莫名其妙的玄学问题。