目标:全面理解大语言模型从预训练到对齐微调的工程实践核心技术,掌握分布式训练、混合精度、显存优化等大规模训练的关键技术。

前置要求:了解 Transformer 架构、注意力机制、损失函数与优化器的基本概念。阅读过本系列第一篇会更有帮助,但不是必须的。

本系列第一篇从因果注意力讲到正则化,系统梳理了大语言模型背后的理论基础——架构设计、损失函数、优化器和泛化控制。然而,当我们将这些理论放到千亿参数、万亿 token 的训练规模时,会面临一系列全新的工程挑战:单张 GPU 装不下模型怎么办?如何让数千张 GPU 高效协作?训练过程中如何防止崩溃?

本文将围绕算力、显存、通信、数据质量四大瓶颈,系统讲解大规模训练的核心工程实践。


1. 预训练——在海量数据中学习语言的规律

1.1 预训练的目标

  大语言模型的预训练本质上是一个自回归语言建模任务:给定前面所有的 token,预测下一个 token。其损失函数就是我们熟悉的负对数似然

$$ \mathcal{L} = -\frac{1}{T}\sum_{t=1}^{T} \log p_\theta(x_t \mid x_1, x_2, \ldots, x_{t-1}) $$

这个目标简洁而强大。通过在海量文本上最小化这个损失,模型被迫学习语法、语义、世界知识、推理模式等一切隐含在文本中的规律。

1.2 数据来源与处理

  预训练数据的质量直接决定了模型的上限。常见的数据来源包括:

  • 网页数据:Common Crawl 等互联网爬虫数据,数量最大但质量参差不齐
  • 书籍与学术论文:高质量长文本,有助于学习长距离依赖
  • 代码:GitHub 等开源代码,提升模型的编程和逻辑推理能力
  • 维基百科与百科全书:结构化的知识密集型文本
  • 对话数据:Reddit、论坛等,有助于学习对话能力

数据清洗流水线是预训练中最被低估但最关键的一环:

  1. 去重(Deduplication):使用 MinHash 或 SimHash 等算法检测近似重复文档。去重的目的是防止模型在重复数据上过拟合,同时提高训练效率。
  2. 启发式过滤:移除过短、过长、语言混杂、包含大量特殊字符的文档。
  3. 质量分类器:训练一个二分类器(如 fastText),用高质量文本(如维基百科)作为正例,随机网页作为负例,过滤低质量内容。
  4. 有害内容过滤:移除色情、暴力、仇恨言论等有害内容。

类比:食材采购。预训练数据就像一家自助餐厅的食材。去重相当于检查是否重复采购了同一种食材;启发式过滤相当于扔掉腐烂变质的食材;质量分类器相当于厨师试吃,只留下好吃的。食材的质量决定了菜品的上限。

1.3 分词(Tokenization)

  在第一篇中我们介绍了 BPE。这里补充几种主流分词算法的对比:

算法代表模型核心思路特点
BPEGPT 系列贪心合并频率最高的字符对简单高效,确定性编码
WordPieceBERT基于似然最大化的合并合并标准更"智能"
SentencePieceLLaMA, T5直接在原始文本上训练,不依赖预分词多语言友好,不依赖空格
Byte-level BPEGPT-2 以后以字节(byte)为基本单位做 BPE无 OOV 问题,支持任意文本

词表大小的选择是一个权衡:

  • 词表太小:每个词被拆分成更多子词 token,序列变长,增加计算量和推理延迟。
  • 词表太大:嵌入矩阵($\mathbb{R}^{|V| \times d}$)参数量膨胀,低频 token 训练不充分。
  • 经验值:现代 LLM 通常使用 32K~128K 的词表大小。

1.4 预训练超参数

  预训练的关键超参数包括:

  • 批次大小(Batch Size):通常以 token 数衡量。GPT-3 使用约 3.2M tokens/batch。大批次有助于训练稳定但需要更多显存。
  • 序列长度(Sequence Length):通常 2048~8192。更长的序列需要更多显存(注意力的 $O(n^2)$ 复杂度)。
  • 学习率调度:Warmup + Cosine Decay 是标配,详见第一篇。
  • 权重衰减:通常 $\lambda = 0.1$,配合 AdamW 优化器。

