跳到主要内容

数据截至 (上游 commit 090253dac668)

02 · 数据管线:4 万亿 token 怎么喂

这一章讲什么: 从一堆 .json.gz 语料到 GPU 里的 batch,OLMo 数据管线的四层结构;以及「数据顺序可复现」这个 OLMo 最看重的性质是怎么实现的。


1. 它要解决的小问题

预训练数据工程有四个硬约束:

  1. 装不下。 4T token × uint32 ≈ 16TB,内存装不下,必须惰性读取。
  2. 要洗牌。 全局随机顺序对收敛至关重要,但 16TB 没法真洗。
  3. 要不重不漏。 几百个 rank × 几十个 worker 并行取数,每个实例恰好被消费一次。
  4. 要可恢复、可审计。 训练挂了重启,得从数据流的同一个点继续;事后要能回答「第 N 步到底读了什么」。

OLMo 的答案是一条四层流水线,每层解决一个问题:

离线制备 训练时
───────── ─────────────────────────────────
*.json.gz ──► ① .npy 文件群 ──► ② MemMapDataset ──► ③ IterableDataset ──► ④ DataLoader+Collator ──► batch
(tokenize, (磁盘/HTTP, (虚拟拼接成 (全局洗牌 + (padding、堆叠、
追加 EOS) 内存映射) 连续 token 流) 三级切片 + 多线程) 生成 mask)

①②解决「装不下」,③解决「洗牌+不重不漏+可恢复」,④解决「凑 batch」。


2. 第一层:离线制备——token 流落盘

scripts/prepare_memmap_dataset.py 把 gzipped JSONL(每行一个含 text 字段的文档)token 化,拼成一个大 numpy 数组写成 .npy关键约定:EOS token 在制备时就已经插进文档之间——MemMapDataset 的 docstring 明说「No special tokens are added to the input IDs so it's assumed that if you want EOS tokens between documents, those will already be in the memory-mapped array」(olmo/data/memmap_dataset.py:29-30)。

这意味着**「一个训练实例 = 4096 个连续 token」可以横跨多个文档**(packing)。文档边界不靠分隔结构表达,而是靠 EOS token 本身——配合 generate_doc_lengthsolmo/data/util.py:122get_document_lengths 按 EOS 位置算文档长度)可以进一步做文档内注意力掩码(见第 1 章的 cu_doc_lens)。

dtype 是 uint16uint32DataConfig.memmap_dtype 默认 uint16olmo/config.py:612),但 OLMo-2 词表超过 65535,配置里必须写 uint32configs/official-1124/OLMo2-7B-stage1.yaml:224)——选错 dtype 数据直接溢出,这是个真实的坑。


3. 第二层:MemMapDataset——把一堆文件虚拟拼接成一个数组

直觉: 每个 .npy 文件是一个 token 数组;把 N 个文件的长度前缀和记下来,就得到一个「全局 index → (文件, 文件内 offset)」的映射,对外表现为一个巨大的连续 token 数组。实例 = 连续 chunk_size(即 max_sequence_length)个 token。

长度发现是并行 + 远程友好的。 offsets 属性(olmo/data/memmap_dataset.py:102-155)用 ThreadPoolExecutor 并发地拿每个文件的长度——文件可能在本地磁盘,也可能是 http://olmo-data.org/... 或 S3 URL,读长度只发一个 HTTP HEAD / range 请求(olmo/util.py:331file_size)。

读一个实例只读需要的字节。 _read_chunk_from_memmapolmo/data/memmap_dataset.py:157-167):

# 摘自 olmo/data/memmap_dataset.py:159-166
item_size = dtype(0).itemsize
bytes_start = index * item_size * self._chunk_size
num_bytes = item_size * self._chunk_size
buffer = get_bytes_range(path, bytes_start, num_bytes)
array = np.frombuffer(buffer, dtype=dtype)

get_bytes_rangeolmo/util.py:368)对本地文件做 mmap 切片、对 URL 发 HTTP Range 请求——这就是为什么官方配置里的数据路径可以是一千多条 http URL:数据可以边训边流式拉取(README 也说明了这一点)。

