← 返回AI变现
🤖 AIGC变现

LLM 训练全景:Pre-train、SFT、RLHF、DPO 与蒸馏

来源:掘金 · 发布于 2026-08-14 21:25:01
以工程视角串联现代 LLM 四阶段训练栈:预训练、中训、SFT 与对齐。覆盖数据、Tokenizer、优化器、精度、Scaling Law 与代表性训练框架,方便对照自己的训练与微调管线。

LLM 训练全景:Pre-train、SFT、RLHF、DPO 与蒸馏

ltl 2026-08-14 25 阅读25分钟

原文发布于 quant67.com,转载请保留出处。

一、为什么要把训练当作一条"流水线"而不是一个训练脚本

很多同学第一次读 torchrun --nproc_per_node=8 train.py 的时候,会产生一种错觉:大模型训练=加大 batch、加大模型、加大数据、跑几个月。真实世界里完全不是这样。一个规模化 LLM 项目的训练栈更像一条炼油厂流水线:

  • 上游是数据工程:抓取、清洗、去重、分类、配比、打包、打 tokenizer 的 tokens。
  • 中游是预训练(Pre-train):几千张 GPU 跑几十天,输出一个 base 模型 checkpoint。
  • 再接中训(Mid-train / Continued-PT):在 base 上注入数学、代码、推理、长上下文等"强化口味"的数据。
  • 下游是SFT(Supervised Fine-tuning):把模型从"补全文本"调成"能听懂指令"。
  • 再经过对齐(Alignment):RLHF、DPO、GRPO、RLAIF 等,把"会听懂"变成"有用、无害、诚实"。
  • 旁路还有蒸馏(Distillation):从大模型蒸出小模型,从推理模型蒸出非推理模型。

这一篇不去钻 3D 并行的细节(那是第 06 篇),也不展开 RLHF 的具体算法(第 09 篇),而是帮你建立一张整体地图:知道训练里每个环节在干什么、卡点在哪、业界主流选型是什么。如果你是团队里新加入的训练工程师,读完这一篇,至少在项目例会里能接得上大家的黑话。

二、四阶段工程栈:从 Pre-train 到 Alignment

2.1 总览

现代大模型(以 2024–2025 年 DeepSeek-V3、Qwen2.5、LLaMA-3.1、Kimi K1.5 的公开披露为参考)典型训练栈是四段:

flowchart TB
  RAW["原始文本 / 代码 / 多模态"] --> PT["Pre-train:数 T ~ 十几 T tokens;causal LM loss;几十天"]
  PT -->|base model| MID["Mid-train / Continued-PT:数百 B ~ 数 T tokens;数学 / 代码 / 推理加权;长上下文扩展"]
  MID -->|enhanced base| SFT["SFT:数十万 ~ 数百万 instruction pairs;学格式 + 学行为"]
  SFT -->|instruct model| ALIGN["Alignment:偏好数据;RLHF / DPO / GRPO / RLAIF / KTO"]
  ALIGN -->|aligned model| DISTILL["Distillation(旁路):teacher → student;推理蒸馏(o1 / R1 范式)"]

不同公司在这四段里投入的算力占比差别很大。粗略规律:

  • Pre-train 吃掉 90%+ 的 GPU-hour。
  • Mid-train 5% 左右,但对下游能力上限至关重要。
  • SFT 通常不到 1%。
  • Alignment 的算力不高,但工程复杂度最高(需要 reward model、PPO actor/critic、online rollout、数据反馈闭环)。

2.2 Pre-train:让模型"见过世界"

预训练阶段目标很朴素:给定一段文本前缀,预测下一个 token(causal LM)。Loss 是标准交叉熵。

工程上的核心挑战不是算法,而是:

  1. 数据够不够、干不干净:一次 13T tokens 的预训练,数据出错一轮,损失的是几千万美金。
  2. 训练稳不稳:loss spike、NaN、grad norm 爆炸,需要 checkpoint + 回滚 + 跳 batch。
  3. 吞吐能不能打满:3D 并行 + 通信重叠 + FP8,MFU(Model FLOPs Utilization)从 30% 抠到 50%+。
  4. 故障能不能容忍:千卡训练几十天,单卡 MTBF 几千小时,意味着每天都有卡挂。