1.5 Scaling Laws——扩展法则

  第一篇已介绍了 Scaling Laws 的基本概念。这里从工程实践的角度深入。

Kaplan 定律(2020):在给定计算预算 $C$ 下,最优策略是将大部分预算分配给更大的模型而非更多的数据。

Chinchilla 定律(2022):修正了 Kaplan 的结论——最优配比是参数量和数据量同步增长,大约 1 个参数对应 20 个训练 token。

  这一定律的工程含义深远:

  • 如果你有 1000 张 GPU 训练 1 个月的算力预算,应该训练一个中等规模的模型(如 7B),配合足够多的数据(如 140B tokens),而不是盲目追求更大的模型。
  • 不足训练(undertrained)的大模型往往不如同等算力训练的更小模型。

1.6 数据课程(Data Curriculum)

  数据的训练顺序也会影响最终性能。常见策略:

  • 先通用后专业:先用大规模通用网页数据训练,后期逐步混入高质量的专业数据(代码、数学、学术论文)。
  • 质量上行:训练后期逐步增加高质量数据的比例,让模型在"精品数据"上精细调整。
  • 领域加权:根据下游任务需求,调整不同领域数据的采样比例。

类比:学生成长。小学阶段学广泛的基础知识(通用网页数据),中学阶段逐步接触专业领域(代码、数学),大学阶段专注于高质量的深度学习(精标数据)。


2. 微调策略——从 SFT 到对齐

2.1 监督微调(SFT)

  预训练后的基座模型(Base Model)虽然有很强的语言能力,但它不会"聊天"——给它一个问题,它可能会继续写文章而不是回答。监督微调(Supervised Fine-Tuning, SFT) 的目标是教会模型遵循指令、以对话形式回答问题。

指令数据格式:每条数据是一个 Prompt-Response 对:

<|user|>请解释什么是量子计算。<|assistant|>量子计算是一种利用量子力学原理进行计算的技术...

关键细节——Masked Loss:训练时,损失仅计算在 Response 部分,Prompt 部分的 token 不参与损失计算。这是因为我们希望模型学习"如何回答",而不是"如何提问"。

$$ \mathcal{L}_{\text{SFT}} = -\frac{1}{|R|}\sum_{t \in R} \log p_\theta(x_t \mid x_{\lt t}) $$

其中 $R$ 是 Response 部分的 token 集合。

SFT 的本质:SFT 不是在"注入新知识"——预训练阶段已经学到了这些知识。SFT 的作用是:

  1. 激活能力:将预训练已学到但不明显的能力"激活"出来。
  2. 格式化输出:教会模型以对话、结构化的形式输出。
  3. 对齐行为:让模型学会拒绝有害请求、承认不确定性等行为模式。

高质量数据的构建

  • 人工标注:质量最高但成本昂贵。
  • 蒸馏(Distillation):用更强的模型(如 GPT-4)生成回复,作为 SFT 数据。
  • Self-Instruct:让模型自己生成指令和回复,人工筛选后使用。

2.2 基于人类反馈的强化学习(RLHF)

  SFT 之后,模型可以回答问题了,但它的回答质量参差不齐——有时冗长、有时不安全、有时事实不准确。RLHF 通过引入人类偏好来进一步优化模型。

RLHF 的三步走

flowchart TD
    A["基座模型(Base Model)"] --> B["第一步:SFT 微调"]
    B --> C["SFT 模型"]
    C --> D["第二步:训练奖励模型(RM)"]
    D --> E["奖励模型(RM)"]
    C --> F["第三步:PPO 强化学习"]
    E --> F
    F --> G["对齐后的模型"]
    H["人类标注偏好数据"] --> D
    H --> F

第一步:SFT——如上所述,用指令数据微调基座模型。

第二步:训练奖励模型(Reward Model, RM)——将"人类偏好"转化为一个标量分数。

  数据收集方式:给定一个 prompt,让模型生成多个回复(如 A、B、C),人类标注员对它们进行排序(如 A > B > C)。这些排序被转化为两两比较的偏好对。

  奖励模型的训练使用 Bradley-Terry 模型

$$ \mathcal{L}_{\text{RM}} = -\log \sigma(r_\theta(x, y_w) - r_\theta(x, y_l)) $$

其中 $y_w$ 是人类偏好的回复(winner),$y_l$ 是不被偏好的回复(loser),$r_\theta$ 是奖励模型输出的标量分数。

