第 7 章 LLM训练优化器

第 7 章 训练基础

第 7 章 训练基础

学习目标

  • 掌握损失函数、优化器、学习率调度、正则化的工程选择;
  • 理解初始化、归一化、残差连接对深层 Transformer 训练的意义;
  • 理解批大小、梯度累积与训练稳定性的关系,会处理 loss spike;
  • 厘清预训练、继续预训练、微调、指令微调的边界;
  • 掌握 RLHF、DPO、RLAIF、Constitutional AI 的原理与工程取舍;
  • 理解蒸馏、自训练、持续学习、增量学习、联邦学习的定位。

7.1 训练的骨架:损失、优化器、学习率

第 1 章定义过:训练 = 在结构下找参数,使目标函数最小。本节展开这个「找」的工程细节。

7.1.1 损失函数

LLM 的预训练损失是交叉熵(Cross-Entropy)。对下一个 token 的预测分布 pθ(x<t)p_\theta(\cdot | x_{<t}) 与真实下一个 token xtx_t

LCE=t=1Slogpθ(xtx<t)\mathcal{L}_{\text{CE}} = -\sum_{t=1}^{S} \log p_\theta(x_t \mid x_{<t})

它惩罚「给正确 token 的概率不够高」。语言模型汇报的困惑度(Perplexity, PPL)就是交叉熵的指数形式:PPL=eL\text{PPL} = e^{\mathcal{L}}——PPL 为 10 意味着模型在每个位置「等效于在 10 个候选里均匀犹豫」(第 9 章评估展开)。

微调阶段的损失仍是交叉熵,只是只对回答部分计算(指令部分被 loss mask 掉,不给梯度)。这是个常被新手忽略的细节:

# SFT 的 loss mask:只对 assistant 回答算损失
labels = tokenizer(...).input_ids
for segment in prompt_segments:        # system + user 部分
    labels[segment.start:segment.end] = IGNORE_INDEX  # -100,不进损失
loss = CrossEntropy(logits, labels)    # 只对回答 token 求平均

7.1.2 优化器:为什么是 AdamW

从 SGD 到 AdamW 的演进是为了对付深度网络的地形:

优化器 核心机制 问题
SGD 沿负梯度走 各维度尺度差异大时震荡
Momentum 梯度指数平均,惯性 仍需精细调学习率
Adam 一阶/二阶矩自适应步长 权重衰减实现有偏(与自适应更新耦合)
AdamW 解耦权重衰减(Decoupled Weight Decay) LLM 事实标准

AdamW 为每个参数维护一阶矩 mtm_t(动量)和二阶矩 vtv_t(梯度尺度),更新规则:

mt=β1mt1+(1β1)gt,vt=β2vt1+(1β2)gt2m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t,\quad v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2

θt=θt1ηm^tv^t+ϵ\theta_t = \theta_{t-1} - \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon}

其中 η\eta 是学习率,m^t,v^t\hat{m}_t, \hat{v}_t 是偏差修正后的矩。第 4 章的「优化器状态 12N 字节」就是这两个矩 + FP32 主权重——每个参数各占 4 字节,共 12 字节。这也是 Adafactor(对 vtv_t 做低秩近似省显存)等变体存在的理由,但主流仍是 AdamW。

超参经验值:β1=0.9\beta_1 = 0.9β2=0.95\beta_2 = 0.95(LLM 惯例,比默认 0.999 更保守,防 loss spike),ϵ=108\epsilon = 10^{-8},权重衰减 0.1。

7.1.3 学习率调度

LLM 训练的标准调度是 Warmup + Cosine 衰减(或 WSD:Warmup-Stable-Decay):

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
flowchart LR
    A[Warmup
线性升到峰值 lr] --> B[Stable/Cosine 主训练段
lr 缓慢衰减] B --> C[退火段
lr 降到 ~0
喂最高质量数据] style C fill:#e8f8e8
  • Warmup(前 0.1-1% 步数):学习率从 0 线性升到峰值(如 3e-4)。初期 Adam 的矩估计不准,大步长会直接把模型推坏。
  • Cosine 衰减:从峰值按余弦曲线降到峰值的 10% 以下。末期小步长让参数「沉淀」进更平坦的极小值——这解释了第 6 章「把最好的数据放在学习率最低的阶段」:退火阶段的小步长 + 高质量数据 = 性能的最后冲刺(Llama 3 的退火实践)。

学习率与模型规模的负相关是经验规律:7B 用 3e-4,70B 用 1.5e-4,越大的模型越「脆」,步长要越小。

7.1.4 正则化