2.3 Mid-train / Continued-PT:把"通才"拉向"硬核"

Mid-train 是 2024 年以后越来越标准化的阶段。在 base 快训完时,调整数据配比,显著加权数学、代码、STEM、推理类数据,同时往往把上下文长度从 4K/8K 扩到 32K/128K/1M。

  • DeepSeek-V3 在后期阶段把上下文从 4K 扩到 32K 再到 128K,配合 YaRN 类方法。
  • Qwen2.5 在 Continued-PT 阶段使用了更大比例的代码/数学数据,base 模型 MATH/HumanEval 分数大幅上升。
  • LLaMA-3 也有类似的 "annealing" 阶段:降低学习率、换数据配比、刷高质量数据。

这一阶段的工程意义是:在不重新花一遍预训练钱的前提下,用 5%~10% 的额外算力,拿到显著的能力跃升。

2.4 SFT:教模型"听人话"

SFT 用指令-回答对(instruction pairs)做监督学习,loss 仅在 response 部分计算(prompt mask 掉)。典型规模:

  • 早年 Alpaca / Vicuna:几万~几十万条;
  • 当下头部开源模型:几百万条,且多轮、多任务、多领域。

SFT 的工程重点:

  • 数据质量 >> 数据数量:一条 GPT-4 生成的高质量答案,胜过十条人工糙活。
  • 多轮对话拼接:loss mask 只打在 assistant turn 上,system/user turn 不算 loss。
  • 长样本打包(packing):把多条短样本拼到一个序列里,但用 attention mask 隔离,以榨干显存利用率。

2.5 Alignment:让模型"对得上人"

对齐阶段有一堆算法,工程上常见:

  • RLHF(PPO):经典三件套——SFT 模型、reward model、PPO actor+critic。复杂、吃显存、对超参敏感,但上限最高。
  • DPO(Direct Preference Optimization):不用 RL,直接在偏好对 (chosen, rejected) 上做对比损失。训练稳定、成本低,是开源社区主力。
  • GRPO(Group Relative Policy Optimization):DeepSeek 提出,去掉 critic,用一组 rollout 的相对 reward 做优势估计,在推理模型训练(R1)中大放异彩。
  • RLAIF:用更强的 LLM 做 reward,替代人工标注,便宜但有偏差。
  • KTO、IPO、SimPO、ORPO:DPO 变体,各家在稳定性与性能上微调。

第 09 篇会展开 RLHF 流水线,这里只需记住:Alignment 的硬件需求小于 pretrain,但工程链路最长——它闭环连接数据、模型、评测、灰度与线上反馈。

2.6 蒸馏:旁路的重要组件

蒸馏有两类:

  1. 能力蒸馏:大模型生成响应,小模型模仿。典型例子是 DeepSeek-R1 把 671B MoE 的推理能力蒸到 7B/14B/32B 的 Qwen/LLaMA 稠密模型。
  2. 行为蒸馏:让"推理模型"把思考链(CoT)蒸给"非推理模型",得到成本可控的 production 模型。

工程上,蒸馏通常复用 SFT 的代码路径,区别是数据来源从"人工/GPT-4"变成"teacher model 在线生成"。

2.7 四阶段的算力 / 数据 / 工程复杂度速查

阶段典型 tokens算力占比主要瓶颈典型 wall-clock
Pre-train数 T ~ 15T90%+3D 并行 + 故障容忍几周 ~ 几个月
Mid-train数百 B ~ 2T3%~8%数据配比 + 上下文扩展几天 ~ 几周
SFT10M ~ 1B<1%数据质量 + packing几小时 ~ 几天
Alignment100M ~ 数 B(含 rollout)1%~5%rollout 吞吐 + reward 稳定性几天 ~ 几周
Distillation10B ~ 数百 B视蒸馏深度teacher 推理吞吐几天

这张表的意义在于预算规划:如果老板问"我给你 1000 张 H100 两个月,能不能训一版模型",你至少知道时间主要花在哪,改哪个阶段能腾出空间做实验。

三、数据工程:训练成败的上限