类比:美食评审。SFT 让模型学会了做菜,但不知道什么菜好吃。奖励模型就像一位美食评审——它通过品尝大量的菜品(回复),学会了判断"这道菜好不好吃"(回复质量高不高)。

第三步:PPO 优化策略模型——用奖励模型的分数作为"奖励信号",通过强化学习优化策略模型。

  PPO(Proximal Policy Optimization)的优化目标:

$$ \mathcal{L}_{\text{PPO}} = \mathbb{E}\left[r_\theta(x, y)\right] - \beta \cdot D_{\text{KL}}\left[\pi_\theta(y|x) \| \pi_{\text{ref}}(y|x)\right] $$

其中:

  • 第一项是奖励最大化:让模型生成高奖励的回复。
  • 第二项是 KL 散度惩罚:防止策略模型偏离参考模型太远。

四模型协作图

模型角色是否更新显存占用
Policy Model当前正在优化的模型完整模型
Reference ModelSFT 后的冻结模型,用于计算 KL 惩罚完整模型(冻结)
Reward Model评估回复质量,输出标量奖励否(已预训练好)通常小于 LM
Critic Model预估每个 token 的长期价值(价值函数)与 Policy 同等大小

  同时维护四个模型对显存的压力巨大——RLHF 训练通常需要至少 4 倍于模型本身的显存,这也是 RLHF 训练成本高昂的主要原因之一。

类比:风筝的线。KL 惩罚就像风筝的线——让模型自由飞翔(优化回复质量),但线不能断(不能偏离太远),否则风筝会飞丢(模型"忘本",生成不自然的内容)。

2.3 DPO——直接偏好优化

  RLHF 的复杂性(四个模型、PPO 的不稳定性、奖励模型的训练)催生了更简洁的替代方案。

  DPO(Direct Preference Optimization) 的核心洞察是:奖励模型可以被解析地消掉,直接用偏好数据优化策略模型。

DPO 的损失函数

$$ \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) $$

直观解释

  • $\log \frac{\pi_\theta(y_w|x)}{\pi_{\text{ref}}(y_w|x)}$:当前模型相对于参考模型,对"好回复"的概率提升了多少。
  • $\log \frac{\pi_\theta(y_l|x)}{\pi_{\text{ref}}(y_l|x)}$:当前模型相对于参考模型,对"坏回复"的概率提升了多少。
  • 损失函数鼓励增大好回复的相对概率,同时减小坏回复的相对概率

DPO 的核心优势

对比项RLHF (PPO)DPO
需要奖励模型
模型数量4 个2 个(策略 + 参考)
训练稳定性较差(PPO 对超参数敏感)好(标准监督学习)
计算成本高(四模型前向 + PPO 采样)低(类似 SFT)
效果上限可能更高通常接近 RLHF

为什么 DPO 更稳定? 从梯度的角度看,RLHF 中 PPO 的梯度信号需要经过奖励模型的"二次翻译"——奖励模型的误差会被放大到策略模型的更新中。而 DPO 的梯度直接来自偏好数据,信号更"纯净":

$$ \nabla_\theta \mathcal{L}_{\text{DPO}} \propto \underbrace{\sigma(-\hat{r}_\theta)}_{\text{权重:难样本给更多}} \cdot \left[\underbrace{\nabla_\theta \log \pi_\theta(y_w|x)}_{\text{提升好回复概率}} - \underbrace{\nabla_\theta \log \pi_\theta(y_l|x)}_{\text{降低坏回复概率}}\right] $$

其中 $\hat{r}_\theta = \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)}$。权重 $\sigma(-\hat{r}_\theta)$ 表示:当模型已经能很好地区分好/坏回复时,梯度自动衰减(不会过度优化)。

2.4 安全对齐与红队测试

  对齐的另一个重要维度是安全性。模型需要学会:

  • 拒绝有害请求:不生成暴力、色情、违法内容。
  • 承认不确定性:不知道的时候说"我不确定",而不是编造答案。
  • 抵抗越狱(Jailbreak):不被精心构造的 prompt 绕过安全限制。

红队测试(Red Teaming) 是一种主动的安全评估方法:让专门的团队(红队)尝试各种方式攻击模型,发现安全漏洞后用这些数据进一步训练模型。


