数据截至 (上游 commit fd01e35c83d8)
03 · DataLoader 分片
这一章讲什么: 4 个进程跑同一个 DataLoader,每个进程怎么拿到不同的四分之一、批数如何对齐、评估指标又怎么不被「对齐手段」污染。这是 Accelerate 里最接地气也最容易踩坑的一块。
1. 它要解决的小问题
分布式训练要求 所有进程步数一致(不然集合通信会挂起),但数据要互不重复。一个普通 DataLoader 直接给 4 个进程用,每个进程都会遍历全量数据——既浪费又让梯度重复计算。而自己写 DistributedSampler 又会遇到:数据集大小不能被「batch size × 进程数」整除时,有的进程多一个 batch、有的少一个,评估指标对不齐。
2. 直觉
Accelerate 提供两种切法:
- 各取所需(shard):改造采样器。把所有 batch 按顺序编号,进程 i 只取第 i、i+P、i+2P…个 batch(P = 进程数)。每个进程独立读数据,零通信。
- 主取后分(dispatch):只有主进程迭代 DataLoader,取出的 batch 广播给所有人,再按进程切片。适合 IterableDataset(无法建索引)或预处理很贵的场景。
步数对齐用一个朴素手段:尾部补齐——最后一个不整除的批次,从数据集开头借样本填满,让所有进程都多走一步。评估时再用「我补了几个」这个记录把重复样本裁掉。
3. 图示
两种策略的数据流(3 个进程,每进程 batch=2):
策略 A:BatchSamplerShard(shard)
采样器产出 batch 序列: [b0] [b1] [b2] [b3] [b4] [b5*尾]
进程0: b0, b3 进程1: b1, b4 进程2: b2, b5(含借来的样本)
—— 各自读数据,无通信
策略 B:DataLoaderDispatcher(dispatch)
主进程: iter → b0,b1,b2 → concatenate → broadcast ─┐
▼
进程0 取切片0 进程1 取切片1 进程2 取切片2(每步一次广播通信)
包装后的 DataLoader 与用户代码之间隔着透明代理层:
for batch in prepared_dl ──▶ DataLoaderShard.__iter__ ──▶ BatchSamplerShard 出本分片的索引
│ │
│ ├─ 预取下一个 batch 判断是否结束
│ └─ 当前 batch 提前 send_to_device
▼
isinstance(prepared_dl, DataLoader) == True (__class__ 被伪装成基类)
4. 原理演示(示意代码)
shard 策略的采样器可以压缩成:
class BatchSamplerShard:
def __iter__(self):
for idx, batch in enumerate(self.batch_sampler):
if idx % self.num_processes == self.process_index: # 轮到我才 yield
yield batch
def __len__(self):
# 向上取整:尾部补齐后每个进程批数相同
return ceil(len(self.batch_sampler) / self.num_processes)
split_batches=True 时则不是「各拿整个 batch」,而是把每个 batch 切成 P 片、进程 i 拿第 i 片——观察到的 batch size 不变,全局 batch size 也不随进程数放大。
评估修尾的逻辑:
def gather_for_metrics(self, data):
gathered = self.gather(data) # 各进程结果拼接
if self.gradient_state.end_of_dataloader: # dataloader 说我到最后一个 batch 了
return gathered[: self.gradient_state.remainder] # 按记录的尾部样本数裁掉补齐部分
return gathered
5. 真实实现
5.1 决策入口:prepare_data_loader
Accelerator.prepare_data_loader(src/accelerate/accelerator.py:2674)把自身配置(num_processes、split_batches、dispatch_batches、rng_types 等)转交给模块级函数 prepare_data_loader(src/accelerate/data_loader.py:1016)。后者的第一个决策是 dispatch_batches 的默认值(src/accelerate/data_loader.py:1113-1117):没显式指定时,put_on_device=False 就关掉 dispatch;否则仅当数据集是 IterableDataset 时默认开启。
5.2 采样器分片:BatchSamplerShard
src/accelerate/data_loader.py:110。两个关键实现:
__len__(src/accelerate/data_loader.py:170-186):split_batches 时长度不变;否则整除时直接除,drop_last时舍尾,even_batches时统一+1(尾部补齐),都不满足时按 process_index 决定谁多走一步。_iter_with_no_split(src/accelerate/data_loader.py:213-271):按idx % num_processes == process_index选 batch;尾部不整除且even_batches=True时,把开头攒下的initial_data自我复制后拼进最后一个 batch 再 yield(src/accelerate/data_loader.py:235-271)——这就是「借样本」的具体代码。_iter_with_split(src/accelerate/data_loader.py:191)则是按切片区间batch[batch_length*rank : batch_length*(rank+1)]取。
5.3 透明代理:DataLoaderAdapter
src/accelerate/data_loader.py:416。两条妙计:
__getattr__把未知属性全部委托给内部的base_dataloader(src/accelerate/data_loader.py:446-451);__class__被定义为 property,返回base_dataloader.__class__(src/accelerate/data_loader.py:458-468)——于是isinstance(prepared, DataLoader)为真,下游库的类型检查全部蒙混过关。副作用是 pickle 被破坏,所以DataLoaderShard专门实现了__reduce__(src/accelerate/data_loader.py:613-620)。
5.4 迭代器:DataLoaderShard
src/accelerate/data_loader.py:510。__iter__(src/accelerate/data_loader.py:577-610)的节奏:
- 先
synchronize_rng_states对齐各进程 RNG(src/accelerate/data_loader.py:578-579),保证 shuffling 一致; set_epoch(self.iteration)级联调用 batch_sampler / sampler / dataset 的 set_epoch(src/accelerate/data_loader.py:622-638),新 epoch 重洗;- 永远预取一个 batch:拿到 current 后立即
next()拿 next,current 先send_to_device再上yield(src/accelerate/data_loader.py:583-609)。这样 StopIteration 在 yield 最后一个 batch 之前就被发现,end_of_dataloader=True与remainder能被准确标记(经DataLoaderStateMixin,src/accelerate/data_loader.py:373)——下游的gather_for_metrics全靠这个标记。
5.5 广播分片:DataLoaderDispatcher
src/accelerate/data_loader.py:723。核心在 _fetch_batches(src/accelerate/data_loader.py:806-868):
- 只有进程 0 真正迭代;
split_batches=False时一次取num_processes个 batch 做concatenate(尺寸不齐在此报错,src/accelerate/data_loader.py:842-849); - batch 的「结构信息」(嵌套结构 + dtype/shape)经
broadcast_object_list发给所有进程(src/accelerate/data_loader.py:859),非主进程据此initialize_tensors造出空壳张量(src/accelerate/data_loader.py:891-893); broadcast(batch, from_process=0)把数据发下去(src/accelerate/data_loader.py:895-896),各进程再用slice_fn(默认slice_tensors)取自己那片。
5.6 评估修尾:gather_for_metrics
src/accelerate/accelerator.py:3068。先 gather(非张量对象走 gather_object),再查 gradient_state.end_of_dataloader(src/accelerate/accelerator.py:3113):若 remainder > 0,把 gather 结果裁到前 remainder 个(src/accelerate/accelerator.py:3121-3131), 把 even_batches 补进来的重复样本去掉。remainder 与 end_of_dataloader 两个属性由 GradientState 转发自当前活跃 dataloader(src/accelerate/state.py:1291-1306)。
6. 坑
- even_batches 的静默复制会污染指标。默认开启的尾部补齐让评估集被重复采样;凡是跨进程聚合指标,必须用
gather_for_metrics而不是裸gather。 split_batches=True的整除约束。batch_size % num_processes != 0在BatchSamplerShard.__init__直接 raise(src/accelerate/data_loader.py:163-167)。- dispatch 模式要求等长 batch。变长 batch(如 NLP 动态 padding 到不同长度)在
concatenate处炸;要么关 dispatch,要么split_batches=True。 - 变长 batch sampler 与 even_batches 冲突。文件里明确警告:变长 batch 必须
even_batches=False(src/accelerate/data_loader.py:140-143),否则补齐逻辑按固定 batch_size 算会出错。 - RNG 同步依赖 rng_types 配置。自定义 sampler/数据集若在
set_epoch之外另有随机源,synchronize_rng_states管不到它,各进程数据仍会漂移。 - skip_batches 的恢复语义。
SkipBatchSampler/SkipDataLoader(src/accelerate/data_loader.py:1332)靠「跳过前 N 个 batch」实现断点续训,对无固定顺序的 IterableDataset 只是近似恢复,不是精确位点。