跳到主要内容

数据截至 (上游 commit cd6a3406572e)

03 · IterableDataset 流式读取

这一章讲什么: 当语料大到不想(或不能)下载全量时的另一条路——streaming=True 拿到的 IterableDataset。读完你会明白它内部那棵「迭代器树」怎么组合、shuffle 为什么用固定大小的 buffer 就能洗开 TB 级语料、以及断点续训和分布式切分是怎么实现的。


1. 它要解决的小问题

Map 式 Dataset 的前提是先有完整的本地 Arrow 文件。三个场景下这个前提不成立:

  • 太大:1 TB 语料,磁盘不够或不想等下载。
  • 用不完:只打算扫前 5% 做实验。
  • 训不完:训练本身就是单遍流式,epochs 的概念都弱化了。

想要的是:像拧水龙头一样取数据——远程文件边下边读,任何时候停下来都不亏。


2. 思路/直觉

把「取数据」建成一棵惰性树:

  • 叶子节点是「数据源」:逐个文件、逐行地产出 (key, example);
  • 每种变换是一层包装节点:MappedExamplesIterable 包一层做 map,FilteredExamplesIterable 包一层做 filter,BufferShuffledExamplesIterable 包一层做 shuffle……
  • 只有 for example in ds 真正迭代时,数据才从树根一路被"拉"出来。

这与 map 式 Dataset 的核心差异:没有任何中间结果落盘,因此也没有指纹缓存——每次迭代都重新流过整条管道。好处是零磁盘占用、零等待;代价是没有随机访问。

shuffle 怎么办?不能全量 permutation(那要先把所有数据拿到手)。答案是三级近似:

① shuffle_data_sources 打乱「先读哪个 shard」
② 交错填充 同时从最多 10 个 shard 轮流取,喂给 buffer
③ BufferShuffled 固定大小 buffer:随机吐一个、补进一个新的

buffer 只需几千条的内存,就把读取顺序洗成了近似随机。


3. 图示:一棵典型的迭代器树

怎么读这张图: 自下而上是数据流向(从文件到训练循环);每个框是一层节点,括号里是真实类名。

训练循环 for example in ds

│ 逐条 yield (example)
┌───────┴────────────────┐
│ 格式化 (FormattedExamplesIterable) │ set_format("torch") 时才有
├────────────────────────┤
│ 洗牌 (BufferShuffledExamplesIterable) │ buffer_size 条的内存 buffer
├────────────────────────┤
│ 交错 (CyclingMultiSourcesExamplesIterable) │ 最多 10 个 shard 轮流取
├────────────────────────┤
│ 映射 (MappedExamplesIterable) │ ds.map(fn) 加的
├────────────────────────┤
│ 叶子:数据源 (ExamplesIterable) │ 逐文件逐行地产出
└────────────────────────┘

│ HTTP range 请求 / 本地逐行
远程或本地的原始文件

对应到代码:每个节点都是 _BaseExamplesIterable 的子类(src/datasets/iterable_dataset.py:202-258),接口就五个——__iter__shuffle_data_sourcesshard_data_sourcesreshard_data_sources_init_state_dict / load_state_dict


4. 原理演示

# 示意,非源码
import random

class BufferShuffled:
def __init__(self, source, buffer_size, rng):
self.source, self.buffer_size, self.rng = source, buffer_size, rng

def __iter__(self):
buf = []
for x in self.source: # 从下一层拉数据
if len(buf) < self.buffer_size:
buf.append(x) # 先填满 buffer
else:
i = self.rng.randrange(self.buffer_size)
yield buf[i] # 随机吐一个
buf[i] = x # 新元素补进刚吐掉的位置
self.rng.shuffle(buf) # 收尾:buffer 里剩的全洗掉
yield from buf

重点看:内存里始终只有 buffer_size 条;第 N 条输入有机会替换 buffer 里任何位置,所以语料开头的数据不会永远堆在开头被先读走——这就是「近似全局打乱」。


5. 真实实现

5.1 叶子:ExamplesIterable