3. 分布式训练——突破单卡算力与显存

  当模型规模达到数十亿甚至数千亿参数时,单张 GPU 已经无法容纳整个模型,更不用说训练了。分布式训练是解决这一问题的核心技术。

3.1 数据并行(Data Parallelism)

  最简单的并行策略:每张 GPU 持有完整的模型副本,但处理不同的数据批次

flowchart LR
    subgraph GPU0["GPU 0"]
        M0["模型副本 0"] --> G0["梯度 0"]
    end
    subgraph GPU1["GPU 1"]
        M1["模型副本 1"] --> G1["梯度 1"]
    end
    subgraph GPU2["GPU 2"]
        M2["模型副本 2"] --> G2["梯度 2"]
    end
    G0 --> AR["AllReduce:梯度求平均"]
    G1 --> AR
    G2 --> AR
    AR --> UP["同步更新参数"]
    UP --> M0
    UP --> M1
    UP --> M2

流程:每个 GPU 独立做前向和反向传播,得到各自的梯度,然后通过 AllReduce 操作求梯度平均值,最后用相同的平均梯度更新各自的模型副本。

瓶颈:当模型太大时,每张卡都要存储完整的模型参数、梯度和优化器状态,显存不够用。数据并行解决的是算力瓶颈(更多 GPU 处理更多数据),但不解决显存瓶颈

3.2 模型并行(Model Parallelism)

  当单张 GPU 装不下模型时,需要将模型本身拆分到多张 GPU 上。

张量并行(Tensor Parallelism)

  将单层的权重矩阵切分到多张 GPU 上。

类比:把一张巨大的工作表拆给多人同时计算。想象你需要计算一个 $1000 \times 1000$ 的矩阵乘法,一个人算太慢。张量并行就像把这张大表拆成几块,每人算一块,最后拼起来。

  以 FFN 层为例:$Y = \text{GeLU}(XA)B$,将权重矩阵 $A$ 按列切分到 2 张 GPU:

$$ A = [A_1 | A_2], \quad Y_1 = \text{GeLU}(XA_1), \quad Y_2 = \text{GeLU}(XA_2) $$

每张 GPU 计算一部分,然后通过 AllGather 操作拼接结果。

流水线并行(Pipeline Parallelism)

  将不同的层分配到不同的 GPU 上。

类比:工厂流水线。GPU 0 负责第 1-10 层,GPU 1 负责第 11-20 层,依此类推。数据像产品一样在流水线上逐层传递。

流水线并行气泡优化

  朴素的流水线并行存在严重的气泡(Bubble) 问题——当一个 GPU 在计算时,其他 GPU 空闲。解决方案是微批次(Micro-batch):将一个大 batch 拆成多个小的 micro-batch,让不同的 micro-batch 在流水线的不同阶段同时执行,减少空闲时间。

3.3 3D 并行

  实际训练大模型时,通常将三种并行策略组合使用:

  • 张量并行:在同一服务器内的 GPU 之间(高速 NVLink 互联)
  • 流水线并行:在不同服务器之间(通信带宽较低)
  • 数据并行:在所有 GPU 之间

这是 GPT-4、LLaMA-3 等大模型训练的标准做法。

序列并行(Sequence Parallelism):与张量并行配合使用。在张量并行中,LayerNorm 和 Dropout 等操作在每张卡上独立计算,但输入是相同的(冗余存储)。序列并行将这些操作的输入沿序列维度切分,消除冗余存储,进一步节省显存。

FSDP(Fully Sharded Data Parallel):PyTorch 原生的 ZeRO-3 实现,功能等价于 DeepSpeed ZeRO-3,但与 PyTorch 生态集成更紧密。对于已在 PyTorch 上开发的项目,FSDP 通常是更方便的选择。

3.4 ZeRO 优化——极致显存节省

  DeepSpeed 的 ZeRO(Zero Redundancy Optimizer) 是数据并行的一种改进,核心思想是消除数据并行中的冗余存储

  在标准数据并行中,每张 GPU 都存储了完整的模型参数、梯度和优化器状态(Adam 需要额外存储一阶矩和二阶矩),这造成了巨大的显存浪费。

ZeRO 的三个阶段