3.1 数据源生态

  • Common Crawl:互联网抓取的原始 HTML,数 PB,脏乱差但"量大管饱",是所有预训练的基石。
  • C4(Colossal Clean Crawled Corpus):Google T5 清洗过的 Common Crawl 子集。
  • RedPajama:开源社区复现 LLaMA-1 数据配方的 1.2T tokens 数据集。
  • The Pile:EleutherAI 出品,825GB,多来源(书籍、代码、论文、Stack Exchange 等)。
  • 书籍:Books3(已因版权争议被下架)、Project Gutenberg、Anna's Archive 等,争议持续存在。
  • 代码:The Stack(Hugging Face + ServiceNow),v2 有 900+ 语言、近 70TB。GitHub 数据是代码能力核心。
  • 多语种:CC-100、mC4、OSCAR;中文专项有 WuDaoCorpora、SkyPile-150B、MAP-CC 等。
  • 学术/问答:arXiv、PubMed、Stack Exchange、Wikipedia。
  • 合成数据:2024 年之后强势崛起——用 GPT-4 / Claude / DeepSeek 生成的高质量 QA、数学、代码是 Qwen、Phi、DeepSeek 等公开承认使用的资源。

国内公开数据集:

  • WuDaoCorpora:智源,中文 5TB。
  • SkyPile-150B:昆仑万维。
  • MAP-CC:开源中文语料联盟。
  • CCI / CCI3:智源+上海 AI Lab 的中文互联网清洗集。

3.2 去重:MinHash 与 SimHash

去重对 loss 与泛化的影响被反复验证。LLaMA、RefinedWeb、Dolma 都把激进去重写进了配方。

两大类技术:

  • MinHash + LSH:对每个文档用 shingles → MinHash 签名 → LSH 分桶找近似重复。LLaMA、RedPajama、Dolma 都用它。阈值一般设 Jaccard ≥ 0.8。
  • SimHash:Google 网页去重经典,签名短、速度快,但召回略低于 MinHash。
# datasketch 做 MinHash LSH 的最小示例
from datasketch import MinHash, MinHashLSH

def minhash(text, num_perm=128):
    m = MinHash(num_perm=num_perm)
    for shingle in (text[i:i+5] for i in range(len(text)-4)):
        m.update(shingle.encode("utf-8"))
    return m

lsh = MinHashLSH(threshold=0.8, num_perm=128)
for doc_id, text in docs:
    lsh.insert(doc_id, minhash(text))

# 查询近似重复
dups = lsh.query(minhash(new_text))

生产里一般不会直接用 datasketch,而是 Spark/Ray + GPU 加速的流水线(如 NVIDIA NeMo Curator、DataComp-LM 工具链)。

3.3 质量过滤与毒性过滤

典型流水线层次(从粗到细):

  1. 语言识别:fastText、CLD3。
  2. 启发式规则:行长度、标点比例、重复率、HTML 残留、关键词黑名单(Gopher rules、C4 rules 都开源)。
  3. 分类器过滤:FastText 训练的质量分类器(以 Wikipedia、书籍为正样本,CC 为负样本),或 perplexity filter(用小 LM 过滤)。
  4. 毒性/NSFW 过滤:Perspective API、自研分类器;性暴力、仇恨言论、PII 打分。
  5. PII 脱敏:邮箱、手机号、身份证号、信用卡号用正则 + NER 匹配替换。
  6. 近似去重:上节的 MinHash/SimHash。
  7. 基准污染过滤(decontamination):扫描训练集里是否混入了 MMLU、GSM8K、HumanEval 等评测题目,必须清除,否则评测分是"偷来的"。

3.4 数据配比:几家公开配方

模型披露 / 推测的配比(摘要)
LLaMA-1CC 67%、C4 15%、GitHub 4.5%、Wikipedia 4.5%、Books 4.5%、arXiv 2.5%、Stack Exchange 2%
LLaMA-3未公开具体比例,但披露"代码占比显著提高、多语种 5%"、总量 15T tokens
DeepSeek-V314.8T tokens,中英双语为主,代码/数学比重高于 V2;FP8 训练
Qwen2.518T tokens,强化代码、数学、多语种;长文本阶段 1M 上下文
Mistral / Mixtral未公开
Phi-3强调"textbook quality"合成数据 + 精选网页