最基础的叶子是 ExamplesIterable(src/datasets/iterable_dataset.py:289-329):包一个用户给的 generate_examples_fn 和它的 kwargs。__iter__(:311-321)就三层循环——遍历 gen_kwargs(每个 shard 一份)、调用生成器、逐条 yield;如果 self._state_dict 已初始化,每 yield 一条就把 shard_example_idx 加一——状态跟踪就埋在迭代循环里,为断点续训服务。

shuffle_data_sources(:323-329)不碰数据,只对 kwargs 列表做一次 _shuffle_gen_kwargs——「打乱 shard 顺序」在叶子层就是打乱生成参数的顺序。

5.2 shuffle 的三级实现

IterableDataset.shuffle(src/datasets/iterable_dataset.py:3738-3834)完整代码就二十几行,一级一级看:

  1. 打乱 shard 顺序:ex_iterable.shuffle_data_sources(generator)(:3811-3812)。如果之前调过 skip/take(顺序已被冻结),SkipExamplesIterable.shuffle_data_sources 会抛 DataSourcesShufflingDisallowed(:2083-2086),shuffle 捕获后把 max_buffer_input_shards 降为 1(:3813-3814)。
  2. 交错填充:若 max_buffer_input_shards > 1,把数据源切成最多 10 份,用 CyclingMultiSourcesExamplesIterable 轮流取(:3817-3825)——不同 shard 的数据在 buffer 里混在一起,洗牌质量显著提升。
  3. buffer 打乱:最外层套 BufferShuffledExamplesIterable(:3826)。

BufferShuffledExamplesIterable.__iter__(src/datasets/iterable_dataset.py:1948-1963)与 §4 的示意几乎逐行对应,两个实现细节:随机索引用 _iter_random_indices 按 1000 个一批预生成(:1943-1946),减少 RNG 调用次数;buffer 用尽前先把剩余元素 rng.shuffle 再吐完。

5.3 多源混合:interleave 的概率采样

interleave_datasets([ds1, ds2], probabilities=[0.7, 0.3])(src/datasets/combine.py:18-165)在流式侧落到 RandomlyCyclingMultiSourcesExamplesIterable(src/datasets/iterable_dataset.py:1206-1328)。它的 _get_indices_iterator(:1238-1265)是个无限迭代器:按给定概率 rng.choice 批量预生成 1000 个「下一个从哪个源取」的决策,逐条消耗。

stopping_strategy 三选一,决定混合数据集何时结束:

策略语义注意
first_exhausted(默认)任一源耗尽即停(欠采样)小源决定了总长
all_exhausted每个源都至少被读完一遍(过采样,源会重开)总长可达 max_len × n_sources,见 combine.py:66-68 的警告
all_exhausted_without_replacement每个样本只被采一次配合 shuffle 内部用

5.4 断点续训:每个节点都记账

IterableDataset.state_dict()(src/datasets/iterable_dataset.py:2560)返回的是整棵树的嵌套状态:每层节点在 _init_state_dict 里声明自己记什么。比如:

  • 叶子 ExamplesIterableshard_idx / shard_example_idx(:307-309);
  • MappedExamplesIterable 记上一层的checkpoint + 之后又吐了几条(:1412-1420);
  • RandomlyCyclingMultiSourcesExamplesIterable 记 RNG 的完整 bit_generator 状态 + 批量决策游标(:1267-1277)。

恢复时 load_state_dict(_BaseExamplesIterable.load_state_dict,:260-273)递归地把保存的状态盖回活树上。例外是 shuffle buffer:buffer 的内容不保存——BufferShuffledExamplesIterable.load_state_dict 检测到状态不一致就打条警告,重新填满 buffer 再接着吐(:1934-1941)。所以续训后洗牌顺序会变,但「读过多少条」是对的。

5.5 分布式与 DataLoader worker:切分在「数据源」这一层