阶段分片内容显存节省通信开销
ZeRO-1优化器状态~4x与 DP 相同
ZeRO-2优化器状态 + 梯度~8x与 DP 相同
ZeRO-3优化器状态 + 梯度 + 参数~$N$x($N$ 为 GPU 数)增加 1.5x

ZeRO 显存分片示意

类比:分摊账单。假设一桌 10 个人吃饭,账单 1000 元。

  • 标准数据并行:每个人都存一份完整的 1000 元账单(冗余存储)。
  • ZeRO-1:每人只记住自己该付的 100 元(优化器状态分片)。
  • ZeRO-2:每人只记住自己的份额和消费明细(梯度也分片)。
  • ZeRO-3:每人只记住自己点的菜和该付的钱(参数也分片),需要用别人的菜时临时借来看一眼。

显存计算实例:以一个 7B 参数的模型为例,使用 AdamW 优化器(FP32):

组成部分每参数字节数7B 模型总显存
参数(FP32)4 bytes28 GB
梯度(FP32)4 bytes28 GB
优化器状态(FP32)8 bytes(一阶矩+二阶矩)56 GB
合计16 bytes112 GB

  一张 A100 80GB 的 GPU 连一个 7B 模型都装不下!而使用 BF16 混合精度 + ZeRO-3 在 8 张 GPU 上,每张卡只需存储 112GB / 8 = 14GB,完全在显存范围内。

  ZeRO-3 的代价是通信量分析:在 Ring AllReduce 中,$N$ 张 GPU 传输总数据量为 $2 \cdot \frac{N-1}{N} \cdot M$($M$ 为参数量),几乎是理论下限的 2 倍。ZeRO-3 在此基础上额外增加约 1.5 倍通信量(因为前向和反向都需要 AllGather 参数)。但通过计算与通信重叠,实际开销通常控制在 20% 以内。

  ZeRO-3 的代价是通信量增加——当需要其他 GPU 上的参数时,需要通过 AllGather 临时获取。但通过计算与通信重叠,可以将这一开销降到最低。

3.5 通信原语

  分布式训练的核心操作是 GPU 之间的通信。三种基本通信原语:

  • AllReduce:所有 GPU 各自有一个值,最终所有 GPU 都得到这些值的求和(或平均)。用于数据并行中的梯度同步。
  • AllGather:每个 GPU 持有一部分数据,最终所有 GPU 都得到完整的拼接数据。用于 ZeRO-3 的参数收集。
  • ReduceScatter:与 AllGather 相反。每个 GPU 有一份完整数据,最终每个 GPU 只得到一部分(归约后的)。用于 ZeRO 的梯度分片。

计算与通信重叠是分布式训练优化的关键。理想情况下,GPU 在计算当前层梯度的同时,应该已经在传输上一层的梯度。这就像一边洗菜一边切菜,让"计算"和"通信"两条流水线并行运转。


4. 混合精度训练——速度与精度的平衡

4.1 为什么需要混合精度

  FP32(32 位浮点数)精度高但慢且占显存。纯 FP16 速度快但容易出现问题:

  • 精度不足:FP16 的尾数只有 10 位,对于很小的梯度值可能被量化为 0(梯度下溢)。
  • 动态范围有限:FP16 的指数位只有 5 位,大数值容易溢出。

  混合精度训练的核心思想:用 FP16/BF16 加速计算,用 FP32 保证精度。

4.2 标准混合精度流程

  1. 维护一份 FP32 主权重副本
  2. 每次迭代时,将主权重转为 FP16,执行前向和反向传播(利用 GPU 的 FP16 Tensor Core 加速)。

混合精度训练流程 3. 得到 FP16 梯度后,转回 FP32 更新主权重。

4.3 损失缩放(Loss Scaling)

  FP16 的最小正规数约为 $6 \times 10^{-8}$,而许多梯度值小于此,会直接变为 0(下溢)。

  损失缩放的解决方案:在计算损失时乘以一个大的缩放因子 $S$(如 1024),让梯度值放大到 FP16 能表示的范围内,更新权重时再除以 $S$ 还原。

$$ \text{scaled\_loss} = S \cdot \mathcal{L}, \quad \text{grad} = \frac{\nabla(\text{scaled\_loss})}{S} $$

动态损失缩放:自动调整 $S$ 的大小——如果连续多步没有出现 Inf/NaN(没有溢出),则增大 $S$;如果出现溢出,则减小 $S$ 并跳过本步更新。