Mid-train 的配比变化是能力突跳的关键:把代码+数学+推理从 20% 拉到 40%+,STEM 基准线性上涨,但会牺牲一些通识问答上的分布。

3.5 数据打包与 tokens 计数

预训练前,数据要被"打包"成固定长度的序列(如 4096 或 8192)。两种做法:

  • Document concat:多篇文档用 <eos> 拼接后切片。训练效率高,但跨文档的 attention 可能引入噪声。
  • In-sample packing with attn mask:拼接但用 block-diagonal attention mask 隔离文档,保证因果但避免跨文档污染。现代框架(Megatron-LM、DeepSpeed、Axolotl)基本都支持。

3.6 数据质量评估的几个实用维度

光看 token 数不够,下面几个维度是头部团队实际在盯的:

  • 有效 tokens(effective tokens):去重 + 过滤后的净 tokens,不是原始抓取量。
  • 语言/领域分布:CC 天然偏英文和新闻,刻意补中文、代码、数学、STEM、长文本。
  • 文档长度分布:过短(< 128 tokens)和过长(> 64K)都要特别处理。
  • perplexity 分布:用一个小 LM 打分,剔除极高(乱码)和极低(重复模板)。
  • 重复 n-gram 率:整体 6-gram 重复率 < 某阈值是常见准入条件。
  • 毒性 / 偏见评分分布:防止后续 alignment 需要花很大力气"洗"。
  • 合成数据占比:2024 年后一个新监控点,过高会放大模型自身的幻觉。

把这些指标做成每批数据的"data card",与训练 ckpt 一起归档,是可审计训练流程的基础。

四、Tokenizer:经常被低估的关键组件

4.1 主流算法

  • BPE(Byte Pair Encoding):GPT-2、LLaMA、Mistral、Qwen、DeepSeek 都在用。从字节/字符出发,贪心合并频率最高的 pair。
  • WordPiece:BERT 系列经典,和 BPE 类似但合并准则是似然而非频率。
  • SentencePiece:Google 的实现,支持 BPE 与 Unigram,直接吃原始字节流(无需预分词),对多语种友好。
  • Unigram LM:SentencePiece 的另一模式,LLaMA tokenizer 之前也用过。
  • Tiktoken:OpenAI 的高性能 BPE 实现(Rust),被 GPT-3.5/4/4o 使用。

4.2 词表大小的权衡

词表优势劣势
小(32K,LLaMA-1/2)embedding 小,学习充分中文、代码切碎,序列变长
中(64K~128K,LLaMA-3、Qwen2、DeepSeek-V3)多语种友好,序列短embedding 显存变大
大(200K+,GPT-4o)极致多语种与符号覆盖embedding 参数暴涨

序列长度与词表的关系是乘法关系:词表翻倍,平均 tokens 数显著降低,训练和推理吞吐直接受益。DeepSeek-V3 的 tokenizer 词表扩到 128K,核心动机之一就是把中文压缩率提上去。

4.3 中文与 Unicode 的坑

  • BPE 要基于字节(byte-level BPE),而非 Unicode 字符,否则 emoji、罕见字会 OOV。GPT-2 以后都是 byte-level。
  • 中文预分词:SentencePiece 不需要,Tiktoken 用正则切分,会把中文切成单字或字+符号,对中文模型并不理想。
  • 组合 emoji、ZWJ 序列:测试 tokenizer 一定要覆盖。
  • 数字处理:LLaMA-1 的 tokenizer 把数字单独拆成 digit,对数学推理更友好;GPT-4 的 Tiktoken 在 o200k 词表里也调整了数字拆分。
# 用 tokenizers 库训练一个 byte-level BPE 的最小骨架
from tokenizers import Tokenizer, models, trainers, pre_tokenizers, decoders

tok = Tokenizer(models.BPE(unk_token="<unk>"))
tok.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
tok.decoder = decoders.ByteLevel()

