恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
PyTorch DistributedSampler 多卡数据分片避坑指南
首页
资讯中心
/
PyTorch DistributedSampler 多卡数据分片避坑指南
PyTorch DistributedSampler 多卡数据分片避坑指南
发布时间:2026/10/1 2:47:30
DistributedSampler 这个类藏在torch.utils.data.distributed底下两百来行代码很多人第一次写多卡训练脚本时都是随手抄个模板就上了DataLoader 里塞个 sampler外面套层 DDP跑起来 loss 在掉就以为万事大吉。我也是这么过来的直到有次八卡训一个千万级样本的排序模型训完在单卡上复现验证指标怎么都对不上差了将近两个千分点排查了两天才发现问题出在分片逻辑和每个 epoch 的随机种子递推上。这篇文章把 DistributedSampler 从定位、参数语义、源码行为到实战踩坑完整拆一遍适合正在写多卡训练脚本的工程师也适合刚接触 DDP、被“每张卡的数据为什么不一样”绕晕的朋友。看完你至少能搞清楚三件事数据是怎么分到每张卡上的、drop_last和set_epoch在什么情况下会咬你、以及哪些指标异常的锅不该甩给模型。1. 把 DistributedSampler 放回它该在的位置1.1 数据并行训练里数据是怎么被“分掉”的先讲大图。DDP 这种数据并行模式本质是把同一份模型复制到 N 张卡上每张卡吃不同的数据各自算梯度反向传播时把所有卡的梯度求和再平均保证各卡上的参数始终一致。这里有个隐含前提每张卡拿到的数据必须不同。否则八张卡各算一遍同样的 batch等于把单卡的 batch 重复算八次白白多付通信开销训练速度反而更慢。把数据分给不同的卡这件事PyTorch 没有帮你自动做需要显式指定。最原始的做法是在 Dataset 里按 rank 切dataset[i]只返回属于当前 rank 的样本但写起来繁琐、容易出错而且切分方式一旦固定后续想换机器数量就得重写。DistributedSampler 正是为了解决这个场景它接收数据集对象和进程组规模自动算出“当前这个 rank 该拿哪些下标”交给 DataLoader 去取数。你可以把它当成发牌员。一副牌洗好按顺序轮流发每个人手里的牌互不重叠发到最后一张为止。发牌员只管分牌不参与出牌也不关心你打得好不好。这个类比后面会反复用到因为它能解释绝大多数困惑为什么调整 world_size 会影响样本数量、为什么小数据集会出现重复样本、为什么改了 batch_size 但每卡步数没变。1.2 不切数据会付出什么代价第一层代价是算力浪费。八张卡跑和一卡跑的速度差不了多少因为你把同样的计算做了八遍。第二层代价是有效 batch 没变。很多人以为上了多卡batch 就自动变大了其实没有每张卡还是处理batch_size条数据全局 batch 变成batch_size × world_size这件事只是“事实上的效果”学习率、warmup 步数、优化器配置都得你自己跟着改。第三层代价最隐蔽梯度方向出现系统性偏差。如果代码里加了梯度累积或者对 loss 做了奇怪的归一化重复数据会让某些样本被反复加权这种情况在小数据集上跑几百步都看不出来等一轮训练结束对比指标才发现曲线很怪。还有一种更常见的翻车方式切了但切重复了。比如多个 rank 用了同一个 rank 值取数据或者进程组压根没初始化好dist.get_rank()抛异常被吞掉后走了默认值。这类 bug 在日志里的表现是 loss 曲线异常平滑因为所有卡的梯度完全一样等于退化成单卡。排查方法很土但有效在每个 rank 上打印第一批数据的下标对比一下有没有交集。1.3 它明确不管的那些事新手最容易犯的错是把 DistributedSampler 当成“多卡训练的解药”。它负责的只有一件事告诉 DataLoader 当前进程该读哪些下标。梯度同步归 DDP 的 autograd hook 管跨卡 BatchNorm 要手动换成SyncBatchNorm日志要在 rank 0 上打印学习率调度、指标聚合、checkpoint 保存全得自己安排。把这几件事混在一起想排查问题时就会找错方向。比如训练 loss 正常但验证指标偏低有人第一反应是 sampler 分片出错实际上更常见的锅是验证集只在一个 rank 上跑、其他 rank 提前退出了导致指标计算覆盖不全。2. 核心参数的逐行拆解2.1 num_replicas 和 rank 从哪里来先把构造签名摆出来逐项对照着看。参数默认值取值范围实际语义num_replicasNone正整数不传时取dist.get_world_size()rankNone[0, num_replicas)不传时取dist.get_rank()shuffleTruebool是否打乱下标顺序seed0int随机种子基值真实种子是seed epochdrop_lastFalsebool是否丢弃尾部不足以整除的样本epoch0int由set_epoch()写入参与种子计算num_replicas和rank这两个参数默认值虽然好用但有个陷阱它们依赖进程组已经初始化。如果你在dist.init_process_group()之前就构造 sampler会直接抛ValueError: Default process group has not been initialized。这个报错信息其实挺明确的但很多人是在 DataLoader 构造链里踩的堆栈很深一眼看不出问题在哪。稳妥的写法是手动传参把world_size和rank显式写出来代码可读性和可控性都更好。另外多机场景下rank是全局排名不是机器内排名。八台机器每台八卡总world_size是 64rank从 0 到 63 全局唯一不要用本地的LOCAL_RANK去构造 sampler否则相邻机器上的分片会完全重合。2.2 shuffle 与 epoch 的种子递推关系这是最容易出问题的地方。源码里的逻辑大概是这样if self.shuffle: g torch.Generator() g.manual_seed(self.seed self.epoch) indices torch.randperm(len(self.dataset), generatorg).tolist() else: indices list(range(len(self.dataset)))关键在于这个self.seed self.epoch。self.epoch默认是 0只有调用set_epoch(epoch)才会更新。这意味着如果你忘了在每个 epoch 开始时调set_epoch那么所有 epoch 用的都是同一个随机序列数据顺序完全一致。后果是什么模型每个 epoch 看到的数据排列一模一样。在数据量小、epoch 数多的任务上这会明显拖慢收敛甚至让模型学到某种与顺序相关的伪规律。这个坑特别难发现因为 loss 曲线看起来只是“收敛慢一点”不会有任何报错。正确的写法是在训练循环开头固定加一行for epoch in range(num_epochs): sampler.set_epoch(epoch) for batch in dataloader: ...顺带说一句set_epoch只影响 shuffle 的随机序列不影响num_samples和分片逻辑所以中途改 epoch 值不会引发形状不一致的问题可以放心调。2.3 drop_last 的真实语义以及它和 DataLoader 那个同名参数的区别这是最容易被混淆的一对参数。DistributedSampler 自己的drop_last管的是“数据集长度除不尽卡数时怎么办”而 DataLoader 的drop_last管的是“最后一个不满 batch 的批次要不要丢”。两者名字一样作用域完全不同。分片逻辑可以简化成下面几步假设数据集长度是L卡数是N。如果drop_lastTrue先算出num_samples L // N然后total_size num_samples * N多出来的L - total_size条样本直接被丢掉最多丢N - 1条。如果drop_lastFalse先算出num_samples ceil(L / N)再算出total_size num_samples * N。因为total_size可能大于L需要补齐补齐的方式是从数据集头部按顺序复制若干条样本填到末尾也就是indices indices[:padding_size]。所以当L不能被N整除时头部会有若干样本被重复采样一次。这两种策略各有取舍。训练场景下我一般用drop_lastFalse让每条样本都被用上重复几条影响可以忽略。但评估场景要注意如果验证集只有几千条样本八卡跑下来头部的重复样本会让指标有轻微的系统偏差样本越少偏差越明显。这时候更推荐的做法是验证时不用 DistributedSampler或者用完整的补零方案后只在 rank 0 上聚合计算。2.4 为什么__len__这个不起眼的接口很重要DistributedSampler 的__len__返回的是num_samples也就是每张卡负责的样本数。所有 rank 的这个值必须相同。为什么因为 DDP 的反向传播里有集合通信操作所有 rank 必须调用相同次数的all_reduce一旦某个 rank 早一步跑完循环退出其他 rank 就会永远卡在通信上。DistributedSampler 通过 padding 或者截断强行让每个 rank 的样本数对齐到num_samples就是为这件事服务。这一点很多人没意识到——它不只是“切分数据”更是“把切分做得足够均匀让所有卡步调一致”。如果你的数据集长度动态变化或者自己写了别的 sampler这个对齐就必须自己保证否则训练会在某个随机时刻挂住而且挂住的位置没有规律极难定位。3. 从零跑通一个多卡训练流程3.1 环境与版本对齐顺带说说那个 DLL 报错多卡训练对版本比较敏感尤其是 CUDA 运行时和框架编译时用的 CUDA 版本必须匹配。安装时把三个包放在一条命令里能最大程度避免版本错配pip install torch torchvision torchaudio --index-url 索引地址--index-url这个参数的作用是覆盖默认的包索引团队内网可以指向自建的索引服务加快下载。三个包要写在一起装原因是它们共享底层的 C 扩展分开装很容易出现某个包被覆盖升级、另外两个还是旧版的情况。有朋友在 Windows 上遇到过这个报错OSError: [WinError 1114] 动态链接库(DLL)初始化例程失败。 error loading ....\site-packages\torch\lib\c10.dll or one of its dependencies.这个错误的字面意思是c10.dll加载失败但对绝大多数人来说真正的原因不在 torch 本身而在于它的依赖项。按我处理过的案例排序概率最高的三种缺少或版本过旧的 Microsoft Visual C Redistributable。装最新版的 x64 运行库重启终端再试这一条能解决大半问题。conda 和 pip 混装导致 DLL 冲突。同一个环境里先 conda 装了 torch后来又 pip 装了一遍两套运行库文件混在一起加载顺序不确定。处理方式是把环境删掉重建安装方式二选一别混着来。PATH里存在同名的其他运行库比如别的软件自带的数学库被优先加载。检查环境变量把可疑的目录临时移出再试。如果这三步都试过还不行就检查 Python 版本和框架版本是否匹配。较新的 Python 大版本通常需要较新的框架版本支持具体对应关系以官方发布说明为准不要凭感觉猜。3.2 单机多卡的最小可运行示例下面这份脚本可以直接拿去当模板跑通之后再往里面加自己的模型和数据处理逻辑。import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import Dataset, DataLoader from torch.utils.data.distributed import DistributedSampler class ToyDataset(Dataset): def __init__(self, size10000): self.size size def __len__(self): return self.size def __getitem__(self, idx): x torch.randn(32) y torch.tensor(idx % 10) return x, y def main(): dist.init_process_group(backendnccl) local_rank int(os.environ[LOCAL_RANK]) global_rank dist.get_rank() world_size dist.get_world_size() torch.cuda.set_device(local_rank) train_set ToyDataset() sampler DistributedSampler( train_set, num_replicasworld_size, rankglobal_rank, shuffleTrue, seed42, drop_lastFalse, ) loader DataLoader( train_set, batch_size64, samplersampler, num_workers4, pin_memoryTrue, drop_lastTrue, ) model torch.nn.Linear(32, 10).cuda(local_rank) model DDP(model, device_ids[local_rank]) opt torch.optim.AdamW(model.parameters(), lr1e-3) loss_fn torch.nn.CrossEntropyLoss() for epoch in range(20): sampler.set_epoch(epoch) model.train() total 0.0 for step, (x, y) in enumerate(loader): x x.cuda(local_rank, non_blockingTrue) y y.cuda(local_rank, non_blockingTrue) opt.zero_grad() out model(x) loss loss_fn(out, y) loss.backward() opt.step() total loss.item() if global_rank 0: print(fepoch {epoch} avg loss {total / len(loader):.4f}) dist.destroy_process_group() if __name__ __main__: main()启动命令torchrun --nproc_per_node4 train.py几个细节值得单独拎出来说。torch.cuda.set_device(local_rank)必须在模型搬到 GPU 之前调用否则所有进程都会默认用 0 号卡显存瞬间打满。non_blockingTrue配合pin_memoryTrue能让数据拷贝和计算重叠在小模型上可能看不出差别大模型上是很实在的加速。global_rank 0的判断是为了避免四个进程抢着往同一个日志文件里写。3.3 单卡代码迁移到 DDP 的改造清单如果你手上已经有一份能跑的单卡训练代码改造成多卡大概需要动这几个地方按重要性排序第一init_process_group和set_device加到最前面destroy_process_group放到最后。第二模型用 DDP 包一层注意包之前要先把模型搬到对应的卡上包之后不要再手动搬。第三DataLoader 里去掉shuffleTrue改成传sampler。这两个参数是互斥的同时传会直接抛出ValueError: sampler option is mutually exclusive with shuffle。第四每卡batch_size要除以world_size如果你想让全局 batch 保持不变的话。第五学习率一般要跟着全局 batch 一起调常见做法是线性缩放或者平方根缩放具体选哪个跟优化器和任务都有关跑几个 ablation 对比一下最靠谱。第六打印日志、保存 checkpoint、写 tensorboard 的地方都加rank 0判断。还有一个容易漏的点如果代码里用了torch.manual_seed注意它只影响当前进程DistributedSampler内部的随机性由seed epoch单独控制两者互不干扰。想要完全复现实验结果两个种子都要固定。3.4 验证集和测试集该怎么取样这是我在实际项目里见过最多分歧的地方。三种常见做法各有适用场景。第一种验证集也用DistributedSampler但设shuffleFalse。优点是实现简单和训练集共用一套代码路径缺点是每个 rank 只看一部分数据要拿全局指标就得做all_gather把各卡的预测收集起来而且当数据集长度不能整除卡数时头部样本会被重复计入指标有轻微偏差。第二种验证集只在 rank 0 上跑其他 rank 用dist.barrier()等待。实现最简单指标最准确缺点是其他卡在等整体评估耗时变长。数据集小的时候推荐这种做法。第三种每个 rank 独立跑全量验证集然后用all_reduce把 loss 求平均、把预测做all_gather后只在 rank 0 上算指标。计算量是world_size倍但完全避免了分片带来的偏差适合验证集不大又追求指标精确的场景。我自己的习惯是验证集超过两万条就用第一种注意精度损失两万条以下就用第二种简单省心。4. 进阶场景与实战取舍4.1 多机多卡下 rank 的语义会变单机多卡时你可能习惯用LOCAL_RANK到处传多机场景下这个习惯会翻车。LOCAL_RANK是单机内的 GPU 序号取值是 0 到 7而rank是全局进程序号取值是 0 到 63。DistributedSampler 必须用全局 rank 构造否则第二台机器上的进程会拿到和第一台一样的分片。数据分片的间隔取样方式也值得留意。源码里用的是indices[rank::num_replicas]也就是每隔num_replicas取一个下标而不是连续切片。这种方式有个附带好处如果数据集是按类别或者按时间排序的间隔取样能让每张卡上的类别分布、时间分布更接近全集训练更均衡。但它也有风险如果数据集本身带有周期性结构恰好周期和卡数相同间隔取样会让每张卡都拿到同一位置的样本分布反而严重偏斜。真遇到这种情况就自己在 Dataset 层面做一次预打乱或者改用连续切片手动分配。4.2 IterableDataset 没法直接用这个 sampler流式数据集、日志逐行读取、在线生成样本这类场景通常是IterableDataset。它没有__len__也没有__getitem__DistributedSampler 那套基于下标的分片逻辑完全用不上强行传进去会直接报错。这类数据的正确切分方式是按 rank 做 shard。常见的做法是在数据源层面按序号取模每个 rank 只处理属于自己的那一条流。如果用的是支持 shard 接口的第三方数据集库通常会有dataset.shard(num_shardsworld_size, indexrank)这样的方法一行就能搞定。要注意的是这种切分方式在卡数变化时数据划分也跟着变断点续训时要格外小心否则会重复消费或者漏掉样本。另外流式数据很难保证每个 rank 的样本数严格相等如果 batch 数量不一致DDP 会在通信阶段挂住。稳妥做法是让数据源本身可控或者手动截断到各 rank 的最小长度。4.3 和 DataLoader 其他参数打架的几种情况除了shuffle还有几组参数是不能共存的。batch_sampler和batch_size、shuffle、sampler、drop_last都互斥因为前者已经完整定义了怎么组 batch。sampler和shuffle互斥前面提过。这些约束 PyTorch 会主动检查并抛异常报错信息还算清楚但第一次看到容易懵。num_workers和 sampler 的配合方式值得单独说明。sampler 生成的 indices 是在主进程算好的然后随 Dataset 一起打包送给各个 worker所以set_epoch只要在主进程调用就能生效不需要在每个 worker 里重复调用。反过来如果你在 Dataset 的__getitem__里做随机数据增强那部分随机性由 DataLoader 的 worker 种子控制和 sampler 的种子是两套独立机制。想完全复现实验两处都要固定。persistent_workersTrue会让 worker 进程跨 epoch 复用省掉反复创建进程的开销对短 epoch 的任务提速明显。它和 sampler 没有冲突每次 epoch 开始 DataLoader 都会创建新的迭代器把新的 indices 序列发给 workerset_epoch照样生效。4.4 断点续训时 sampler 的状态要一起存模型参数、优化器状态、学习率调度器这些大家都记得存但 epoch 号经常被漏掉。而 epoch 号恰好是set_epoch的输入直接决定了后续每个 epoch 的 shuffle 序列。如果你在第 7 个 epoch 中断恢复时从第 8 个 epoch 继续那就得把 epoch 号存进 checkpoint恢复时读出 8 再调用sampler.set_epoch(8)。少了这一步shuffle 序列会从 0 重新开始训练数据的顺序和中断前不一致。对最终精度的影响通常不大但如果你的实验需要严格复现这个细节就必须处理。还有一个容易被忽略的限制中途改变卡数会导致分片方式不兼容。比如用 8 卡存了 checkpoint用 4 卡恢复虽然参数能加载上但每个 rank 拿到的数据集合完全不同了。这种情况没有通用的完美解法只能重新开始一轮或者接受一定的训练轨迹偏移。5. 报错与异常排查速查5.1 样本重复、缺失、顺序错乱的排查路径这类问题不会报错只能靠观察。我一般用下面这套流程定位。第一步在每个 rank 上打印len(dataloader)和前三个 batch 的样本下标看看各卡长度是否一致、下标有没有交集。第二步写个临时脚本把sampler的__iter__结果收集起来做并集和range(len(dataset))对比能直接看出哪些样本被丢了、哪些被重复采样了。第三步检查rank和num_replicas的取值是否和实际进程组匹配这一步能抓到绝大部分低级错误。如果训练 loss 曲线异常平滑八成是所有卡拿到的数据一样优先怀疑进程组没初始化或者 rank 传错。如果 loss 震荡比单卡明显可能是分片导致每卡 batch 内的类别分布太偏考虑换连续切片或者调整数据布局。5.2 训练中途卡住不动的几种原因卡住是分布式训练里最折磨人的问题因为没有任何报错只能看到 GPU 利用率掉到零。按我的经验概率从高到低排一下各 rank 步数不一致是最常见的。如果 sampler 的drop_lastTrue而 DataLoader 的drop_lastFalse最后一个不满 batch 的批次会被某些 rank 发出去、被另一些 rank 跳过通信就配不上对了。第二个原因是某个 rank 上的数据出现坏样本抛了异常退出其他 rank 还傻等在all_reduce上。这种情况去看退出进程的 stderr 就能找到线索关键是别只看 rank 0 的日志。第三个原因是验证阶段的不对称。如果只用 rank 0 算指标其他 rank 没调用barrier就直接进了下一轮训练通信次数又会错位。第四个原因是num_workers太大导致内存不足worker 被系统杀掉主进程一直等数据。这个在容器里跑的时候尤其常见因为容器有额外的内存限制。排查手段上torch.distributed提供了一个环境变量可以打印集合通信的调用详情同时在启动命令里加上超时参数能让卡住的进程主动报错退出而不是无限等下去。日志层面我习惯在训练循环里每隔 N 步打印一次带 rank 前缀的时间戳哪张卡停了、停在哪个位置一目了然。5.3 从报错信息快速定位问题类型下面这张表是我自己攒的工具按报错关键词对应处理方向。报错/现象关键词大概率原因处理方向sampler option is mutually exclusive with shuffleDataLoader 同时传了sampler和shuffleTrue删掉shuffle参数Default process group has not been initialized在init_process_group之前构造 sampler调整构造顺序或手动传rank/num_replicasWinError 1114/c10.dll加载失败运行库缺失、环境混装、PATH 冲突补装 VC 运行库、重建环境、统一安装方式训练中途无报错卡死各 rank 步数不一致或某 rank 异常退出对齐drop_last、检查所有 rank 日志指标比单卡明显偏低验证集分片导致覆盖不全或重复计数换验证集取样策略做全局聚合loss 曲线异常平滑所有 rank 拿到相同数据核对rank与num_replicas每个 epoch 数据顺序不变忘了调set_epoch训练循环开头补上5.4 我在实际项目里踩过的坑第一个坑也是最费时间的一个验证集用了 DistributedSampler 但忘了关 shuffle导致每次评估的样本划分都不一样。指标抖得厉害一开始以为是模型不稳定调了半天超参最后打印 sampler 的配置才反应过来。评估类 sampler 一定要shuffleFalse没有例外。第二个坑是drop_last的连锁反应。我把 sampler 的drop_last设成了 True卡数是 8数据集长度恰好是 8 的倍数加 3于是每个 epoch 都有 3 条样本永远学不到。这个数据集里那 3 条恰好是少数类样本训了两天才发现模型对那个类的召回一直是零。数据量小或者类别不均衡的任务drop_last建议设成 False。第三个坑和 worker 数量有关。num_workers16在本地机器上跑得好好的换到容器里直接 OOM因为每个 worker 都会复制一份数据集对象大数据集下内存开销是成倍增长的。我的经验值是默认给 4根据 IO 压力再往上调调整时盯着内存监控。第四个坑是多机训练时 rank 传错。我用LOCAL_RANK构造了 sampler单机测试完全正常上到两台机器之后一半的数据被重复处理了而且因为每台机器内部看起来“没问题”排查方向一开始完全错了。从那以后我养成习惯在训练脚本开头把rank、local_rank、world_size三个值一起打出来一眼就能确认。最后分享一个小技巧。调试分片逻辑时不要用真实数据集写个Dataset只返回下标本身、长度设成 23 或者 37 这种质数然后在每个 rank 上跑一个 epoch 把所有下标打印出来。你能一次性看清重复样本出现在哪个位置、哪些样本被丢了、各 rank 的样本数是否相等。这个小脚本我到现在还在用比读源码快得多也直观得多。真正把 DistributedSampler 用明白靠的不是背参数表而是知道每个 epoch 那些数字是怎么一步步算出来的。