示意图: 实例切分

文件 A (6 chunk) 文件 B (4 chunk) 文件 C (5 chunk)
┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄┄
全局实例 index: 0 1 2 3 4 5 | 6 7 8 9 | 10 11 12 13 14
offsets 表: (0,6) (6,10) (10,15)
└── index=8 → memmap_index=1, local=8-6=2

真实查找逻辑在 __getitem__olmo/data/memmap_dataset.py:179-196):线性扫 offsets 表定位文件,再 _read_chunk_from_memmap 按字节读。

一个实例还能带这些附加字段__getitem__ 后半段,:198-219):

字段触发条件用途
instance_mask配了 instance_filter周期性重复检测(见 §5)
label_mask配了 label_mask_paths哪些 token 计入 loss(SFT/annealing 用)
attention_maskgenerate_attention_mask屏蔽 padding
doc_lensgenerate_doc_lengths文档内注意力掩码
metadata默认实例来源(评测时按域分桶靠它)

4. 第三层:IterableDataset——可复现性的心脏

这是整个数据管线最值得抄的部分(olmo/data/iterable_dataset.py:20)。

4.1 全局洗牌:seed + epoch 决定一切

_build_global_indicesiterable_dataset.py:91-114):

# 摘自 olmo/data/iterable_dataset.py:93-98
indices = np.arange(len(self.dataset), dtype=np.uint32)
if self.shuffle:
# Deterministically shuffle based on epoch and seed
rng = np.random.Generator(np.random.PCG64(seed=self.seed + self.epoch))
rng.shuffle(indices)

要点:洗牌种子 = seed + epoch。只要 (seed, epoch) 相同,全世界任何机器算出的全局顺序逐位相同。注释还顺手吐槽了 torch 内置随机数不够随机(:96),所以用 numpy 的 PCG64。

洗牌后还要补齐/裁到 total_sizeworld_size 的整数倍,:100-113)——不整除时从头部循环补样本,保证每个 rank 拿到一样多。

4.2 洗牌结果落盘,全机共享

_build_and_save_global_indices:75-89)把洗牌后的索引写成 train_data/global_indices.npy 内存映射文件:每个节点的 fs_local_rank == 0 进程算一次写一次,其余进程 barrier 后直接 mmap 只读打开(get_global_indices:116-120)。几亿条索引不用每个进程重算,也不占内存。

4.3 三级保序切片

__iter__:127-182)把全局索引切成「本 worker 该读的子序列」,分三级:

全局索引 [洗牌后]
│ ① 截断到 max_examples / start_index(恢复时用)
│ ② 按 rank 抽: indices[rank :: world_size] (:141)

本 rank 索引
│ ③ 按 worker 切: reshape(-1, device_batch_size)
│ [worker_id :: num_workers] (:148-161)

本 worker 索引
│ ④ 按线程交错: indices[i::num_threads] + roundrobin (:168-180)

逐实例 yield

第 ③ 级有个容易写错的细节:worker 间不能简单地 indices[worker_id::num_workers] 交错切——那样一个 batch 的实例会被拆到不同 worker,DataLoader 按 round-robin 收集时就保不住顺序。OLMo 先把索引按 device_batch_size 一捆一捆 reshape,再整捆分给 worker(:154-160),注释里写明了理由。任何「rank/worker 数变了顺序也不变」的保证都靠这个 reshape。

4.4 可恢复性从哪来

  • 续跑start_index(= 已见的实例数)直接切掉已读前缀(:136-138),训练器恢复时设置它(olmo/train.py:403-406)。
  • 换 epochreshuffle(epoch):122-125)换种子重洗。
  • 审计:每个实例被读时都带着自己的全局 index(_get_dataset_item 塞进 index 字段,:184-191),训练器把它写进 TSV(见第 3 章)。

5. 第四层之外:实例级重复过滤

小问题: 网页数据里满是「关注公众号关注公众号关注公众号……」式的周期性垃圾,混进 batch 会污染梯度。