trainer = trainers.BpeTrainer(
    vocab_size=128_000,
    special_tokens=["<|endoftext|>", "<|im_start|>", "<|im_end|>"],
    initial_alphabet=pre_tokenizers.ByteLevel.alphabet(),
)
tok.train(files=["corpus/*.txt"], trainer=trainer)
tok.save("tokenizer.json")

五、训练目标:不止 Causal LM

5.1 Causal LM

主流 decoder-only 模型的 loss:

L=−1T∑tlog⁡P(xt∣x<t)L = -\frac{1}{T} \sum_{t} \log P(x_t \mid x_{<t})L=−T1​t∑​logP(xt​∣x<t​)

实现上就是把 input_ids 右移一位作为 labels,用交叉熵。

5.2 Masked LM

BERT 系列,15% 随机 mask 预测原 token。现在很少用于大模型预训练,但在 embedding 模型、检索模型、编码器上仍然是主力(BGE、E5、GTE 等)。

5.3 MoE 路由损失

MoE(Mixture of Experts)模型除了主 loss,还有负载均衡损失(load balancing loss)和router z-loss。DeepSeek-V3 在这一块提出了 Auxiliary-Loss-Free Load Balancing:不再靠额外 loss 强推 expert 均衡,而是在 gating 时对每个 expert 加一个动态 bias,运行时根据实际负载调整。这避免了辅助 loss 对主 loss 的扰动,是 V3 训练稳定的一个重要原因。

5.4 Multi-Token Prediction(MTP)

DeepSeek-V3 还引入了 MTP:每一步不仅预测下一个 token,还预测未来 k 个 token。这给了三重好处:

  1. 训练信号更稠密,数据利用率提高;
  2. 推理时可以作为**推测解码(speculative decoding)**的 draft,提升解码吞吐;
  3. 对长距离依赖有轻微正则作用。

MTP 的 loss:

Ltotal=LCE(t+1)+λ1LCE(t+2)+λ2LCE(t+3)+⋯L_{\text{total}} = L_{\text{CE}}(t+1) + \lambda_1 L_{\text{CE}}(t+2) + \lambda_2 L_{\text{CE}}(t+3) + \cdotsLtotal​=LCE​(t+1)+λ1​LCE​(t+2)+λ2​LCE​(t+3)+⋯

第 15 篇会详细讲推测解码与 MTP 的推理侧应用。

5.5 几种训练目标的组合

现代 frontier 模型的 loss 很少只有一项。一个典型的 DeepSeek-V3 风格 loss:

L=LCE(next)⏟主 causal LM+λmtp∑kLCE(next+k)⏟Multi-Token Prediction+λz(logsumexp⁡(logits))2⏟router z-loss(MoE)L = \underbrace{L_{\text{CE}}(\text{next})}_{\text{主 causal LM}} + \underbrace{\lambda_{\text{mtp}} \sum_k L_{\text{CE}}(\text{next}+k)}_{\text{Multi-Token Prediction}} + \underbrace{\lambda_z \big(\operatorname{logsumexp}(\text{logits})\big)^2}_{\text{router z-loss(MoE)}}L=主 causal LMLCE​(next)​​+Multi-Token Predictionλmtp​k∑​LCE​(next+k)​​+router z-loss(MoE)λz​(logsumexp(logits))2​​

其中负载均衡走 aux-loss-free 方案(gating bias 动态调整),不进 loss。对 LLaMA-3 这样的稠密模型则简单很多,几乎只有主 CE loss。多项 loss 之间的权重 λ 是工程玄学,一般会在小规模(1B~7B)上扫一次,然后 scale up 沿用。

六、优化器与学习率调度

6.1 Adam / AdamW:老将仍在

大模型训练事实标准是 AdamW:Adam + decoupled weight decay。原因:

  • 自适应二阶动量对不同参数量级鲁棒;
  • decoupled weight decay 避免和动量混淆,泛化更好。

代价是显存 2x:每参数除了 FP32 master weight 外,还有 m、v 两个状态。

典型超参(沿用 GPT-3/LLaMA 配方):

AdamW(lr=peak_lr, betas=(0.9, 0.95), eps=1e-8, weight_decay=0.1)

