Skip to content

预训练 ​

预训练是 AI Infra 最极致的负载:​一次跑三个月,一步都不能错。

预训练在算什么 ​

预训练的目标朴素而昂贵:给互联网级别的文本,用下一个 token 预测任务把参数 Ψ 调到最优。它的负载特征由三个量决定:

  • 计算量​:C=6ND(见芯片架构篇),训练 GPT-3 级模型(N=175B、D=300B)约需 3.1×1023 FLOP;
  • 数据量​:现代预训练用 10∼20T token,原始网页数据在 PB 级;
  • 时间尺度​:数月连续运行,跨越数千次故障(见容错与 Checkpoint篇)。

Scaling Law:先算账再开机 ​

训练前最贵的决策是"模型多大、数据多少"。Chinchilla 定律给出的经验结论:​计算预算 C 给定时,最优配置满足 N 与 D 大致等比例增长​(约每增加 1 个参数配 20 个 token)。工程上的推演方式:

D≈20Ψ⇒最优 Ψ≈C/120, D≈20C/6

决定训练哪一档模型后,万卡集群篇的时间公式直接给出工期与成本。​预训练 Infra 的第一步不是写代码,是算这本账​——这也是"量化分析与系统设计"能力的第一个用武之地。

数据管道:GPU 断粮的预防针 ​

万亿 token 从 PB 级原始网页到训练样本,是一条多级流水线:

flowchart LR
    A["原始网页<br/>(PB 级)"] --> B["去重 / 清洗 / 过滤"]
    B --> C["分词 tokenize"]
    C --> D["分片打包<br/>(预写盘)"]
    D --> E["训练集群<br/>(流式读取)"]
    B -.质量过滤.-> F["丢弃 50%+"]

Infra 视角的两个关键点:

  1. 清洗决定质量上限,工程决定质量下限​。去重(MinHash)、质量分类器过滤会把数据量砍掉一半以上,但这些计算是一次性的 CPU/离线作业,可以用便宜资源慢慢跑。
  2. 训练侧只求稳定供给​:数据预分片打包成顺序读友好的格式(如 packed 序列),训练时流式读取,保证每步 I/O 时间远小于计算时间——否则上万张卡集体等数据,每秒都是烧钱。

训练循环与损失稳定性 ​

单个训练步:前向 → 反向 → 优化器更新,配合显存层次篇的 ZeRO/激活重算把状态装下、用并行策略把计算摊开。预训练特有的 Infra 问题是数值稳定性​:

  • loss spike​:训练数周后 loss 突然跳升,常见诱因是坏数据批次或数值上溢。一线做法:检测到 spike 后回滚到健康 checkpoint 并跳过可疑数据段(Llama 3、MegaScale 均内置);
  • 混合精度护栏​:BF16 计算天然稳于 FP16(无需 loss scaling),梯度裁剪兜底;
  • 学习率与 warmup​:前期 warmup 防止初期大梯度把参数打进饱和区。
深入推导:为什么是 6ND,注意力项什么时候不能忽略

一次前向中,每个参数恰好参与一次乘加(每 token):2Ψ FLOP/token。反向传播计算各参数梯度及输入梯度,约两倍于前向,故 C≈6ND。

注意力项​:序列长 s 时,注意力矩阵 QK⊤ 与 softmax⋅V 各需约 2Ls FLOP/token(L 为层数),合计约 4Ls。它与 2Ψ 之比决定能否忽略:

4Ls2Ψ=2LsΨ

以 7B 模型(L=32、Ψ=7×109):比值 ≈9×10−9×s,s=4096 时约 4%,可忽略;但 s=105(超长上下文)时超过 90%,​必须计入,且注意力从矩阵乘退化为平方复杂度​——这正是 Ring Attention 与序列并行要解决的问题。

(据 Kaplan et al. 2020 与 Hoffmann et al. 2022 的口径。)

思考题 ​

  1. 你有 1024 张 H100(MFU 按 45%),训练一个 34B 模型、8T token 需要多少天?
  2. Chinchilla 最优是每参数 20 token,但 Llama 3 405B 用了约 38 token/参数(15.6T/405B)。为什么一线团队普遍"过训练"?
  3. 数据管道中"去重"为什么同时提升质量和效率?
参考答案
  1. C=6×34×109×8×1012=1.63×1024;吞吐 =1024×989×1012×0.45=4.55×1017 FLOP/s;T≈3.6×106 s≈41 天。
  2. 推理成本与参数量挂钩,模型越大服务越贵。​推理生命周期内要被调用万亿次的模型,多训一倍的 token(训练成本 +33%)换取更小的参数量是划算的​——这是"推理感知的 scaling"(overtraining),Llama 3 系列是典型。
  3. 去重后有效 token 密度更高(同样的训练预算学到更多),且重复序列会诱发记忆化;同时数据量减少直接降低了后续所有训练 I/O。

小结 ​

  • 预训练由 6ND 与 scaling law 定盘:先算工期成本,再开机器。
  • 数据管道是一次性离线工程,训练侧只求流式稳定供给。
  • loss spike 的标准处置:回滚 + 跳过可疑数据;数值稳定性靠 BF16 与梯度裁剪兜底。
  • 长上下文时注意力项不能忽略,引出序列并行(后续)。

参考资料 ​

最近更新