LLM 训练里的正则化比想象中少:主要就是权重衰减(0.1 左右)+ 数据本身的多样性。Dropout 在现代 LLM 预训练里基本弃用(第 3 章的 GQA、第 8 章的混合精度已经引入足够的不确定性),Label Smoothing 也少用——因为它直接损害 PPL 的可比性。

7.2 初始化、归一化、残差:深网可训练的三板斧

千层网络能训得动,靠三个结构性设计(第 3 章从结构角度讲过,这里从训练动力学角度):

  1. 残差连接:梯度可以走「高速通道」直达浅层,缓解梯度消失。现代 LLM 的深径(residual stream)视角:主干是恒等映射,每层只学「增量」。
  2. 归一化:LayerNorm/RMSNorm 把每层激活拉回稳定分布,防数值爆炸。Pre-Norm(归一化在子层之前)取代 Post-Norm,因为 Post-Norm 在深层会梯度衰减,Pre-Norm 让千层训练稳定(GPT-2 起的共识)。
  3. 初始化:权重按宽度缩放初始化(如 N(0,1/d)\mathcal{N}(0, 1/\sqrt{d}) 或scaled init),保证前向激活方差与反向梯度方差不爆炸/消失。残差分支的输出投影常乘 1/2L1/\sqrt{2L} 缩放(GPT-2 的做法),让深层的累积方差可控。

工程含义:当你的微调训练 loss 不降时,先查数据与学习率,再查数值稳定性(NaN、梯度范数爆炸);结构层面这三板斧在现代架构里基本不用动。

7.3 批大小、梯度累积与训练稳定性

7.3.1 全局批大小

「批大小」在分布式训练里指的是全局批大小(所有卡上的 micro-batch × 梯度累积步数 × 数据并行度)。关键约束来自第 4 章的算力账:Critical Batch Size——批太小则 GPU 利用率低(通信占比大),批太大则样本效率下降(同样 token 数的收益递减)。现代 LLM 预训练的全局批大小通常在 1M-16M token 量级(注意单位是 token 不是样本),且随训练进程逐步增大。

7.3.2 梯度累积

单卡装不下大 batch 时,把一个逻辑 batch 拆成多次前向反向、累积梯度后再更新:

# 梯度累积:等效 batch = micro_batch × accum_steps
for i, batch in enumerate(loader):
    loss = model(batch) / accum_steps        # 除以累积步数,平均梯度
    loss.backward()                          # 梯度累加进 .grad
    if (i + 1) % accum_steps == 0:
        optimizer.step()                     # 一次真正的参数更新
        optimizer.zero_grad()

它与混合精度(第 8 章)配合时要注意:梯度累积的两次 forward 之间,BatchNorm 统计会变——LLM 用 LayerNorm/RMSNorm 所以无此问题,但如果你在微调里引入了 BN 层(罕见),要小心。

7.3.3 训练稳定性:Loss Spike

大模型训练最凶险的工程问题是 loss spike(损失突然飙升后可能恢复也可能发散)。已知诱因:数据批次异常(乱码、超长重复)、学习率过大、混合精度的数值溢出、硬件故障导致的梯度异常。业界公开的应对(Llama 3、Qwen 技术报告都有描述):

  1. 监控梯度范数(grad norm):spike 前常有先兆;
  2. 跳批:发现坏批次直接跳过(数据质量实时检测);
  3. 回滚:从 spike 前的 checkpoint 重启,跳过问题数据并降低学习率;
  4. BF16 优先:数值范围大(第 8 章),比 FP16 稳得多。

这条经验把第 4、6 章串起来了:训练稳定性 = 数据质量 × 数值精度 × 优化器超参,三者都会以 loss spike 的形式爆发。

7.4 训练谱系:预训练、继续预训练、微调、指令微调

现在把「训练」这个动词按生命周期排开:

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
flowchart TD
    A[随机初始化] -->|万亿 token 自监督| B[基座模型
Base Model] B -->|领域语料 CPT| C[领域基座] C -->|指令数据 SFT| D[对话模型] B -->|指令数据 SFT| D D -->|偏好数据 RLHF/DPO| E[对齐模型
产品级] style B fill:#e8f4f8 style E fill:#e8f8e8
阶段 数据 成本 改变什么
预训练(PT) 万亿 token 原始语料 万卡月级 世界知识与语言能力
继续预训练(CPT) 百亿-千亿 token 领域语料 千卡天级 领域知识注入
微调(FT/SFT) 万-百万条指令数据 单机-集群天级 行为格式与指令跟随
对齐(RLHF/DPO) 万-百万条偏好数据 与 SFT 同级 有用/无害/可靠的平衡