beta2=0.95(而非默认 0.999)是大模型实践的共识,降低二阶动量的惯性可以避免长训练中的后期发散。

6.2 Lion:更省显存的黑马

Google 2023 年提出的 Lion 优化器只保留动量 mmm,用 sign-based 更新:

update=sign⁡(β1m+(1−β1)g)\text{update} = \operatorname{sign}\big(\beta_1 m + (1 - \beta_1) g\big)update=sign(β1​m+(1−β1​)g)

显存比 AdamW 省一半,在 ViT、语言模型上表现接近或更好。但对学习率和 weight decay 的调参窗口更窄,社区采用率不如 AdamW。

6.3 Muon:2024 的新秀

Muon(Keller Jordan 等)基于**矩阵正交化(Newton-Schulz 迭代)**对梯度做预处理,再施加动量更新。在 nanoGPT-speedrun 社区刷榜,Kimi K2 在公开技术报告中明确使用 Muon 作为预训练优化器,这是 Muon 在大规模生产里的重要背书。

工程上,Muon 只适合 2D 参数矩阵(全连接层、attention 投影),对 embedding、LayerNorm、1D bias 仍用 AdamW。这种"混合优化器"写法:

muon_params, adamw_params = [], []
for n, p in model.named_parameters():
    if p.ndim == 2 and "embed" not in n and "lm_head" not in n:
        muon_params.append(p)
    else:
        adamw_params.append(p)

opt_muon = Muon(muon_params, lr=0.02, momentum=0.95)
opt_adam = torch.optim.AdamW(adamw_params, lr=3e-4)

6.4 学习率调度

主流两种:

  • Warmup + Cosine Decay:前 1%~3% æ­¥ warmup 到峰值,然后 cosine 降到 10% 峰值。GPT-3、LLaMA、Qwen 等都用它。
  • WSD(Warmup-Stable-Decay):warmup → 长时间恒定 → 最后短衰减。MiniCPM、DeepSeek 等用过,好处是"decay 前可以当作 mid-train 的 base",继续训练或 annealing 都方便。