4.4 BF16——大模型训练的首选

  BF16(Brain Float 16) 是 Google 专门为深度学习设计的浮点格式:

特性FP16BF16FP32
总位数161632
指数位588
尾数位10723
动态范围$\pm 6.5 \times 10^4$$\pm 3.4 \times 10^{38}$$\pm 3.4 \times 10^{38}$
精度较高较低最高
需要 Loss Scaling通常不需要不需要

  BF16 的动态范围与 FP32 相同(同为 8 位指数),这意味着梯度几乎不会溢出或下溢,因此通常不需要损失缩放。虽然尾数精度不如 FP16,但在实践中对训练的影响很小。BF16 已成为现代大模型训练的默认选择。

4.5 自动混合精度(AMP)

  PyTorch 的 AMP(Automatic Mixed Precision)框架可以自动管理混合精度:

  • 自动将适合的操作(矩阵乘法等)转为 FP16/BF16 执行。
  • 自动插入类型转换(FP32 ↔ FP16)。
  • 自动管理损失缩放(FP16 模式下)。
# PyTorch AMP 示例
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
    output = model(input)
    loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

5. 梯度累积——用小 batch 模拟大 batch

5.1 动机

  大 batch 训练对稳定性有益,但显存装不下那么大的 batch。梯度累积是一种用时间换空间的技巧。

5.2 原理

  将一个大 batch 拆成 $K$ 个小的 micro-batch,连续计算 $K$ 步梯度并累加,在第 $K$ 步才执行一次参数更新:

$$ g_{\text{accumulated}} = \frac{1}{K}\sum_{i=1}^{K} g_i, \quad w \leftarrow w - \eta \cdot g_{\text{accumulated}} $$

等效 batch 大小 = micro-batch 大小 $\times$ 累积步数 $K$。

梯度累积示意

类比:攒快递。每次收到一个包裹(micro-batch 梯度),先放在储物间不拆(累积),等攒够了一起拆(参数更新),相当于一次性收到了一个大包裹(大 batch)。

5.3 与学习率的关系

  线性缩放规则:当 batch 大小变为 $N$ 倍时,学习率也应该线性缩放为 $N$ 倍。

$$ \eta_{\text{new}} = N \cdot \eta_{\text{base}} $$

原因直觉:batch 变大 $N$ 倍,梯度的方差减小 $N$ 倍(更稳定),所以可以走更大的步幅。

5.4 与分布式训练的协同

  在分布式设置中,等效 batch 大小还要乘以 GPU 数量:

$$ \text{等效 batch} = \text{micro-batch} \times K_{\text{accum}} \times N_{\text{GPU}} $$

例如:micro-batch=2, 累积步数=8, 64 张 GPU → 等效 batch = 2 × 8 × 64 = 1024 个序列。


6. 显存优化的其他利器

6.1 梯度检查点(Gradient Checkpointing)

  在标准的前向传播中,每一层的中间激活值都需要保存,用于反向传播计算梯度。对于 96 层的模型,这些激活值可能占用总显存的 60%~70%。

  梯度检查点(Gradient Checkpointing) 的策略:只保存部分层的激活值,其余的在反向传播时重新计算

flowchart LR
    subgraph 标准训练
        F1["前向:保存全部激活"] --> B1["反向:直接使用"]
    end
    subgraph 梯度检查点
        F2["前向:只保存检查点激活"] --> B2["反向:重新计算中间激活"]
    end

代价:前向传播的计算量增加约 33%(每个检查点的层需要前向计算两次)。 收益:显存占用从 $O(L)$ 降到 $O(\sqrt{L})$,其中 $L$ 是层数。

类比:考试复习。标准训练相当于考试前把所有笔记都摊在桌上(保存全部激活);梯度检查点相当于只记几个关键点(检查点),需要其他内容时翻书查找(重新计算)。

6.2 FlashAttention

  第一篇已介绍 FlashAttention 的基本原理。从显存角度看,它的关键优势是避免存储完整的 $n \times n$ 注意力矩阵

  • 标准注意力:显存 $O(n^2)$,需要存储所有 $QK^T$ 的注意力分数。
  • FlashAttention:显存 $O(n)$,通过分块计算,在 SRAM 中完成 Softmax 后直接丢弃中间结果。