从基座到产品的典型顺序是 PT →(CPT)→ SFT → 对齐。各阶段解耦的价值:领域适配时不动 PT 权重(太贵),对齐时不动 SFT 之前的权重(避免破坏能力)。

7.4.1 灾难性遗忘与继续预训练

CPT 的头号敌人是灾难性遗忘(Catastrophic Forgetting):领域数据训多了,通用能力掉。缓解手段:

  • 数据回放(Replay):领域语料里混入通用语料(如 7:3);
  • 小学习率 + 少 epoch:CPT 学习率通常比 PT 低一个量级;
  • 参数高效微调(第 8 章 LoRA):冻结主干,天生防遗忘;
  • 模型合并(Model Merging):领域 LoRA 与通用权重按比例插值。

7.5 LLM 对齐:RLHF、DPO、RLAIF、Constitutional AI

对齐(Alignment)让模型从「会说话」变成「说该说的话」。第 3 章给了概念,这里展开技术路线。

7.5.1 RLHF:三段式

RLHF(Reinforcement Learning from Human Feedback)是 ChatGPT 的功臣,流程三段:

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
flowchart LR
    A[SFT 模型] --> B[奖励模型 RM]
    B -->|打分| C[RL 优化
PPO] C -->|KL 约束| A
  1. SFT:先得到会指令跟随的初始策略;
  2. 奖励模型(RM):用偏好对(A 优于 B)训练打分器。损失是 Bradley-Terry 成对排序损失:

LRM=logσ(r(x,yw)r(x,yl))\mathcal{L}_{\text{RM}} = -\log \sigma\big(r(x, y_w) - r(x, y_l)\big)

其中 ywy_w 是更好的回答(win),yly_l 是更差的(lose)。RM 学的是「相对好坏」而非绝对分数。

  1. RL 优化(PPO):策略模型生成回答,RM 打分作为奖励,最大化期望奖励。关键约束是KL 惩罚——不许策略偏离 SFT 初始分布太远,否则模型会「reward hacking」(学会讨好 RM 的怪癖,比如无脑变长、堆砌好词):

maxθ ExD,yπθ[r(x,y)]βKL(πθπSFT)\max_\theta\ \mathbb{E}_{x \sim \mathcal{D},\, y \sim \pi_\theta}\big[r(x, y)\big] - \beta\, \mathrm{KL}\big(\pi_\theta \,\|\, \pi_{\text{SFT}}\big)

7.5.2 DPO:去掉 RM 的直通车

PPO 的工程痛点:四个模型同时在显存里(策略、参考、RM、价值网络)、训练不稳、超参敏感。DPO(Direct Preference Optimization) 用数学变换证明:上面的 RL 目标有闭式解,可以直接用偏好数据优化策略,不需要显式 RM 和 RL 循环

LDPO=logσ(βlogπθ(ywx)πref(ywx)βlogπθ(ylx)πref(ylx))\mathcal{L}_{\text{DPO}} = -\log \sigma\left(\beta \log \frac{\pi_\theta(y_w | x)}{\pi_{\text{ref}}(y_w | x)} - \beta \log \frac{\pi_\theta(y_l | x)}{\pi_{\text{ref}}(y_l | x)}\right)

直觉:提高「胜过参考模型的程度」在赢方与输方之间的差距。只需策略 + 参考两个模型,监督学习式的稳定性。

维度 RLHF/PPO DPO
模型数 4(策略/参考/RM/价值) 2(策略/参考)
稳定性 难调、易 reward hacking 类监督学习,稳
上限 理论更高(在线探索) 偏静态偏好
在线 vs 离线 在线(边生成边学) 离线(固定数据集)
工程门槛 低,开源社区标配

实践格局:闭源大厂主打 PPO 及其变体(在线性带来上限优势),开源社区与多数业务场景用 DPO/变体(IPO、KTO、SimPO 持续涌现)。选择逻辑与全书主题一致:不是谁更先进,是谁与你的数据规模、工程能力、稳定性要求匹配。

7.5.3 RLAIF 与 Constitutional AI

RLHF 的标注瓶颈(第 6 章)催生 RLAIF:用强模型代替人类标偏好。Constitutional AI(Anthropic)是完整方法论:给模型一部「宪法」(一组原则),让模型先按原则自我批评(critique)、再自我修正(revise),用修正后的回答构造偏好数据做 RLAIF。价值:标注成本降一个量级 + 对齐标准显式可审计(宪法文本就是治理文档)——第 23 章会回到这个治理视角。