flowchart LR
  subgraph COS[&#34;Cosine 方案&#34;]
    C1[&#34;warmup 升到峰值&#34;] --> C2[&#34;cosine 衰减&#34;] --> C3[&#34;降到约 10% 峰值&#34;]
  end
  subgraph WSD[&#34;WSD 方案&#34;]
    W1[&#34;warmup 升到峰值&#34;] --> W2[&#34;stable:长时间恒定&#34;] --> W3[&#34;最后短衰减&#34;]
  end

6.5 WSD 为什么适合现代训练

WSD 的一个隐性好处是"stable 阶段的 checkpoint 可以当 base"。你可以:

  • 在 stable 阶段末尾存 ckpt,作为 continued-PT / mid-train 的起点;
  • 不同方向的 mid-train(代码强化 / 数学强化 / 多语强化)基于同一 stable ckpt 分支实验;
  • 最终 decay 阶段可以针对不同业务做多次独立 decay(用不同的数据配比),形成多个线上模型。

cosine scheduler 则要求你在训练开始时就决定好总步数,中途改步数会破坏几何性质,分支实验成本高。这也是为什么 MiniCPM、DeepSeek 之后越来越多团队切到 WSD。

七、精度:FP32 → BF16 → FP8 → FP6/FP4

精度演化直接决定训练成本:

精度典型硬件代表备注
FP32所有 GPU2017 以前稳,慢
FP16 混合精度Volta 以后GPT-3、早期 LLaMA需 loss scaling;动态范围小
BF16 混合精度A100/H100LLaMA-2/3、Qwen、GPT-4 训练主流动态范围大 = FP32,精度略低
FP8H100 HopperDeepSeek-V3 全流程 FP8、Llama-3 部分 FP8E4M3 / E5M2 两种格式
FP6 / FP4B200 Blackwell2025 年起推理主导,训练探索中

7.1 混合精度的基本结构

各张量的精度分工:

  • forward / backward 计算用 BF16(或 FP8);
  • 梯度 all-reduce 用 BF16;
  • 优化器状态(mmm、vvv)用 FP32;
  • master weight 用 FP32。

每一步的流程是:FP32 master weight 先 cast 成 BF16 权重参与 forward,反向得到 BF16 梯度,最后在 FP32 下做优化器更新写回 master。

7.2 FP8 训练的工程要点

DeepSeek-V3 是第一个在完整预训练里把 GEMM、通信、激活、梯度 大面积切到 FP8 的公开案例。关键技巧:

  • 每 block / 每 tile 动态 scaling:粗粒度 per-tensor scaling 范围太小,per-token 或 per-128-elements scaling 更稳。
  • 选择性回退:对 LayerNorm、softmax、优化器 state 保留 BF16/FP32。
  • 通信也用 FP8:all-to-all、all-reduce 的带宽压力减半。

FP8 训练的工程难点不是"能不能跑起来",而是"能不能全程无 spike 跑完 14T tokens"。

7.3 FP8 的两种格式:E4M3 vs E5M2

FP8 有两个 IEEE 近似变体:

  • E4M3:4 位指数 + 3 位尾数,精度高、动态范围小,用于前向激活和权重。
  • E5M2:5 位指数 + 2 位尾数,范围大、精度低,用于反向梯度(梯度的动态范围更广)。

NVIDIA Transformer Engine、Microsoft MS-AMP、DeepSeek 自研 FP8 kernel 都是基于这套分工。工程上最容易踩的坑是忘了给 gradient 用 E5M2,导致小梯度被 flush 到 0,训练后期 loss 不再下降。

7.4 精度选型的决策树

flowchart TB
  Q1{&#34;训练目标是 base 预训练?&#34;}
  Q1 -->|是| Q2{&#34;有 H100+ 和 FP8 工程能力?&#34;}
  Q1 -->|&#34;否(SFT / RLHF)&#34;| A3[&#34;BF16 即可,FP8 收益小风险大&#34;]
  Q2 -->|有| A1[&#34;FP8(DeepSeek 风格),省 30%+ 成本&#34;]
  Q2 -->|无| A2[&#34;BF16,稳妥首选&#34;]

SFT/RLHF 阶段样本量小、迭代快,FP8 带来的吞吐收益有限,反而 debug 成本高,一般不建议。

八、批大小、学习率与训练稳定性

8.1 批大小 scaling

大 batch 的好处是通信代价被摊薄;坏处是有效学习率变大,容易发散。

  • 线性缩放律:batch 翻倍,lr 翻倍(Goyal 2017,ImageNet)。
  • 平方根缩放律:理论更稳,但 LLM 社区多数仍沿用线性,只是给足 warmup。
  • Critical batch size:Kaplan 的论文给出"超过某个 batch,收益递减甚至负面"的临界点。大模型的临界 batch 在几百万到几千万 tokens 级别。

LLaMA-3、DeepSeek-V3 的全局 batch 常见于 4M~16M tokens。

8.2 损失 spike 的工程处理

长训练不可避免会遇到 loss spike。处理手段:

  1. 梯度裁剪(grad norm clip):clip 到 1.0 是事实标准。
  2. skip batch:遇到 NaN/Inf,丢弃当前 batch,回滚 optimizer 状态到上一步。
  3. checkpoint + 回滚:如果 spike 不可恢复,从之前的 checkpoint 载回,跳过问题数据区段。
  4. 数据怀疑论:spike 80% 的原因是数据(长重复、乱码、错误标注),应先看数据。
  5. embedding 归一化 / weight decay 微调:部分 spike 由 embedding 爆炸触发。
  6. 监控 z-loss、router entropy:MoE 模型专项。

8.3 一段值班日志的"标准 SOP"

假设你半夜收到告警:grad_norm > 20, loss increased by 1.5。标准排查:

  1. 看最近 200 step 的 loss / grad_norm / lr 曲线,确认不是 scheduler 阶段性变化。
  2. dump 当前 batch 的 input_ids,反 tokenize 看内容;统计 token 熵、重复 n-gram。
  3. 检查 NCCL / IB 是否有 retransmit、是否掉卡。
  4. 如果只是瞬时 spike 且后续 recover,不动;超过 3 个 batch 没 recover,从最近一个 ckpt rollback,跳过这段 data shards。
  5. 记录事件:时间、step、影响 tokens、措施、结论,进事故库。

头部团队的训练事故库动辄几百页——这是最值钱的工程资产。

九、3D 并行的组合策略(概览)

详细内容在第 06、07 篇。这里只给一张速记表:

并行方式切什么通信何时用
DP(Data Parallel)切 batchall-reduce grad永远用
TP(Tensor Parallel)切 weight 矩阵all-reduce activation单层太大装不下单卡时
PP(Pipeline Parallel)切 layerP2P send/recv模型层数很多、机间带宽不够时
SP(Sequence Parallel)切 seqall-gather / reduce-scatter长上下文训练
EP(Expert Parallel)切 MoE expertsall-to-allMoE 专用
ZeRO(1/2/3)切 optimizer / grad / paramreduce-scatter + all-gatherDP 的显存优化

千亿-万亿规模的典型组合(以 DeepSeek-V3 / Qwen2.5-72B / LLaMA-3-405B 为参考):

  • TP = 8(单机内,NVLink 带宽足)
  • PP = 8 ~ 16(跨机,容忍 IB 带宽)
  • EP = 8 ~ 64(MoE 专用)
  • DP / ZeRO:剩下的 GPU 数都分给 DP 维度
  • SP:长上下文时打开

9.1 一条朴素的选型经验

  • 单卡显存塞不下一层 → 开 TP。
  • 单机(8 卡 NVLink)塞不下一整份模型 → 开 PP 或 ZeRO-3。
  • 机间 IB 带宽紧张、PP bubble 不可接受 → 优先 ZeRO + 较大 micro-batch。
  • 序列超过 32K → 打开 SP / context parallel。
  • MoE 模型 → EP 先于其他维度考虑,all-to-all 是最贵的通信。

9.2 并行维度选择对通信量的直觉

定性上:

  • TP 通信量 = O(batch × seq × hidden),每一层都有,最怕跨机。
  • PP 通信量 = O(batch × seq × hidden),只有相邻 stage,不怕跨机但怕 bubble。
  • DP/ZeRO 通信量 = O(params),每步一次,带宽敏感、延迟不敏感。
  • EP 通信量 = O(batch × seq × hidden × topk),每层两次 all-to-all,对对称带宽和 crossbar 敏感。

所以"TP 放在 NVLink 域内,DP/ZeRO 放在 IB/RoCE 跨机域上"是黄金法则。

十、Scaling Laws:花钱怎么花最合算

10.1 Kaplan 2020

OpenAI 的 Kaplan 等人给出了第一个系统的 scaling law:loss 是模型参数 N、数据量 D、算力 C 的幂律函数。结论鼓舞人心——但它低估了数据的作用。

10.2 Chinchilla(Hoffmann 2022)

DeepMind 重新做实验后发现,Kaplan 对数据和模型的最优比例错了。最优比例大约是:

D∗≈20ND^* \approx 20 ND∗≈20N

即每个参数配 ~20 tokens。GPT-3(175B 参数、300B tokens)被严重"undertrained",而 Chinchilla(70B 参数、1.4T tokens)在同等算力下效果更好。

这之后,开源社区"小而多数据"成为共识:LLaMA-1 用 7B/13B/33B/65B + 1T~1.4T tokens,LLaMA-3 干脆把 8B/70B 训到 15T tokens——远超 Chinchilla 最优,因为推理时代,推理成本 >> 训练成本,把更多训练算力砸进去换来推理侧便宜是划算的。

10.3 推理时 scaling(o1 范式)

2024 年 OpenAI o1 带来新范式:推理时算力也是 scaling 维度。模型可以在推理时生成长 CoT、反思、回溯,用更多 tokens 换更高正确率。DeepSeek-R1、Kimi K1.5、Qwen QwQ 都跟进了这条路线。

对训练基础设施的影响:

  • RL 阶段需要大规模 online rollout,推理集群和训练集群边界变模糊。
  • 长 CoT 训练样本动辄几万 tokens,对长上下文训练压力大。
  • 奖励建模变成"可验证答案(数学、代码)优先"的新范式。

10.4 三代 scaling 的对照

范式ä