_prepare_ex_iterable_for_iteration(src/datasets/iterable_dataset.py:2764-2821)是每次迭代前的装配点,按顺序处理三件事:

  1. epoch:若设过 set_epoch(n),用 epoch 做种子重洗数据源并把树里所有 RNG 平移(shift_ex_examples_rngs,:2769-2771)——每个 epoch 顺序不同但可复现。
  2. 多机分布式(self._distributed):shard 数能整除 world_size 时,每个 rank 领 num_shards/world_size 个 shard(shard_data_sources,:2776-2783);不能整除就退化为 StepExamplesIterable——每个 rank 读全部 shard、只 yield 第 rankrank+world_size…条(:2794),日志里明说这不划算(:2785-2793)。
  3. 格式化:需要时套 RebatchedArrowExamplesIterable / FormattedExamplesIterable(:2796-2811)。

PyTorch DataLoader 的 num_workers_iter_pytorch(src/datasets/iterable_dataset.py:2695-2751):每个 worker 进程按 worker_info.id 领走一组 shard(split_shard_indices_by_worker + shard_data_sources,:2716-2726)。worker 数超过 shard 数时,多余 worker 直接空转并打警告(:2705-2714)——shard(文件)数是流式并行的硬上限,所以官方建议数据拆成足够多的文件,或 reshard() 按 Parquet row group 重切。

5.6 从哪来:streaming 模式的 builder 旁路

load_dataset(..., streaming=True)load.py:1717-1718 直接走 builder_instance.as_streaming_dataset()(src/datasets/builder.py:1113-1129):用 StreamingDownloadManager(fsspec 打开远程文件)调 _split_generators,拿到生成器后包成 ExamplesIterable——完全绕过 Arrow 落盘和指纹缓存,这就是第 2 章机制在流式侧一概不适用的原因。


6. 关键细节/坑

  • 没有随机访问、没有可靠长度。 len(ds) 对 IterableDataset 通常不可用;ds[i] 不存在。要抽样用 take(n)
  • shuffle 是近似的。 buffer 越小局部性越强;语料级训练够用,评测/复现要注意。docstring 明说「perfect shuffling 需要 buffer ≥ 全数据集大小」(src/datasets/iterable_dataset.py:3748-3750)。
  • skip/take 会冻结 shard 顺序,之后 shuffle 退化为单 shard 填 buffer(§5.2 第 1 步)。先 shuffle 后 take,别反过来。
  • worker 数 > shard 数 = 白起的进程。 文件太少时 reshard()(对 Parquet 按 row group 切)能救(shuffle docstring 里的示例,src/datasets/iterable_dataset.py:3795-3803)。
  • 续训恢复的是"进度"不是"洗牌内容"(§5.4),对洗牌顺序敏感的任务要注意。
  • fsspec 异步锁在 worker 里要重置。 _iter_pytorch 开头就 fsspec.asyn.reset_lock()(:2698-2699),注释指向 gcsfs 的挂死 issue——自己封装流式数据源时这是个已知的坑。

7. 代码地图

主题文件路径符号名
节点基类src/datasets/iterable_dataset.py_BaseExamplesIterable
叶子数据源src/datasets/iterable_dataset.pyExamplesIterableArrowExamplesIterable
map/filter 节点src/datasets/iterable_dataset.pyMappedExamplesIterableFilteredExamplesIterable
buffer 洗牌src/datasets/iterable_dataset.pyBufferShuffledExamplesIterable_iter_random_indices
多源交错src/datasets/iterable_dataset.pyCyclingMultiSourcesExamplesIterableRandomlyCyclingMultiSourcesExamplesIterable
skip/take/步进src/datasets/iterable_dataset.pySkipExamplesIterableTakeExamplesIterableStepExamplesIterable
数据集门面src/datasets/iterable_dataset.pyIterableDatasetshuffleset_epochstate_dictload_state_dict
迭代前装配src/datasets/iterable_dataset.py_prepare_ex_iterable_for_iteration_iter_pytorchshift_ex_examples_rngs
多源混合入口src/datasets/combine.pyinterleave_datasets
流式构建旁路src/datasets/builder.pyDatasetBuilder.as_streaming_dataset