7.6 蒸馏、自训练、持续学习、增量学习、联邦学习

五个「训练的变体」,各自解决生命周期不同环节的问题:

知识蒸馏(Knowledge Distillation):教师模型(大而强)指导学生(小而快)。两种信号:

  • 软标签:教师的输出概率分布(含类间关系信息,比硬标签信息量大)——白盒蒸馏;
  • 生成数据:教师回答问题,学生做 SFT——黑盒蒸馏(没有教师权重也能做,闭源 API 可当教师)。

这是第 4 章「小模型价值」的主力工艺:端侧模型的行业标准路径。

自训练(Self-Training):模型给自己生成伪标签再自我提升(STaR:模型生成推理链,答对的保留回训)。与蒸馏的区别:教师就是自己(或自己的更强采样)。风险同样是错误累积。

持续学习(Continual Learning):时间维度上的扩展——模型随新数据不断更新而不忘旧能力。与 CPT 的灾难性遗忘斗争是同一战场,研究热点是参数隔离(给新任务分配专属参数)与回放策略。

增量学习(Incremental Learning):工程语境常指「新数据到达时只训增量」,比如新类别的检测头。在 LLM 语境下退化成「继续预训练 + 回放」。

联邦学习(Federated Learning):数据不出本地,各方训练本地模型、只上传梯度/参数更新到中心聚合。价值是隐私(医疗、金融、端侧输入法预测)。LLM 时代的现实:全参联邦不现实(通信量=参数量,见第 4 章),可行的是联邦微调 LoRA(只传低秩增量,通信量降 99%)。第 20、23 章的隐私部署会再遇到它。


本章要点回顾

  1. 预训练损失是交叉熵;SFT 的关键是 loss mask(只对回答算损失);PPL = e^损失。
  2. AdamW 是事实标准;β₂=0.95 是 LLM 惯例;「优化器状态 12N 字节」就是它的两个矩 + FP32 主权重。
  3. 学习率调度三段式 Warmup-Stable-Decay;「最好的数据放在最低学习率阶段」是退火的本质。
  4. 深网可训练三板斧:残差、Pre-Norm、缩放初始化;训练不稳步排查顺序:数据 → 学习率 → 数值。
  5. 训练谱系 PT → CPT → SFT → 对齐各司其职;CPT 的敌人是灾难性遗忘,靠回放与小学习率压制。
  6. RLHF = RM + PPO + KL 约束;DPO 用闭式解去掉 RM,工程友好;RLAIF/CAI 用模型反馈规模化对齐。
  7. 蒸馏(软标签/生成数据)、自训练(STaR)、持续学习(回放)、联邦学习(LoRA 增量上传)各解决生命周期不同环节的问题。

习题

  1. 手推:为什么 SFT 里把 prompt 部分 label 设为 -100 后,交叉熵只在回答部分平均?如果不禁用 prompt 损失,模型会学成什么样?
  2. Adam 的 β₂ 从 0.999 改成 0.95,对训练稳定性的影响机制是什么?什么场景下必须改回 0.999?
  3. 你的 CPT 模型领域评测涨了 8 分,通用评测掉了 12 分。给出至少三条带优先级的修复动作。
  4. DPO 训练中 β(KL 强度)调大调小分别会发生什么?分别对应什么业务偏好?
  5. 为什么联邦学习全参更新在 LLM 上不可行,而 LoRA 联邦可行?用第 4 章的通信量公式算一遍 7B 模型的差距。

延伸阅读

  • Loshchilov & Hutter, Decoupled Weight Decay Regularization(AdamW), 2019
  • Ouyang et al., Training Language Models to Follow Instructions with Human Feedback(InstructGPT/RLHF), 2022
  • Rafailov et al., Direct Preference Optimization: Your Language Model is Secretly a Reward Model(DPO), 2023
  • Bai et al., Constitutional AI: Harmlessness from AI Feedback, 2022
  • Hinton et al., Distilling the Knowledge in a Neural Network, 2015(蒸馏开山)
  • Zelikman et al., STaR: Bootstrapping Reasoning With Reasoning(自训练), 2022
  • Zhang et al., Why Should You Trust Your RLHF Policy? A Survey of Trustworthy RLHF(对齐风险综述), 2024
  • Grattafiori et al., The Llama 3 Herd of Models, 2024(退火与 loss spike 实践)

下一章预告

第 8 章聚焦「高效」:混合精度(FP16/BF16/FP8)、梯度检查点、ZeRO/FSDP 显存分片,以及 LoRA/QLoRA 参数高效微调——在单卡和集群上把训练成本打到可承受的范围。