【Bug已解决】[Question] Manual Dataset Sharding per GPU Rank with Accelerate + DistributedSampler (Avoid

发布时间:2026/8/1 3:46:48

【Bug已解决】[Question] Manual Dataset Sharding per GPU Rank with Accelerate + DistributedSampler (Avoid 【Bug已解决】[Question] Manual Dataset Sharding per GPU Rank with Accelerate DistributedSampler (Avoid Double DataLoader Length Split) 解决方案一、现象长什么样用 Accelerate 做分布式训练想手动按 rank 分片数据集而不是完全交给框架于是自己用DistributedSampler又叠加了一些手工切分结果出现len(dataloader)比预期小了一半或 N 分之一的平方本以为每 rank 看到1/N数据实际却只有1/N²。训练 epoch 提前结束 / 样本数对不上每个 rank 的 batch 数只有正确值的1/N一个 epoch 里大量样本没被训练到。日志困惑print(len(dataloader))在 rank0 打出 250而总数据 1000、4 卡本应每 rank 250——看起来「对」但若你之前还手动dataset dataset[rank::world]实际就变成 62静默丢数据。特征只在「手动分片 DistributedSampler Accelerate.prepare」三者叠用时出现。不报错是「长度算错、数据静默变少」这类难发现的错。用户的核心疑问就是标题怎么正确手动分片、又不被二次分割长度。本质DistributedSampler本身已经把数据集按 rank 切成1/N并让len(dataloader)反映「每 rank 的长度」如果你在创建 sampler 之前又手动把 dataset 切了一遍dataset[rank::world]或者 Accelerate.prepare 再分一次长度就被除了两次 →1/N²。二、背景要理清「为什么长度会被除两次」得先知道DistributedSampler做了什么给定数据集D长度L和world_size NDistributedSampler会给每个 rank 分配一个不重叠的下标子集大小约ceil(L/N)或floor(L/N)。包装后的DataLoader的len()返回的是该 rank 分到的样本数即≈ L/N而不是L。所以「正确的手动分片」其实只需要用DistributedSampler不要再去切 dataset 本身。常见双重分割的来源手动预切 datasetmy_ds full_ds[rank::world]再用DistributedSampler(my_ds)。此时my_ds已经是L/Nsampler 又把它切成1/N→ 最终L/N²。这是最典型的「双重分割」。Accelerate.prepare 再分一次如果my_ds已用 DistributedSampler 分好prepare 又识别为「需分片」再裹一层 → 同上与上一个 Feature 议题同源但这里聚焦长度计算。drop_last与长度取整的误解drop_lastTrue会丢弃不能整除的尾部若用户同时手动处理了尾部又叠一层长度再次错位。一句话DistributedSampler已负责按 rank 切片并修正len额外手动切片或 prepare 再分都会让长度二次除以 N。三、根因根因是对DistributedSampler的「已切片 已修正 len」语义理解错位又叠加了一层手动切片或 prepare 分片三层第一层主因手动预切 dataset 与 sampler 切片重复。full_ds[rank::world]已是1/N再套DistributedSampler→1/N²。用户误以为「sampler 只是在 dataset 内打乱不改长度」其实它会按 rank 重排并让len反映每 rank 长度。第二层Accelerate.prepare 的再分片未关。若 dataset 已带 DistributedSamplerprepare 默认再裹一层分布式逻辑长度再除一次。需要让 prepare「识别已分片、跳过」见相关议题否则双重分割。第三层长度校验缺失错误静默。没有任何地方断言「每 rank 数据量 L/N」。于是1/N²静默发生训练照跑但样本大量未训练loss 异常却找不到原因。一句话手动切片 sampler 切片 prepare 再分叠加长度被多次除以 N且缺断言导致静默丢数据。四、最小可运行复现下面用纯 Python 模拟「DistributedSampler已切片后再手动预切 → 长度二次除以 N」的控制流不需要 GPUfrom dataclasses import dataclass dataclass class Split: total: int per_rank: int def distributed_sampler_len(total: int, world: int) - int: # DistributedSampler 让 len(dataloader) ceil(total / world) return (total world - 1) // world def scenario_double_split(total, world, rank): 错误先手动切 dataset[rank::world]再套 sampler。 manual total // world # 手动切后长度 total/N sampler_len distributed_sampler_len(manual, world) # 再除一次 return sampler_len def scenario_correct(total, world, rank): 正确只用 sampler不手动切 dataset。 return distributed_sampler_len(total, world) def main(): total, world 1000, 4 wrong scenario_double_split(total, world, 0) right scenario_correct(total, world, 0) print(f总数据 {total}, {world} 卡) print(f双重分割后每 rank 长度: {wrong} (应为 {right}, 实际少 {right - wrong})) print(f正确每 rank 长度: {right}) if __name__ __main__: main()跑出来双重分割后每 rank 长度 62、正确应为 250——直观展示了「长度被二次除以 N、静默丢 3/4 数据」。五、解决方案第一层最小直接修复最省事的救火只用DistributedSampler分片绝不手动预切 dataset并让 Accelerate.prepare 跳过再分。这样长度天然是1/Nfrom torch.utils.data import DataLoader, DistributedSampler from accelerate import Accelerator accelerator Accelerator() full_ds MyDataset(...) # 完整数据集不要手动切 # 只用 DistributedSampler 负责按 rank 分片长度自动 1/N sampler DistributedSampler( full_ds, num_replicasaccelerator.num_processes, rankaccelerator.process_index, shuffleTrue, ) dl DataLoader(full_ds, batch_size8, samplersampler, shuffleFalse) # 关键告诉 prepare 别再分片否则长度再除一次 prepared accelerator.prepare_data_loader(dl, split_batchesFalse, # 视版本用 make_sharded_dataloaderFalse ) print(每 rank 长度:, len(prepared)) # 应为 1000/4 250如果你确实想手动预切例如要按哈希分片而非顺序那就不要再套 DistributedSampler而是用RandomSampler/SequentialSampler在已切数据集上# 手动切后用普通 sampler不再按 rank 分 my_ds full_ds[accelerator.process_index::accelerator.num_processes] dl DataLoader(my_ds, batch_size8, shuffleTrue) # 无 DistributedSampler二选一不要两者都要。六、解决方案第二层结构性改进第一层是「二选一不要叠」第二层是「封装一个分片工厂强制单一分片来源并断言长度正确」从设计上消灭双重分割from dataclasses import dataclass from typing import Literal dataclass class ShardingPlan: mode: Literal[sampler, manual] # 唯一分片来源二选一 def build(self, dataset, world, rank): if self.mode sampler: # 只用 DistributedSampler不手动切 dataset from torch.utils.data import DistributedSampler sampler DistributedSampler(dataset, num_replicasworld, rankrank) return dataset, sampler # dataset 仍是完整 else: # 手动切 dataset用普通 sampler不再按 rank 分 from torch.utils.data import RandomSampler sub dataset[rank::world] return sub, RandomSampler(sub) def assert_length_ok(dataloader, total, world): 断言每 rank 长度 ceil(total/world)杜绝静默双重分割。 expected (total world - 1) // world actual len(dataloader) assert abs(actual - expected) 1, ( f每 rank 长度 {actual} 与期望 {expected} 差距过大 f疑似被多次分割双重分片。请只使用一种分片方式。 ) return actual # 用法 plan ShardingPlan(modesampler) ds, sampler plan.build(full_ds, world4, rank0) dl DataLoader(ds, batch_size8, samplersampler, shuffleFalse) assert_length_ok(dl, total1000, world4) # 250断言通过关键ShardingPlan.mode强制「sampler 或 manual」二选一代码层面禁止叠加。assert_length_ok在训练前断言每 rank 长度双重分割立刻暴露而非静默。七、解决方案第三层断言 / CI 守护把「单一分片来源」「长度不被二次除」「手动与 sampler 互斥」固化成测试import pytest def test_sampler_mode_no_manual_split(): plan ShardingPlan(modesampler) ds, sampler plan.build(range(1000), world4, rank0) # dataset 仍是完整的未被手动切 assert len(list(ds)) 1000 def test_manual_mode_uses_random_sampler(): plan ShardingPlan(modemanual) ds, sampler plan.build(range(1000), world4, rank0) assert len(list(ds)) 250 # 手动切后 1/N def test_length_ok_for_sampler(): from torch.utils.data import DataLoader, DistributedSampler ds range(1000) sampler DistributedSampler(ds, num_replicas4, rank0) dl DataLoader(ds, batch_size8, samplersampler) assert_length_ok(dl, total1000, world4) # 250 def test_double_split_detected(): # 模拟双重分割手动切后再套 sampler manual range(1000)[0::4] # 250 sampler DistributedSampler(manual, num_replicas4, rank0) dl DataLoader(manual, batch_size8, samplersampler) with pytest.raises(AssertionError): assert_length_ok(dl, total1000, world4) # 实际 62断言失败 def test_no_overlap_between_ranks(): plan ShardingPlan(modemanual) seen set() for r in range(4): ds, _ plan.build(range(1000), world4, rankr) seen | set(ds) assert len(seen) 1000 # 不重叠、不遗漏再加一个端到端回归4 卡训练每 rank 长度正确、数据不重不漏def test_distributed_length_correct(): total, world 1000, 4 plan ShardingPlan(modesampler) for rank in range(world): ds, sampler plan.build(range(total), world, rank) dl DataLoader(ds, batch_size8, samplersampler) assert_length_ok(dl, total, world)八、排查清单看len(dataloader)是否远小于total/N或 epoch 样本数对不上 → 是双重分割。检查是否既手动dataset[rank::world]又套DistributedSampler只留一个。检查 Accelerate.prepare 是否又分了一次应用 opt-out 跳过。临时救火只用 DistributedSampler不手动切 dataset或只手动切不用 DistributedSampler。训练前assert_length_ok(dl, total, world)断言每 rank 长度提前暴露双重分割。升级 accelerate 到合了「已分片跳过」的版本并跑上面的「双重分割检测」用例。若用drop_last长度期望用total//world//batch* batch口径断言时留容差。九、小结手动按 rank 分片时长度被「双重分割」不是 Accelerate 坏了而是**DistributedSampler已按 rank 切片并修正len为1/N你又手动预切 dataset 或让 prepare 再分一次长度被二次除以 N、静默丢 3/4 数据**。最小修复是「只用 sampler 或只手动切二选一」结构性修复是封装ShardingPlan强制单一分片来源 assert_length_ok训练前断言最后用 pytest 把「单一来源」「长度不被二次除」「不重不漏」锁死。抓住「分片只能有一处来源、且len必须断言等于ceil(total/N)」这条所有分布式数据集长度错乱都能照此排查。

相关新闻