ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch Sampler深度解析:从数据采样到分布式训练的关键技术

PyTorch Sampler深度解析:从数据采样到分布式训练的关键技术 1. 先把概念理清Sampler到底是干什么的先说个题外话标题里那个Pytorh就是PyTorch估计是手滑打错了不影响阅读。真正让你头疼的应该是后半截——Sampler。说实话我在最早用PyTorch训练模型的时候整整半年都没主动碰过Sampler因为DataLoader默认行为已经够用了。数据读进来、打乱顺序、切成batch、喂给模型看起来一切都很顺畅。直到我开始做不平衡样本分类、做多卡分布式训练、自定义训练策略才意识到Sampler才是数据管线的灵魂DataLoader只是台面上的调度员。Sampler直译叫采样器它的职责就是回答一个问题下一个应该取哪个样本别小看这个取哪个的问题它直接决定了模型能看到什么、以什么顺序看、看不看得到重复内容进而影响训练收敛速度、最终精度甚至多卡环境下会不会出现数据泄露或重复。这套机制不光PyTorch有TensorFlow的tf.data、MindSpore的Dataset也都做了类似设计但PyTorch的Sampler做得最独立、最灵活也最容易让人困惑。困惑主要来自两个方面一是Sampler和DataLoader的分工边界不清晰二是内置的好几种Sampler名字相近、参数相近但行为差别很大。1.1 从一次训练循环说起数据是怎么一步步走进模型的要理解Sampler先看一条完整的数据流。假设你写了一个最普通的训练脚本from torch.utils.data import DataLoader, TensorDataset dataset TensorDataset(x_data, y_data) dataloader DataLoader(dataset, batch_size32, shuffleTrue) for epoch in range(10): for inputs, labels in dataloader: outputs model(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step()这段代码里shuffleTrue其实是一个快捷开关。当你把shuffle设为True时PyTorch内部会在DataLoader里自动创建一个RandomSampler并用它来产生样本索引。当你把shuffle设为False时用的则是SequentialSampler。所以Sampler的产出物不是数据本身而是样本索引indices。DataLoader拿到这串索引后再去dataset里做dataset[i]操作把对应位置的样本取出来聚合成一个batch。这就是为什么网上有些人管它叫索引生成器这个说法很准确。如果你把DataLoader的生命周期拆开看它是这样一个链条Sampler生成索引序列比如[3, 0, 2, 1]DataLoader按照索引序列逐一取出样本dataset[3]、dataset[0]、dataset[2]、dataset[1]默认情况下DataLoader会对这批样本调用默认的collate_fn把单个样本堆叠成tensor一批数据被送到模型里也就是说Sampler决定的是顺序collate_fn决定的是怎么拼装dataset决定的是单条样本长什么样。三者各管一段最好不要越权。1.2 Sampler的三个核心职责顺序、重复、分配把Sampler做的事情往深了挖其实就是三件小事但每件小事都能玩出花。第一是控制顺序。最简单的SequentialSampler按0、1、2、3……顺序取RandomSampler打乱后取。顺序为什么重要因为模型训练时如果每个epoch都用同样的样本顺序模型会记住这种顺序导致收敛变慢甚至过拟合。反过来在验证集和测试集上我们反而希望顺序固定保证评价结果可复现。第二是控制是否重复取样。默认训练时是不放回采样也就是一个epoch内每个样本刚好出现一次除非replacementTrue。但有些场景需要放回采样比如类别极度不平衡时你想让少数类的样本反复出现就让Sampler按概率放回地抽抽中哪个算哪个。第三是控制样本分配。这主要针对分布式训练。假设你有8张卡总共有1000个样本如果每张卡都各自跑一个RandomSampler很可能出现两张卡拿到相同样本的情况导致梯度更新时某些样本被重复贡献另外一些样本被整体忽略。DistributedSampler就是为了解决这个分配问题而存在的。搞清楚了这三件事后面所有Sampler的代码你都能看懂它无非就是在回答按什么顺序要不要重复怎么分给每张卡这三个问题。2. 内置Sampler逐个拆解原理、参数与适用场景PyTorch在torch.utils.data里提供了几个开箱即用的Sampler数量不算多但覆盖了绝大多数需求。我把它们一个个拆开讲每个都说说工作原理、参数坑点和适用场景。2.1 SequentialSampler最容易被忽略的老实人from torch.utils.data import SequentialSampler sampler SequentialSampler(data_source)这个类简单到有点无聊。你传入一个data_source只要它实现了__len__它就从0开始按顺序返回range(len(data_source))的全部索引。源码大概只有几行class SequentialSampler(Sampler[int]): def __init__(self, data_source): self.data_source data_source def __iter__(self): return iter(range(len(self.data_source))) def __len__(self): return len(self.data_source)但它非常有用。什么时候用顺序采样验证集和测试集你对模型的评估应该是在稳定、可复现的条件下进行顺序采样保证每次eval结果一致不会因为随机打乱而出现上下波动。需要对齐样本顺序的场景比如可视化模型预测结果时你想把预测输出和原始数据集的下标一一对应顺序采样是最省事的。调试和数据检查当你想复现一个bug顺序采样能让你精确知道每一步取的是哪个样本。这里有个小细节data_source并不强制要求必须是Dataset任何实现了__len__的对象都能传进去。因为在Sampler眼里它只关心你有多少样本至于样本内容是什么它不管。所以你可以用一个裸的range(100)当data_source照样能工作。2.2 RandomSampler训练集默认选择的背后逻辑from torch.utils.data import RandomSampler sampler RandomSampler(data_source, replacementFalse, num_samplesNone, generatorNone)先看参数replacementFalse默认不放回每个epoch内每个样本只出现一次只是顺序随机。num_samples当replacementTrue时它指定总共采样多少个样本如果replacementFalse这个参数会失效因为最多只能把样本取完。generator传入一个torch.Generator对象用来控制随机种子。这个参数在追求可复现实验时极其重要。很多用户第一次看到这个类会疑惑shuffleTrue不就够了吗为什么还要自己动手建RandomSampler 答案在于灵活性。直接用shuffleTrue时你无法控制随机种子而且它背后用的就是RandomSampler。手动创建的好处是g torch.Generator() g.manual_seed(42) sampler RandomSampler(data_source, generatorg) dataloader DataLoader(dataset, samplersampler)这样做你可以在每次实验时复现完全相同的数据读取顺序。对调参和对比实验来说这个能力至关重要——否则你很难判断模型效果提升是来自超参数调整还是仅仅因为运气好碰上了一个更有利的数据顺序。再聊聊replacementTrue。很多人不理解训练集本来就有几万个样本为什么要放回典型场景是过采样少数类。比如二分类里正负样本比例是1:99你不做任何处理模型学到的决策边界会严重偏向负类。但如果用WeightedRandomSampler后面细说或者手动把少数类样本复制几份放进dataset效果常常不如按概率放回抽样来得稳定。放回抽样有个数学上的好处它等价于从原始分布中反复独立采样每个样本被抽中的期望次数正比于它的权重样本之间的相关性更低。2.3 SubsetRandomSampler一行代码划分数据集的利器from torch.utils.data import SubsetRandomSampler indices list(range(len(dataset))) random.shuffle(indices) train_sampler SubsetRandomSampler(indices[:8000]) val_sampler SubsetRandomSampler(indices[8000:]) train_loader DataLoader(dataset, samplertrain_sampler, batch_size32) val_loader DataLoader(dataset, samplerval_sampler, batch_size32)这是我最喜欢的划分数据集方式之一。它接受的参数只有一个indices列表迭代时会把indices打乱并逐个返回。跟上一种方式配合使用你根本不需要提前做train_dataset和val_dataset的切分只要维护两个索引列表就行。它有几个隐藏优点省内存不需要真的拷贝数据去生成两份dataset只是记录索引。适合交叉验证你做K折验证时每一折只需要重新生成一组train/val的indices然后套到同一个dataset上即可代码量极小。配合增量学习当数据集动态变化时比如新增了样本你只需要更新索引不需要重建整个dataset。但也要提醒一句如果你用的是torch.utils.data.Dataset并且dataset[i]内部做了非常重的预处理那SubsetRandomSampler并不能帮你省预处理的时间它只管理索引不管数据读取。预处理该多慢还是多慢。2.4 WeightedRandomSampler解决类别不平衡的救星算法from torch.utils.data import WeightedRandomSampler # 假设每个样本的权重已经算好 weights torch.tensor([1.0, 0.1, 0.1, 1.0, ...]) sampler WeightedRandomSampler(weights, num_sampleslen(weights), replacementTrue)这个Sampler干的事情很简单按照权重概率分布有放回地抽取样本。它的三个核心参数weights一维数组长度必须等于数据集大小。值越大对应样本被抽中的概率越高。它会被内部转换成概率分布源码里会自动做归一化但你传入时不归一化也没关系因为它基于multinomial采样会自动处理。num_samples想要采样的总样本数。这个参数很关键它决定了len(dataloader)是多少。replacement必须设为True才有过采样意义。如果设为False它只做一次不重复采样类别平衡的效果会大打折扣。最经典的使用场景是类别不平衡分类。假设有三类样本数分别是1000、100、10。最简单的权重方案是让每个样本的权重等于其所属类别样本数的倒数class_counts torch.tensor([1000, 100, 10]) class_weights 1.0 / class_counts.float() # [0.001, 0.01, 0.1] sample_weights class_weights[labels] # labels是每个样本的类别索引 sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue)这样少数类样本的单次被抽概率会大幅提升在训练中每个epoch里少数类样本被抽中的次数会显著增多。我测试过在正负比1:50的文本分类任务上用这个Sampler配合focal loss比单纯用focal loss还能再涨两三个点的F1。但这里有个容易搞混的地方WeightedRandomSampler的num_samples通常应该设为多少如果设成len(weights)那么一个epoch内的batch数跟普通训练一样如果你觉得还想让少数类被抽得更狠可以设成len(weights) * oversample_factor比如原来的1.5倍、2倍。但要注意num_samples太大意味着每个epoch更长、迭代更多训练时间更长。我个人的经验是先用len(weights)跑一版观察验证集loss曲线如果少数类欠拟合明显再逐步调大num_samples。2.5 BatchSamcher批次的排班员from torch.utils.data import BatchSampler, SequentialSampler sequential_sampler SequentialSampler(dataset) batch_sampler BatchSampler(sequential_sampler, batch_size32, drop_lastFalse)严格来说BatchSampler不是采样样本而是把其他Sampler产生的索引装进batch。它的输入是一个底层sampler比如SequentialSampler或RandomSampler然后按batch_size把索引分组返回的是一个个索引列表比如[[3, 0, 9, ...], [4, 1, 8, ...], ...]。结合前面讲的逻辑你现在应该能看出DataLoader在shuffleTrue背后的完整流程了创建RandomSampler将其传给BatchSampler设置batch_sizeBatchSampler每次迭代返回一批indicesDataLoader按这批indices从dataset取数据drop_lastTrue这个参数很多人在DataLoader里用过它其实就是BatchSampler的行为当最后剩余的样本不足一个batch时是丢掉还是保留。保留时最后一个batch可能比batch_size小这在很多框架的模型里会引发问题比如BatchNorm在batch太小的时候统计量非常不稳定。所以训练时我通常建议drop_lastTrue。但也不绝对如果整个数据集size刚好能被batch_size整除那设为True还是False无所谓。2.6 DistributedSampler多卡训练的必备品这是所有Sampler里最容易出错的一个没有之一。用DataParallel还相对简单因为它在单进程里通过多线程/多卡并行计算数据读取还是在主进程完成的一般不涉及Sampler。但如果你换成DistributedDataParallelDDP每个GPU上跑的是独立的进程每个进程都有自己的DataLoader此时如果各进程都用独立的RandomSampler就会出现严重的数据重复或遗漏。DistributedSampler的工作机制是这样的它把原始数据集的所有索引按GPU数量和当前进程的rank做均分。比如有8个进程总样本1000个那么每个进程大约负责125个样本的索引。关键是这8份索引互不重叠合起来恰好覆盖全部样本。它还会在shuffleTrue时利用epoch做种子让每个epoch的打乱顺序都不同但不同进程之间的样本仍然不会重复。from torch.utils.data import DataLoader, DistributedSampler sampler DistributedSampler(dataset, num_replicasworld_size, rankrank, shuffleTrue, seed42) dataloader DataLoader(dataset, batch_size64, samplersampler) for epoch in range(num_epochs): sampler.set_epoch(epoch) # 这个必须调保证每个epoch打乱方式不同 for inputs, labels in dataloader: ...这里有两个巨坑我重点标出来必须调用set_epoch(epoch)。如果你不调用每个epoch的shuffle结果都一样等于所有epoch都用相同的数据顺序训练模型会严重过拟合这个顺序。官方文档特别标注了这个要求但实际项目中我见过太多人漏掉。shuffle和sampler不能同时设置。如果你在使用DistributedSampler的情况下还在DataLoader里写shuffleTruePyTorch会直接抛异常。原因很简单shuffleTrue的内部实现就是创建一个RandomSampler而DistributedSampler也是sampler一个DataLoader只能有一个sampler策略。另外还有一个容易忽略的点BatchSampler在分布式下怎么工作答案是DistributedSampler和BatchSampler是叠加关系。DataLoader内部会先拿DistributedSampler做进程间分配然后由BatchSampler在单个rank内组batch。所以你在DDP里通常不需要自己套BatchSampler直接用batch_size参数即可系统会自动处理。3. 实操自定义Sampler的完整方法论内置的Sampler覆盖了90%的常规需求但总有那10%的奇怪需求只能自己写。好在PyTorch的接口设计得相当干净你只需要继承Sampler基类实现两个方法。3.1 自定义Sampler的接口与最小实现from torch.utils.data import Sampler from typing import Iterator class MySampler(Sampler[int]): def __init__(self, data_source): self.data_source data_source self.num_samples len(data_source) def __iter__(self) - Iterator[int]: # 返回一个可迭代对象元素是样本索引 return iter(range(self.num_samples)) def __len__(self) - int: return self.num_samplessamples决定迭代器产出多少个索引__len__决定这个Sampler告诉DataLoader我总共有多少个样本。要注意__len__和__iter__产出的元素数量不一定要一致但强烈建议保持一致否则len(dataloader)和你实际能迭代的batch数会不一致训练循环经常因此卡出各种奇怪问题。有个细节值得注意Sampler里不要直接访问dataset内容它只需要索引。如果你真需要根据样本内容决定采样策略比如按样本难度采样你可以在构造函数里传入样本难度向量而不是把整个dataset传进去。这样职责清晰调试方便。3.2 场景一困难样本优先采样有个我实际做过的需求人脸识别训练时样本难度差异很大简单样本占了大多数模型很快就在简单样本上过拟合对困难样本几乎没学到东西。常规的随机采样下困难样本出现频率太低。我当时的做法是维护一个难度分数列表每训练一个epoch后用模型在验证集上算出每个样本的loss把loss当作难度分数。然后自定义一个HardExampleSampler让loss大的样本在下个epoch更大概率被抽中class HardExampleSampler(Sampler[int]): def __init__(self, num_samples: int, difficulty: torch.Tensor, temperature: float 2.0): self.num_samples num_samples self.difficulty difficulty # 每个样本的难度分数 self.temperature temperature def __iter__(self) - Iterator[int]: # 用softmax把难度转成概率温度越高分布越平滑 probs torch.softmax(self.difficulty / self.temperature, dim0) # 有放回抽样量级保持和原数据集一致 indices torch.multinomial(probs, num_samplesself.num_samples, replacementTrue) return iter(indices.tolist()) def __len__(self) - int: return self.num_samples实现逻辑不复杂但有几个关键决策为什么要用torch.multinomial而不是random.choices因为multinomial直接支持GPU而且和PyTorch的随机数管理机制一致更容易复现。温度参数temperature的作用是什么温度越高概率分布越平缓困难样本的优势被削弱温度越低越聚焦于少数极端困难样本。实践中温度太高会导致采样过于随机太低会导致少数样本被反复抽中模型反而在困难样本上过拟合。我调下来觉得1.5到3之间比较稳妥。输入困难度向量需要每个epoch更新。这意味着你得在训练循环里获取最新loss并对Sampler的属性做更新。Sampler实例可以不是静态的它可以带有状态只要每次__iter__被调用前更新好状态就行。3.3 场景二多数据集按比例混合采样另一个常见需求在预训练阶段出现你手上有两个数据集A和B但A大得多B小得多。如果你直接把两个dataset拼接模型会严重偏向A。你想要的是每个batch里来自A和B的样本数占比符合某个固定比例比如1:1。一种做法是在线对B重复采样但那样代码很丑。更优雅的是自定义一个混合Samplerclass MixedRatioSampler(Sampler[int]): def __init__(self, len_a: int, len_b: int, ratio_a: float 0.5): self.len_a len_a self.len_b len_b self.ratio_a ratio_a self.num_samples len_a len_b def __iter__(self) - Iterator[int]: # 先分别对两个数据集的索引做随机shuffle a_indices torch.randperm(self.len_a).tolist() b_indices torch.randperm(self.len_b).tolist() # 目标数量 num_a int(self.num_samples * self.ratio_a) num_b self.num_samples - num_a # 如果某一方不够就用重复采样补足 if num_a len(a_indices): a_indices (a_indices * (num_a // len(a_indices) 1))[:num_a] if num_b len(b_indices): b_indices (b_indices * (num_b // len(b_indices) 1))[:num_b] # 混合后打乱防止batch内数据全部来自同一数据集 mixed list(range(self.len_a))[:num_a] list(range(self.len_a, self.len_a self.len_b))[:num_b] # 上面这行写错了更正如下 mixed a_indices [i self.len_a for i in b_indices] random.shuffle(mixed) return iter(mixed) def __len__(self) - int: return self.num_samples这里最核心的坑在于当两个数据集大小不一致时你需要用负数的索引方式来区分它们来自哪个数据集或者在batch组装时单独处理。如果两个dataset类型不同你还需要配合自定义collate_fn否则DataLoader会把两个数据集的样本混在一起后无法正确堆叠。这个Sampler的价值是让比例控制发生在索引层面而不是在数据组装层面逻辑更清晰也方便调试。3.4 自定义Sampler的调试技巧自定义Sampler最痛苦的是它返回的是索引你很难直观看到取出来的数据长什么样。我习惯在写完一个自定义Sampler后立刻做一个小实验sampler MyCustomSampler(dataset) print(list(iter(sampler))[:20]) # 打印前20个索引如果发现索引超出len(dataset)范围或者顺序不符合预期大概率就是__iter__里的索引计算写错了。另一个技巧是你可以把Sampler的输出直接传给一个临时DataLoader把batch_size设成1然后打印每个batch的样本label用肉眼验证采样分布是否符合预期。4. 常见问题与排查技巧实录这一节全是实战踩坑我把那些在Stack Overflow上反复出现的问题以及我自己项目中踩过的坑整理成一张速查表再逐个展开说。4.1 踩坑DataLoader同时指定sampler和shuffle报错信息大概是ValueError: sampler option is mutually exclusive with shuffle。原因前面讲过shuffleTrue内部就是创建RandomSampler而DataLoader不允许同时存在两个Sampler。解决办法很简单要么删掉shuffleTrue要么删掉sampler参数。如果你用SubsetRandomSampler或DistributedSampler做随机打乱就不需要再设shuffleTrue。相反如果你就想用shuffleTrue的快捷方式那就别自己传Sampler。4.2 踩坑WeightedRandomSampler的replacement设成False这个坑特别隐蔽。很多人把WeightedRandomSampler当成按权重不放回地打乱顺序然后设置replacementFalse结果发现少数类样本依然很少出现甚至和普通RandomSampler效果差不多。原因在于replacementFalse时每个样本最多被抽一次权重只能影响抽中的优先顺序但最终每个样本还是都会出现一次。少数类样本数量没有变化自然无法缓解类别不平衡。如果要真正过采样必须把replacementTrue。我甚至建议在写这段代码时加个注释防止自己和同事以后又踩一遍# 注意这里replacement必须为True否则过采样不生效 sampler WeightedRandomSampler(weights, num_sampleslen(weights), replacementTrue)4.3 踩坑len(dataloader)和实际batch数不一致这个通常出现在自定义Sampler里。假设你的__len__返回100但你实际产出了103个索引那么PyTorch会用len来计算len(dataloader)和进度条导致训练循环中进度条走完了循环还没跑完或者在分布式环境里不同rank的batch数量不一致导致进程间同步等待甚至死锁。要避免这个问题最好在__iter__里强制对齐def __iter__(self): indices self._generate_indices() # 截断或补零到目标长度 indices indices[: self.num_samples] return iter(indices)4.4 踩坑DistributedSampler忘了set_epoch这个问题如果发生在小规模实验里不容易被发现因为loss曲线还是正常的只是收敛速度变慢了。我见过一个多卡训练任务跑了20个epoch每个epoch的样本顺序完全一样模型收敛速度比预期慢了一大截。排查半天才发现是忘了sampler.set_epoch(epoch)。养成习惯只要用了DistributedSampler训练循环里务必加上for epoch in range(epochs): train_sampler.set_epoch(epoch) ...4.5 性能建议num_workers、prefetch和Sampler的关系Sampler本身是在主进程里执行的num_workers开再多Sampler也只有一个实例在工作。它的计算量通常很小无非就是randperm、multinomial之类一般不会成为瓶颈。但有一点要注意如果自定义Sampler的__iter__里有非常重的计算比如对全量样本算一次前向传播来获取难度分数它会完全阻塞DataLoader的准备工作导致GPU空转。我的建议是把重的计算移到训练循环之外提前算好权重、难度、索引缓存成numpy数组或tensor在Sampler里只做轻量的抽样操作。prefetch_factor这个参数可以调高比如8、16让后台worker提前准备数据但前提是Sampler能及时产出索引。如果你发现训练时GPU利用率上不去先用torch.profiler看一下数据加载时间别一上来就盲目加num_workers。4.6 排查工具一行代码可视化Sampler行为最后分享一个我一直在用的排查技巧。当我觉得某个数据加载行为不符合预期时我会写一个极简脚本把DataLoader产出的索引类别分布打出来from collections import Counter counter Counter() for batch_indices in dataloader.batch_sampler: for idx in batch_indices: label dataset[idx][1].item() counter[label] 1 print(counter)这段代码绕过了collate_fn和tensor堆叠的干扰直接看Sampler产出的原始索引对应的标签分布。第一次跑通之后我建议把这段脚本存成单独文件以后每次调整Sampler都能快速验证行为是否符合预期非常实用。单说技术细节Sampler这块其实不难难的是搞懂它和Dataset、DataLoader、collate_fn之间的协作关系。一旦你把这些拆清楚了后续做任何自定义数据策略都会顺手很多。就我个人的经验来说花一晚上把Sampler相关的源码读一遍绝对比看十篇博客有用——PyTorch的源码注释写得相当友好照着源码自己捋一遍你会突然觉得它真的没那么神秘。最后再分享一个小技巧写自定义Sampler时先在CPU上用小数据集把索引打印出来确认无误后再上大数据集和GPU训练。这个习惯帮我省下了太多排查时间。数据加载是深度学习里最枯燥但也最不能出错的环节Sampler就是你控制这里的核心抓手用好它你的训练流程会顺手一大截。
返回列表