跳到主要内容

数据截至 (上游 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_loadersrc/accelerate/accelerator.py:2674)把自身配置(num_processes、split_batches、dispatch_batches、rng_types 等)转交给模块级函数 prepare_data_loadersrc/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_splitsrc/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_splitsrc/accelerate/data_loader.py:191)则是按切片区间 batch[batch_length*rank : batch_length*(rank+1)] 取。

5.3 透明代理:DataLoaderAdapter

src/accelerate/data_loader.py:416。两条妙计:

  • __getattr__ 把未知属性全部委托给内部的 base_dataloadersrc/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)的节奏:

  1. synchronize_rng_states 对齐各进程 RNG(src/accelerate/data_loader.py:578-579),保证 shuffling 一致;
  2. set_epoch(self.iteration) 级联调用 batch_sampler / sampler / dataset 的 set_epoch(src/accelerate/data_loader.py:622-638),新 epoch 重洗;
  3. 永远预取一个 batch:拿到 current 后立即 next() 拿 next,current 先 send_to_device 再上 yieldsrc/accelerate/data_loader.py:583-609)。这样 StopIteration 在 yield 最后一个 batch 之前就被发现,end_of_dataloader=Trueremainder 能被准确标记(经 DataLoaderStateMixinsrc/accelerate/data_loader.py:373)——下游的 gather_for_metrics 全靠这个标记。

5.5 广播分片:DataLoaderDispatcher

src/accelerate/data_loader.py:723。核心在 _fetch_batchessrc/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_dataloadersrc/accelerate/accelerator.py:3113):若 remainder > 0,把 gather 结果裁到前 remainder 个(src/accelerate/accelerator.py:3121-3131),把 even_batches 补进来的重复样本去掉。remainderend_of_dataloader 两个属性由 GradientState 转发自当前活跃 dataloader(src/accelerate/state.py:1291-1306)。

6. 坑

  1. even_batches 的静默复制会污染指标。默认开启的尾部补齐让评估集被重复采样;凡是跨进程聚合指标,必须用 gather_for_metrics 而不是裸 gather
  2. split_batches=True 的整除约束batch_size % num_processes != 0BatchSamplerShard.__init__ 直接 raise(src/accelerate/data_loader.py:163-167)。
  3. dispatch 模式要求等长 batch。变长 batch(如 NLP 动态 padding 到不同长度)在 concatenate 处炸;要么关 dispatch,要么 split_batches=True
  4. 变长 batch sampler 与 even_batches 冲突。文件里明确警告:变长 batch 必须 even_batches=Falsesrc/accelerate/data_loader.py:140-143),否则补齐逻辑按固定 batch_size 算会出错。
  5. RNG 同步依赖 rng_types 配置。自定义 sampler/数据集若在 set_epoch 之外另有随机源,synchronize_rng_states 管不到它,各进程数据仍会漂移。
  6. skip_batches 的恢复语义SkipBatchSampler/SkipDataLoadersrc/accelerate/data_loader.py:1332)靠「跳过前 N 个 batch」实现断点续训,对无固定顺序的 IterableDataset 只是近似恢复,不是精确位点。