6.3 高效数据加载

  大规模预训练的另一个工程挑战是数据加载效率。如果 GPU 训练速度快于数据加载速度,GPU 就会空等(I/O 瓶颈)。

常见优化:

数据加载预取优化

  • 预取(Prefetching):在 GPU 计算当前 batch 的同时,CPU 已经在准备下一个 batch 的数据。
  • 流式处理(Streaming):不将整个数据集加载到内存,而是按需从磁盘或网络读取。
  • 在线数据增强:在数据加载过程中完成 tokenization、padding 等预处理,避免训练时的额外计算。

7. 训练稳定性与监控

7.1 梯度裁剪

  梯度裁剪是训练稳定性的第一道防线(详见第一篇)。在大规模训练中,由于 batch 更大、训练步数更多,梯度爆炸的风险更高,梯度裁剪几乎总是开启的(典型阈值 $\tau = 1.0$)。

7.2 损失尖峰(Loss Spike)

  大规模训练中一个令人头疼的现象:损失突然飙升,然后缓慢恢复。

常见原因

  • 数据质量问题:某个 batch 包含异常数据(如极长文本、乱码)。
  • 学习率过高:在某个参数区域导致不稳定更新。
  • 数值溢出:FP16 训练中的梯度/激活溢出。

应对策略

  1. 跳过异常更新:当检测到梯度范数异常大时,跳过本步的参数更新。
  2. 回滚检查点:如果损失持续异常,回滚到之前的检查点重新训练。
  3. 降低学习率:在损失尖峰后临时降低学习率,待恢复后再逐步回升。

7.3 关键监控指标

指标含义正常范围
训练损失当前 batch 的损失值应该稳步下降
验证困惑度模型在验证集上的 PPL比训练损失更可靠
梯度范数所有参数梯度的 L2 范数应该稳定,不应剧烈波动
吞吐量每秒处理的 token 数(token/sec)反映训练效率
MFU模型浮点运算利用率(Model FLOPs Utilization)理论峰值的 40%~60% 为正常

MFU 的定义

$$ \text{MFU} = \frac{\text{实际吞吐量} \times \text{模型参数量} \times 6}{\text{GPU 理论 FLOPS} \times \text{GPU 数量}} $$

其中 6 是因为每个 token 的前向+反向传播大约需要 6 倍参数量的浮点运算。MFU 反映了硬件的利用效率——如果 MFU 只有 10%,说明有 90% 的算力被浪费在通信、等待、冗余计算上。

训练损失尖峰与恢复

实际案例——LLaMA 训练中的 Loss Spike:Meta 在训练 LLaMA 65B 时报告,训练过程中出现了约 10 次明显的损失尖峰。他们的处理策略是:检测到尖峰后回滚到尖峰前的检查点,跳过产生尖峰的数据批次,然后继续训练。值得注意的是,这些尖峰并不影响最终模型质量——只要及时处理,模型可以恢复正常。


8. 全文总结——四大瓶颈与解决方案

flowchart TD
    B1["算力瓶颈"] --> S1["数据并行 / 3D 并行"]
    B2["显存瓶颈"] --> S2["ZeRO / 梯度检查点 / 混合精度 / 梯度累积"]
    B3["通信瓶颈"] --> S3["张量并行(NVLink)/ 计算通信重叠"]
    B4["数据质量瓶颈"] --> S4["去重 / 质量分类器 / 数据课程"]
    S1 --> T["大语言模型训练"]
    S2 --> T
    S3 --> T
    S4 --> T
    T --> FT["微调对齐:SFT → RLHF/DPO"]

预训练阶段:数据工程(去重、质量过滤、数据课程)解决数据质量瓶颈;Scaling Laws 指导算力的最优分配。

分布式训练:数据并行解决算力瓶颈;ZeRO、梯度检查点、混合精度、梯度累积解决显存瓶颈;张量并行与计算通信重叠缓解通信瓶颈。

微调对齐:SFT 激活预训练能力并格式化输出;RLHF 通过人类偏好进一步优化回复质量;DPO 提供了更简洁稳定的替代方案。

工程稳定性:梯度裁剪、损失尖峰处理、MFU 监控等确保训练过程不崩溃、不浪费算力。


系列导航