L2.8 数据集采样、分片与课程学习
三维坐标
layer: L2(数据与算子层)|level: Engineer|pillar: 数据L2.1 讲数据如何进入管道;L3 训练章讲并行与显存。本章聚焦训练循环的数据入口:如何把 TB 级语料切成每个 GPU 每 step 的 batch,如何避免重复与热点,以及何时用课程学习提升收敛效率。这一层的 Bug 最阴险——它不报错、不 crash,只是让你花了同样的算力账单,却训出一个更差的模型。
学习目标
- 前置知识:读过 L2.1 数据管道、L3.1 数据并行概念;会 PyTorch
DataLoader。无需分布式训练经验。 - 学完产出:① 能区分 采样(sampling)、分片(sharding)、混洗(shuffle) 三者的职责边界;② 能设计 DP=N 时「每 rank 读不同 shard、全局无重复」的文件划分方案,并会算「shard 数不整除 worker 数」时的吞吐长尾与重复率;③ 能解释 Megatron sample-based 与 token-based 训练预算的差异;④ 能描述课程学习(curriculum learning)多阶段(易→难)的设计、阶段切换时 loss 曲线的预期形态与回滚策略;⑤ 亲手实现一个带
DistributedSampler的最小多进程数据加载 demo(CPU 即可)。 - 阅读姿势:数据供给层 Bug 的表现是 loss 曲线「看起来能训但泛化差 / 重复样本 / 各 rank 步数不一致」——常在数周后才被发现。带着一个问题读全篇:「我怎么证明每个样本恰好被训了它该被训的次数?」 采样、分片、resume、curriculum 的所有设计,最终都是在回答这一句。
背景与现状
大模型预训练数据流:
- 离线:PB 级 raw → 清洗 → 去重 → 按 token 长度分桶 → 写成 shard 文件(JSONL / Parquet / MDS)。
- 在线:每个 training step,各 GPU rank 从自己的 shard 采样一个 micro-batch,经 tokenizer → collate → forward。
三类工程事故:
| 事故 | 症状 | 根因 |
|---|---|---|
| 重复采样 | 有效 token 数虚高、过拟合 | shard 划分重叠或未去重 |
| rank 不同步 | NCCL hang / step 计数漂移 | 各 rank dataset 长度不一致 |
| 热点文档 | 少数 URL 占比异常 | 采样未加权或未 cap per-source |
业界信号:Megatron-LM、HuggingFace
datasets、LitData(Lightning)均提供 deterministic sharding + resume 语义;预训练团队把「数据 mix 配方」与超参同等保密。
📅 时效口径(2026 年初):主流开源模型的预训练 token 量已普遍进入 15T~36T 量级——Llama 3 系列公布约 15T,Qwen2.5 约 18T,2025 年发布的新一代开源模型(Llama 4、Qwen3 等)公布口径已达 30T+。在这个量级下,「数据几乎只过一个 epoch、靠 mix 配比而非重复来堆有效 token」成为主流;多阶段 curriculum(中后期退火阶段提高高质量/数学/代码占比)也已从各家秘方变成公开实践(Llama 3、MiniCPM、OLMo 2 等技术报告均明确描述了退火/多阶段配比)。本章讲的采样、分片与阶段切换,正是支撑这些实践的工程地基——具体数字随新模型发布会继续变化,但「一遍过 + 分阶段配比」的范式在可见的未来是稳定的。
原理与架构
2.1 采样 vs 分片 vs 混洗
- 分片:静态或动态把数据划给 rank/worker,保证并行无 overlap。它回答「谁读哪些文件」。
- 采样:从分片内按分布抽样本(均匀 / 加权 / 温度混合)。它回答「每个来源出现多少次」——加权采样(weighted sampling)本质是给每个数据源指定一个目标占比
w_i,训练器按w_i抽取,而不管该源的自然大小。 - 混洗:epoch 边界重排,避免顺序偏差。它回答「以什么顺序出现」。
三者职责正交:分片错 → 重复/遗漏;采样错 → mix 失衡(见思考题 1);混洗错 → 顺序偏差(如同域文档连续出现导致梯度短期偏移)。排查数据问题时先问「是哪一层的锅」,能省一半时间。
2.2 分布式分片模式
| 模式 | 做法 | 适用 |
|---|---|---|
| 文件级分片 | rank i 读 files[i::world_size] | 预训练 shard 巨大、顺序读 |
| 样本级分片 | DistributedSampler(dataset, shuffle=True) | 中小数据集、微调 |
| Iterable + 状态 | WebDataset / Megatron 内置 cursor | 无限流、需精确 resume |
文件级分片有一个容易被忽视的前提:shard 数最好是「总消费者数 = rank 数 × 每 rank worker 数」的整数倍。不整除时,有的 worker 分到的 shard 比别人多一个——所有人都要等最慢的那个,产生吞吐长尾;而若框架用「补齐重复(pad by repeat)」来对齐长度,则会静默引入重复样本(见思考题 2 的完整算账)。
Resume 关键:checkpoint 必须保存 data cursor(epoch、file offset、sample index、sampler RNG 状态),否则恢复训练重复或跳过数据。
2.3 Token 预算与 sample 预算
- Sample-based:固定 epoch 数 × 数据集大小。
- Token-based:训练直到消耗
N个 token(Chinchilla 法则给出「参数量 × 20 token」的下界量级;2026 年主流实践早已远超此值,做的是 over-training 换推理性价比)。
工程上 token-based 更常见——需在线统计 tokens_per_step × steps 并对齐 lr schedule。注意一个隐蔽耦合:token 预算 + 加权采样 = 每个源各自的等效 epoch 数不同。总预算固定时,小体量、高权重的源会被反复过很多遍——这正是 mix 配方要与去重、memorization 风险一起审计的原因(思考题 1 会把这笔账算出来)。
2.4 课程学习(Curriculum Learning)
按难度递增组织训练:
- 阶段 A:短序列、高质量子集(教科书 / 代码 / 维基精选)
- 阶段 B:混合中等长度、多样域
- 阶段 C:全长、含噪声域(论坛、用户生成)
收益:早期稳定收敛、减少 loss spike;代价:需维护难度标注或 proxy。常用难度 proxy 对比:
| 难度 proxy | 计算成本 | 优点 | 局限 |
|---|---|---|---|
| 序列长度 | 几乎为零 | 与显存/算力曲线天然耦合,warm-up 期用短序列还能提高吞吐 | 长 ≠ 难:长的模板化网页可能比短数学题简单 |
| 参考模型 perplexity | 需跑一遍小模型推理 | 与「模型视角的难度」最贴近 | 依赖参考模型质量;对参考模型没见过的域会系统性高估难度 |
| 质量分类器分数 | 训练 + 推理一个小分类器 | 可按「教育价值」等业务定义打分(FineWeb-Edu 式做法) | 分类器偏好会整体注入训练分布,需抽样人审校准 |
| 来源先验 | 零 | 简单可解释(教科书 < 论坛噪声) | 粒度粗,同源内部难度方差大 |
工程上通常组合使用:来源先验定大盘,分类器分数细筛,长度控制吞吐。
一个必须提前建立的预期:阶段切换点的 loss 上跳是数学必然,不是事故。loss 是「模型能力」与「数据分布难度」的合成读数——切到更难的分布,同一个模型的平均 loss 自然抬高。区分「良性台阶」与「真事故」的判据、以及切换失败时的回滚策略,见思考题 3。
2.5 两个真实业务症状复盘
数据供给层的事故几乎从不以「数据错误」的面目出现,而是伪装成训练问题。两个典型场景:
症状 一:resume 后 loss 突刺,三天后才发现是数据重复。
某 7B 续训任务在 step 42k 因节点故障重启,resume 后 loss 先小幅下降、随后在评测集上 perplexity 反而变差。排查 lr schedule、优化器状态均正常,最后 diff 数据侧才发现:checkpoint 只保存了 step,没保存 sampler 的 epoch 与 RNG 状态(2.2 节的 data cursor)。resume 后 DataLoader 从 epoch 开头重新迭代,前 42k step 见过的样本被原样重训一遍——loss 看起来「降得很快」(因为在背已见过的数据),实际是 memorization。修复方式:cursor(epoch、offset、RNG state)随 checkpoint 一起落盘,并在 resume 后用「前 100 个 batch 的样本 hash 与故障前日志比对」做验证。
症状二:加了新数据源后英文能力「无声退化」。 团队往 mix 里加入 20% 的中文语料并等比压缩其他源,两周后内部 benchmark 显示英文推理任务掉了 2 个点,但 loss 曲线全程平滑无异常。根因是采样权重改了,per-source token 计数 dashboard 没有跟着建——没有人能回答「英文 web 数据这两周实际被训了多少 token」。这类退化不报错、不突刺,只能靠「按源统计消耗 token + 定期分域评测」的审计闭环兜住(2.1 节:采样层回答「每个来源出现多少次」,就必须有对账机制证明它答对了)。
动手实践:DistributedSampler 与 shard 划分
# shard_demo.py — 单机模拟 4 rank 文件分片
import json
from pathlib import Path
shards = [f"shard_{i:03d}.jsonl" for i in range(16)]
world_size = 4
def files_for_rank(rank: int):
return shards[rank::world_size]
for rank in range(world_size):
assigned = files_for_rank(rank)
print(f"rank {rank}: {len(assigned)} files -> {assigned[:2]}...")
# rank 0: shard_000,004,008,012 — 无重叠
# distributed_sampler_demo.py — 需 torchrun 2 进程
import torch
import torch.distributed as dist
from torch.utils.data import DataLoader, TensorDataset
from torch.utils.data.distributed import DistributedSampler
def main():
dist.init_process_group("gloo")
rank = dist.get_rank()
world = dist.get_world_size()
data = TensorDataset(torch.arange(100))
sampler = DistributedSampler(data, num_replicas=world, rank=rank, shuffle=True)
loader = DataLoader(data, batch_size=8, sampler=sampler)
batch = next(iter(loader))
print(f"rank={rank} batch={batch[0].tolist()}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
torchrun --nproc_per_node=2 distributed_sampler_demo.py
# 预期:两 rank 打印的 index 无交集,合起来覆盖 0..99
课程学习配置示例(YAML 示意)
curriculum:
stages:
- name: warm-up
max_seq_len: 2048
data_mix: {high_quality: 1.0}
token_budget: 50B
- name: main
max_seq_len: 8192
data_mix: {high_quality: 0.5, web: 0.5}
token_budget: 500B
踩坑预警
DistributedSampler.set_epoch(epoch)忘记调用:每个 epoch shuffle 相同,泛化差。- drop_last=False 导致 rank 步数差 1:DP 训练需
drop_last=True或 padding dummy batch。 - 全局 batch 与 lr 线性缩放:改采样权重或 batch 组成未同步调 lr。
- 去重只做 document 级:段落级 duplicate 仍会导致 memorization。
- checkpoint 只存模型不存 data cursor:resume 后重复或跳过数据(2.5 节症状一),且
DistributedSampler的 shuffle 依赖set_epoch,RNG 状态也要一并落盘。 - shard 数与消费者数不整除:吞吐长尾或静默重复(思考题 2);离线切 shard 时就应按「rank × worker 的公倍数」规划 shard 数。
深入思考
下面三题每题先给题干,再用
<details>折叠一份图文并茂的参考答案。建议先合上答案自己想 3 分钟,再 展开对照。
思考题 1:加权采样的混比账——「配比差一点」如何变成「能力差一截」?
假设总预算 1000B token,语料自然体量为 web 800B、code 100B、books 50B。产品要求加强代码能力,于是把采样权重定为 web:code:books = 0.5:0.3:0.2。结合 2.1 节(采样的职责)与 2.3 节(token 预算与等效 epoch 的耦合),算一算:这套权重下每个源实际被训多少遍(等效 epoch)?哪个源有 memorization 风险、哪个源被训不够?如果不能改总预算,你有哪些工程手段把这套配比「安全地」落地?
展开参考答案(含混比-等效 epoch 流向图 + 算一遍)
结论:加权采样把「总预算 × 权重」硬性分给每个源,而各源自然体量差异巨大,于是同一套权重下不同源的等效 epoch 可以差出一个数量级——小体量高权重的源被反复背诵(memorization 与过拟合风险),大体量低权重的源连一遍都过不完(欠训练)。mix 配方的真正难点不在「定比例」,而在「验证每个源被训的遍数落在安全区间」,这需要 per-source token 计数与上采样上限(cap)机制兜底。
用具体数字算一遍(等效 epoch = 分得预算 ÷ 自然体量):
| 源 | 自然体量 | 权重 | 分得预算 | 等效 epoch | 判定 |
|---|---|---|---|---|---|
| web | 800B | 0.5 | 500B | 0.625 | 欠训练:37.5% 的 web 数据从未被见过 |
| code | 100B | 0.3 | 300B | 3.0 | 重复 3 遍:可接受偏上限(经验上高质量数据重复 2~4 遍尚可) |
| books | 50B | 0.2 | 200B | 4.0 | 重复 4 遍:长文本逐字 memorization 风险显著,且边际收益递减 |
对照自然配比(web 84.2%、code 10.5%、books 5.3%)可见:这套权重把 books 上采样了近 4 倍、code 近 3 倍。权重每偏离自然配比一分,就在「重复遍数」上付一分代价——而重复带来的收益是递减的:数据重复研究(如 Muennighoff 等的 scaling 分析)的公开结论是重复到 4 遍左右收益开始明显衰减,十几遍后接近白费算力。
不改总预算的安全落地手段:
- 设 per-source 重复上限(epoch cap):如 books 上限 2.5 epoch,超出的配额按剩余权重归还给 web——权重表达的是「意图」,cap 表达的是「底线」。
- 温度采样代替硬权重:按
p_i ∝ size_i^α(α 取 0.5~0.7)在自然配比与均匀配比之间插值,天然抑制小源被过度上采样。 - 扩充小源:给 code 补充合成数据/新抓取,把「靠重复凑配额」变成「靠新数据凑配额」。
- 建审计闭环:per-source token 计数 dashboard + 分域评测(2.5 节症状二的教训)——mix 是否被 faithfully 执行,必须有对账数字,不能靠配置文件自证。
这正是 2.1 节「采样回答每个来源出现多少 次」与 2.3 节「token 预算 + 权重 = 各源等效 epoch」两条线的交汇:配比是一等公民超参,它的验证成本不低于 lr。
思考题 2:100 个 shard 喂 32 个消费者——不整除的代价是长尾还是重复?
预训练数据被切成 100 个大小相近的 shard,训练拓扑是 8 个 DP rank、每 rank 4 个 dataloader worker(共 32 个消费者),采用 2.2 节的文件级分片。请算一算:shard 如何分到 32 个消费者头上?一个 epoch 的墙钟时间被谁决定?如果框架为了让各消费者步数对齐而「补齐重复」,会引入多少重复样本?结合 2.2 节(分片模式与整除前提),给出离线与在线两侧的修复方案。
展开参考答案(含分片不均衡示意图 + 算一遍)
结论:100 ÷ 32 不整除时,4 个消费者分到 4 个 shard、28 个消费者分到 3 个,同步训练的 epoch 墙钟时间由最慢者决定——集群整体吞吐利用率只有约 78%,等价于白扔两成算力;若框架改用补齐重复来对齐长度,则要把 100 个 shard 垫成 128 份,静默引入 28% 的重复样本。长尾与重复是同一枚硬币的两面,根治办法是在离线切数据时就让 shard 数是消费者数的公倍数,或改用样本级/流式分片把粒 度打细。
用具体数字算一遍:
- 分配:
100 = 32 × 3 + 4→ 前 4 个消费者拿 4 个 shard,其余 28 个拿 3 个。 - 长尾:同步 DP 下每个 step 都要 allreduce 对齐,epoch 墙钟时间由「4 shard 组」决定,慢速组比快速组多干
4/3 ≈ 1.33倍的活;整体吞吐利用率 = 平均负载 ÷ 最大负载 =(100/32) / 4 = 78.1%——买了 8 台机器,约 1.75 台在等数据。 - 重复代价:若框架选择把每个消费者的 shard 数补齐到 4(复用已有 shard 垫上),总消费份数变成
32 × 4 = 128,其中128 − 100 = 28份是重复——28% 的重复率,比长尾更糟:它不显示在任何监控上,直接污染有效 token 统计与 loss(对照 2.2 节 resume 一段:重复样本正是数据侧最难事后察觉的事故)。 - 更隐蔽的变体:worker 数从 4 改成 5(消费者 40 个,100/40 同样不整除)这类「只调了个 dataloader 参数」的改动,也会悄悄改变长尾/重复结构——分片健康度必须在每次拓扑变更后重新审计。
修复方案:
- 离线侧(首选):切 shard 时按「常用拓扑的公倍数」规划数量,如 128、256、512 个 shard——2.2 节表格中文件级分片的「整除前提」应写进数据构建规范,而不是训练现场再补救。
- 在线侧:改用样本级分片(
DistributedSampler直接按样本 index 轮转,粒度细到单样本,天然均衡);或用流式分片(WebDataset / Megatron cursor 模式,消费者从全局队列动态领取 shard,快者多劳,无长尾也无重复),代价是 resume 状态管理更复杂。 - 兜底监控:per-rank「本 epoch 已消费样本数」指标——长尾与重复在这个指标上都会现形。
思考题 3:curriculum 阶段切换后 loss 跳升 0.6——这是事故还是数学必然?
按 2.4 节三阶段方案训练:阶段 A 用高质量子集(模型在其上 loss 已收敛到 1.8),切到阶段 B 后混入 50% web 数据(模型在 web 上的 loss 约 3.0)。切换当步,loss 从 1.8 跳到约 2.4,值班同学按「loss spike 事故」流程回滚了 checkpoint。结合 2.4 节(阶段切换的预期形态),判断这次回滚是否正确;给出「良性台阶」与「真事故」的区分判据;并设计一套让切换更平滑的 ramp 方案与真正该有的回滚策略。
展开参考答案(含切换形态判别图 + 算一遍)
结论:这次回滚是误操作。loss 是模型能力与数据难度的合成读数,切到更难的混合分布,期望 loss 跳到各源 loss 的加权平均(0.5×1.8 + 0.5×3.0 = 2.4)是数学必然——它是「台阶」不是「尖刺」。判据看三点:跳变高度是否可由混合比例预测、跳变后是否立刻恢复下降趋势、梯度范数是否同步爆炸。真正的工程改进是把硬切换改成线性 ramp 摊平台阶,并把回滚预案绑定在「切换后 N step 内 loss 不降反升 / 梯度异常」的条件上,而不是绑定在「loss 变高」本身。
用具体数字算一遍(线性 ramp:4000 step 内把 web 占比从 0 升到 0.5,每 1000 step 提高 12.5%):
| step | 高质量占比 | web 占比 | 期望 loss(加权平均) | 相邻台阶高度 |
|---|---|---|---|---|
| 0(切换前) | 1.000 | 0.000 | 1.80 | — |
| 1000 | 0.875 | 0.125 | 1.95 | +0.15 |
| 2000 | 0.750 | 0.250 | 2.10 | +0.15 |
| 3000 | 0.625 | 0.375 | 2.25 | +0.15 |
| 4000 | 0.500 | 0.500 | 2.40 | +0.15 |
硬切换是一次 +0.6 的悬崖;ramp 把它摊成 4 级 +0.15 的缓坡——每一步的分布偏移变小,优化器动量与 lr schedule 有时间适应,监控上也更容易把「预期台阶」与「异常尖刺」区分开(实际训练中模型还在持续变强,真实曲线会比这张静态表更低、且各段内持续下降)。
区分「良性台阶」与「真事故」的操作化判据:
- 可预测性:切换前先用当前 checkpoint 在新 mix 上跑离线评测,得到预测 loss;实际跳变落在预测值 ±10% 内即属预期。
- 趋势:良性台阶在切换后 100~500 step 内恢复下降;真事故(数据污染、脏样本、lr 不匹配)表现为持续上行或反复尖刺。
- 梯度侧证:良性切换的 grad norm 只小幅抬升;若 grad norm 数量级爆炸或出现 NaN,才走事故流程。
真正该有的回滚策略:① 切换前强制落一个「阶段边界 checkpoint」(含 2.2 节的完整 data cursor,保证回滚后数据不重不漏);② 回滚触发条件写成规则——「切换后 500 step,loss 高于预测台阶 10% 以上,或 grad norm 超过历史 P99 的 3 倍」;③ 回滚后的重试动作不是「放弃切换」,而是放缓 ramp(拉长到 2 倍 step)或临时下调 lr再切。把这套规则写进值班手册,才能避免本题中「见 loss 变高就回滚」的误操作——那次回滚浪费的不只是重训的算力,还有阶段 B 本该注入的新分布信号(呼应 2.4 节:切换点的 loss 上跳是数学必然,不是事故)。
延伸阅读
1. 核心 Paper / 技术报告
- Scaling Data-Constrained Language Models(Muennighoff et al., 2023)— 数据重复多少遍开始白费算力,思考题 1 的定量依据。
- DoReMi: Optimizing Data Mixtures Speeds Up Language Model Pretraining(2023)— 用小模型代理自动学 mix 权重。
- Llama 3 / OLMo 2 / MiniCPM 技术报告 — 多阶段数据配比与退火(annealing)阶段的公开实践细节。
2. 相关高 Star 仓库与源码必读路径
NVIDIA/Megatron-LM— 看megatron/core/datasets/下BlendedDataset与 sample index 构建,理解工业级加权采样 + deterministic resume。mosaicml/streaming(MDS 格式)— 看 shard 与 canonical nodes 的映射逻辑,是思考题 2「整除规划」的工程化范本。webdataset/webdataset— 流式分片与split_by_node / split_by_worker的组合语义。huggingface/datatrove— FineWeb 系列的数据处理流水线,看质量分类器打分如何接入 curriculum 的难度 proxy。
3. 站内关联
下一篇 → 数据供给层到此闭环:你已经能证明「每个样本恰好被训了它该被训的次数」。接下来去 L3 训练章,看这些 batch 进入 GPU 之后,并行策略与显存如何接棒。