解法不是删数据(那会改变数据集长度,破坏可复现性),而是读出来打个掩码find_periodic_sequencesolmo/data/util.py:41)的做法很巧:

把序列 reshape 成 (n_rows, period) 的矩阵
→ 每行和上一行比,整行相等 = 可能是一个周期段
→ 连续的「整行相等」聚成段,首尾再按半行修正偏移
→ 重复次数 ≥ 3 才报(少于 3 次无法可靠判定周期)

MemMapDataset._validate_instanceolmo/data/memmap_dataset.py:236-247)对 period 1~13 各扫一遍(InstanceFilterConfig 默认值,olmo/config.py:601-605),命中且重复 ≥32 次的实例返回 instance_mask=False。训练侧在 get_labelsolmo/train.py:727-728)把被掩实例的 label 全部置 -100,loss 直接归零——数据还在、顺序不变,只是它不产生梯度。OLMo-2 7B stage1 配置里这个过滤是开着的(configs/official-1124/OLMo2-7B-stage1.yaml:226-229)。


6. 数据混合的真相:没有采样器,只有文件清单

很多框架的数据混合是「按概率在线采样」,OLMo 不是。混合比例 = 文件列表里各类文件的数量比build_memmap_datasetolmo/data/__init__.py:22-53)把 data.paths(或 data.datasets 字典)里所有路径平铺给 MemMapDataset,想上采样某类数据就把它的路径重复列几遍——olmo/data/named_data_mixes.py 顶部注释里 wiki 就明确写着「repeated twice to up-sample」并真的列了两次。

真实配方在两个地方:

  • configs/official-1124/OLMo2-7B-stage1.yaml:230-1366:stage1 的全部训练数据,一千多行 URL,分段的注释标着来源(ProofPile 2、pes2o、Starcoder、DCLM……)。
  • configs/official-1124/provenance.csv:README 指向的逐来源统计。

两阶段训练的数据切换也是配置级的:stage2 配置(configs/official-1124/OLMo2-7B-stage2-seed42.yaml)换成 Dolmino 高质量混合 + load_path 指向 stage1 终点 + restore_dataloader: false:75-77)——新数据、旧权重、从头计数。


7. 坑与边界

  • dtype 选错会静默溢出。 词表 >65535 必须 memmap_dtype: uint32;代码里 effective_memmap_dtypeolmo/config.py:627-635)只校验它是合法 numpy 类型,不校验够不够装词表。
  • IterableDataset 要求数据集长度 < 2^32iterable_dataset.py:92 的断言)——索引是 uint32。
  • worker 切片只在 num_workers>0 时发生;单进程时退化到多线程(默认猜 4 个线程,:162-165)。
  • build_train_dataloader 的 work_dir 保护olmo/data/__init__.py:148-155):train_data/ 目录已存在且没开 save_overwrite 就报错——防止两次 run 的全局索引互相覆盖。
  • eval 用另一套build_eval_dataloaderolmo/data/__init__.py:92)用 PyTorch 自带的 DistributedSampler,不走 IterableDataset——评测不需要跨 epoch 续跑语义。
  • 改动文件列表 = 改动整个数据顺序。想在 run 中途调配比而保持可复现,没有内建机制(inferred:配置即真相,中途改配置只能新开 run)。

8. 代码地图

主题文件路径符号名
训练 dataloader 装配olmo/data/__init__.py:127build_train_dataloader
memmap 数据集olmo/data/memmap_dataset.py:20MemMapDataset
字节级读取olmo/data/memmap_dataset.py:157olmo/util.py:368_read_chunk_from_memmapget_bytes_range
全局洗牌与切片olmo/data/iterable_dataset.py:91:127_build_global_indices__iter__
collatorolmo/data/collator.py:14DataCollator
周期重复检测olmo/data/util.py:41find_periodic_sequences
文档长度olmo/data/util.py:122get_document_lengths
数据制备scripts/prepare_memmap_dataset.py(jsonl.gz → .npy)
命名混合olmo/data/named_data_mixes.py:1DATA_PATHS
真实配方configs/official-1124/OLMo2-7B-stage1.yaml:216-1366(YAML data: 段)