第 8 章 LLM混合精度ZeRO

第 8 章 高效训练与微调

第 8 章 高效训练与微调

学习目标

  • 理解混合精度训练的数值细节:FP16 的损失缩放、BF16 为什么更稳、FP8 的分级量化;
  • 掌握梯度检查点(重计算)的显存-计算权衡;
  • 理解 ZeRO 三级分片与 FSDP 的关系,会按显存账选择策略;
  • 掌握 LoRA/QLoRA/Adapter/Prefix Tuning 的原理与选择;
  • 理解长上下文训练(位置插值、序列并行)的挑战;
  • 了解算子优化、编译、通信重叠等训练加速手段。

8.1 高效训练 = 省显存 + 省时间

第 4 章算过账:7B 全参训练单卡需要 ~112GB + 激活,80GB 卡根本放不下;70B 更是要几百 GB。本章所有技术都归两个目标:

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
mindmap
  root((高效训练))
    省显存
      混合精度
        FP16 损失缩放
        BF16 宽范围
        FP8 分级
      梯度检查点
        重计算换显存
      ZeRO/FSDP
        分片不复制
      LoRA/QLoRA
        少训参数
        小优化器状态
    省时间
      算子优化
        FlashAttention
      编译
        torch.compile
      通信重叠
        计算通信并行

一句话总纲:显存不够 → 分片(ZeRO/FSDP)或重算(检查点)或减少状态(LoRA);速度不够 → 降精度(FP8)或融合算子或重叠通信

8.2 混合精度:FP16、BF16、FP8

8.2.1 三种格式的数值本质

浮点数的精度由「指数位 + 尾数位」决定,两者各管一件事:

格式 符号 指数 尾数 动态范围 精度 备注
FP32 1 8 23 ~103810^{38} 基准,主权重用
FP16 1 5 10 ~10510^{5} 范围小是死穴
BF16 1 8 7 ~103810^{38} 低(同范围同 FP32) 范围大是救星
FP8 (E4M3/E5M2) 1 4/5 3/2 更窄/更宽 很低 H100 时代新增

FP16 的问题:尾数多但指数少——能表示很精细的数,但范围只有 6 万量级。训练里梯度经常小到 10810^{-8} 以下(下溢为 0)或激活大到溢出为 inf。BF16 用尾数换指数:精度粗一点,但范围与 FP32 相同,上下溢风险大幅降低。这就是第 7 章「BF16 优先防 loss spike」的数值根源。现代硬件(Ampere 之后 NVIDIA、多数国产卡)都支持 BF16,新项目无脑选 BF16

8.2.2 混合精度的实现:Loss Scaling 与主权重

混合精度训练不是「全部用低精度」,而是分角色用不同精度

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
flowchart LR
    A[FP32 主权重
optimizer 里的 master weights] -->|cast| B[BF16/FP16 前向] B --> C[BF16 反向
得到梯度] C -->|FP16 时先放大| D[Loss Scaling
防梯度下溢] D -->|unscale + cast| E[FP32 优化器更新] E --> A
  • 前向/反向:低精度(快,Tensor Core 只在低精度下满血);
  • 梯度:低精度传递(FP16 需要先乘 loss scale 放大,更新前再除回);
  • 主权重与优化器矩:FP32(累加精度,低精度累加会漂移)。

第 4 章「优化器状态 12N 字节 = 4+4+4」里那个 FP32 主权重,就是为混合精度存在的。

8.2.3 FP8 训练

FP8 在 H100 上把 Tensor Core 吞吐再翻倍,但两种子格式(E4M3 精度优先用于前向权重/激活、E5M2 范围优先用于反向梯度)都不够稳,所以需要 per-tensor 或 per-channel 的动态 scaling(DeepSeek-V3 的 FP8 训练把细粒度分组缩放做到 block 级)。工程现状:FP8 预训练已在头部团队生产化,微调场景仍以 BF16 为主——精度收益与稳定性风险要按团队工程能力权衡

8.3 梯度检查点:重计算换显存

不存全部前向激活,反向时需要哪层、重算哪层:

# PyTorch 原生一行开启(transformer 层内部自动分段重算)
from torch.utils.checkpoint import checkpoint_sequential
outputs = checkpoint_sequential(layers, segments=8, input=x)
# 或模型级开关
model.gradient_checkpointing_enable()

权衡账:显存:激活从 O(L) 降到 O(√L)(分段存储);时间:多一次前向,约 +30% 训练时间。它是「单卡塞下更大 batch/模型」的第一选择,且与 ZeRO 正交(可叠加)。经验法则:显存紧张先开检查点,吞吐不够再关。

8.4 ZeRO 与 FSDP:显存分片

8.4.1 数据并行的冗余问题

朴素数据并行(DDP)里每张卡都完整持有权重、梯度、优化器状态——三份冗余。第 4 章的账:7B 模型这三样共 112GB/卡,明明总共也是 112GB,八卡却占了 896GB。ZeRO(Zero Redundancy Optimizer)的思想:把这三样东西切分到各卡,用时再聚合

级别 分片什么 显存节省 通信代价
ZeRO-1 优化器状态 基本不变
ZeRO-2 + 梯度 更大 略增
ZeRO-3 + 权重 最大(每卡只存 1/N) +50%(权重 all-gather)

ZeRO-3 的通信账:前向需要 all-gather 权重分片,反向也要一次——换来的是「单卡 N 分之一的显存」。DeepSpeed 的 ZeRO 与 PyTorch 原生 FSDP(Fully Sharded Data Parallel)本质是同一思想,FSDP 可理解为 PyTorch 版 ZeRO-3(社区生态更活跃)。

8.4.2 怎么选

%%{init: {"themeVariables": {"primaryTextColor": "#000000", "textColor": "#000000", "labelColor": "#000000", "nodeTextColor": "#000000", "labelTextColor": "#000000", "scaleLabelColor": "#000000"}}}%%
flowchart TD
    A[模型塞得下单卡?] -->|是| B[DDP + 梯度累积
最简单最快] A -->|否| C[开梯度检查点后塞得下?] C -->|是| D[DDP + checkpointing] C -->|否| E[FSDP/ZeRO-3 分片] E -->|仍不够| F[+ CPU Offload
慢但能跑] E -->|要训 100B+| G[+ 张量/流水并行
第 10-11 章]

微调 7B-70B 的现实路径:7B 单卡(QLoRA)或 FSDP;70B 用 FSDP 多机或 DeepSpeed ZeRO-3。Offload(把优化器状态/权重卸到 CPU 内存)是最后手段——NVMe/CPU 带宽远低于 GPU,速度掉 2-5 倍。

8.5 参数高效微调:LoRA、QLoRA 与家族

8.5.1 全参微调的问题

全参微调(Full Fine-Tuning)动全部参数:7B 模型训练态要 112GB+(第 4 章),每个任务还要存一份完整模型(部署时 7B × N 任务)。参数高效微调(PEFT)的问题意识:冻结原权重,只训一个很小的增量

8.5.2 LoRA:低秩增量

核心假设:微调引起的权重变化 ΔW\Delta W 是低秩的。于是不直接学 ΔWRd×d\Delta W \in \mathbb{R}^{d \times d},而是分解成两个瘦矩阵:

W=W+ΔW=W+BA,BRd×r, ARr×d, rdW' = W + \Delta W = W + BA,\quad B \in \mathbb{R}^{d \times r},\ A \in \mathbb{R}^{r \times d},\ r \ll d

前向时 h=Wx+BAxh = Wx + BAxBB 初始化为 0,保证起点等于原模型)。可训参数从 d2d^2 降到 2dr2dr——以 7B 模型、r=16r=16、只对注意力投影加 LoRA 为例,可训参数约 0.1%,训练显存从 112GB 级降到 ~20GB 级(权重冻结后没有梯度与优化器状态,只有激活与小矩阵状态)。

# HuggingFace PEFT 的 LoRA 微调骨架
from peft import LoraConfig, get_peft_model

config = LoraConfig(
    r=16, lora_alpha=32,              # 常用 r=8/16/32,alpha≈2r
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
    lora_dropout=0.05, task_type="CAUSAL_LM",
)
model = get_peft_model(base_model, config)
model.print_trainable_parameters()
# trainable params: ~0.1% || all params: 100% 

LoRA 的工程红利不止省显存:

  • 多任务即插即用:一份基座 + N 份 LoRA 适配器(每份几十 MB),部署时动态加载(第 17 章多租户服务的标配);
  • 防遗忘:原权重不动,通用能力天然保留;
  • 训练快:优化器状态小一个量级,步速更高。

超参经验:r=16r=16 起步,任务复杂(领域差异大)升到 64-128;lora_alpha 设为 2r2r;target_modules 至少覆盖全部注意力投影,复杂任务加 FFN(gate/up/down_proj)。

8.5.3 QLoRA:4bit 基座 + LoRA

QLoRA 在 LoRA 上再压一层:基座权重量化到 4bit(NF4)存储,LoRA 增量用 BF16 训练,反向时反量化计算。三个关键技术:

  1. NF4(4-bit NormalFloat):按正态分布分位数设计的量化格点,对呈高斯分布的权重更友好;
  2. 双重量化:把量化的缩放常数本身也量化,再省 0.4 bit/参数;
  3. 分页优化器:显存峰值时把优化器状态临时换页到 CPU。

效果:7B 微调降到单张 24GB 消费卡可跑(如 4090),48GB 卡可跑 33B。代价:比 BF16 LoRA 慢 30-70%(反量化开销)。定位明确:资源受限场景的平民神器,生产训练集群上仍以 BF16 + FSDP 为主。

8.5.4 家族其他成员与选择

方法 思路 参数量 推理是否加开销
LoRA/QLoRA 权重旁路低秩增量 ~0.1-1% 可合并进权重,零开销
Adapter 层间插入小 MLP ~1-5% 有串行小层,+延迟
Prefix/Prompt Tuning 虚拟 token 前缀 ~0.01% 占用上下文窗口
BitFit 只训 bias ~0.01% 零开销,效果有限

选择逻辑:默认 LoRA(可合并、效果好、生态最全)→ 显存极限上 QLoRA → 多任务软提示探索用 Prefix Tuning → 学术对比才考虑 Adapter/BitFit。LoRA 可合并(W=W+BAW' = W + BA 合并后与原模型同构)是它统治地位的工程原因——部署侧无任何额外成本。

8.6 长上下文训练

第 3 章讲过上下文窗口与 RoPE。把窗口从 4K 扩到 128K+,训练侧有三重障碍:

  1. 显存:激活与注意力矩阵随 SS 平方增长——FlashAttention(第 14 章详述)把显存从 O(S2)O(S^2) 降到 O(S)O(S),是长上下文的前提;
  2. 数据:长文档语料稀缺(网页平均长度远小于 128K),要做长文档拼接与「文档注意力掩码」防跨文档泄漏;
  3. 位置编码外推:训练时见过的相对距离有限,推理遇到更长序列会崩。

主流扩展方案(都作用于 RoPE):

方案 思路 代价
位置插值(PI) 把长序列的位置索引线性压缩回训练范围 需少量继续训练
NTK-aware 缩放 改 RoPE 基频,低频外推高频插值 免训练或轻训练
YaRN 分频段差异化缩放 PI 与 NTK 目前常用折中

序列并行(DeepSpeed-Ulysses、Ring Attention)解决「单卡放不下整条序列」:把序列切到多卡,注意力计算时环形传递 K/V。它与张量并行(第 10 章)的区别:切的是序列维度而不是权重维度,激活显存随卡数线性下降。

工程路径总结:短窗口预训练 → 高质量长文档 + 位置插值继续训练 → 长上下文评测。全程 FlashAttention 必开。

8.7 训练加速:算子、编译、通信

三个「不用改算法」的提速手段:

算子融合:把多个小 kernel 合成一个大 kernel,省的是 HBM 往返。代表作 FlashAttention:把注意力的 softmax 分块在 SRAM 里完成,数学等价但 IO 减少 ~10×,训练推理通吃。类似思想遍布 LayerNorm 融合、GeLU 融合等。

编译torch.compile(PyTorch 2.x)把 eager 执行图捕获编译优化,模式有 reduce-overhead(CUDA Graph 消 kernel 启动开销)。训练场景常见 1.3-2× 提速,代价是首次编译等待与动态形状的重新编译——固定 shape 的预训练收益最大,变长微调要 mark_dynamic

model = torch.compile(model, mode="max-autotune")  # 训练前一行

通信重叠:梯度同步(通信)与反向计算(计算)重叠——反向算完一层就把该层梯度 all-reduce 发出,不等全网算完。ZeRO/FSDP 实现里都有 prefetch 与 overlap 开关,第 10-11 章展开原理。


本章要点回顾

  1. 高效训练两大目标:省显存(分片/重算/减状态)、省时间(低精度/融合/重叠)。
  2. BF16 用尾数换指数,范围同 FP32,是稳定性的来源;FP16 要 Loss Scaling;FP8 要细粒度缩放,生产化仍在推进。
  3. 混合精度分角色:前向反向低精度、优化器与主权重 FP32。
  4. 梯度检查点 +30% 时间换 √L 级激活显存;ZeRO-3/FSDP 分片冗余状态,通信 +50% 换 N 分之一显存。
  5. LoRA 用低秩增量把可训参数降到 0.1%,可合并进权重零部署开销;QLoRA 用 4bit 基座把 7B 微调塞进 24GB 卡。
  6. 长上下文 = FlashAttention + 位置插值(PI/NTK/YaRN)+ 序列并行 + 长文档数据。
  7. torch.compile 与通信重叠是「不改算法」的最后两级加速。

习题

  1. 7B 模型全参 BF16 训练,8 卡 FSDP(ZeRO-3)后每卡显存的理论下限是多少(忽略激活与通信缓冲)?加上激活后为什么实际要留 20% 余量?
  2. 为什么 LoRA 的 BB 矩阵初始化为零?如果 AABB 都随机初始化会发生什么?
  3. QLoRA 为什么把基座量化到 NF4 而不是 INT4?从权重分布角度解释。
  4. 你的 70B LoRA 微调吞吐只有理论 MFU 的 15%。列出至少 5 个排查点(提示:从数据、算子、通信、精度四个层面)。
  5. 位置插值 PI 和 NTK-aware 的区别是什么?为什么 NTK 能「免训练」外推?

延伸阅读

  • Micikevicius et al., FP8 Formats for Deep Learning, 2022
  • Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, 2020
  • Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel, 2023
  • Hu et al., LoRA: Low-Rank Adaptation of Large Language Models, 2021
  • Dettmers et al., QLoRA: Efficient Finetuning of Quantized LLMs, 2023
  • Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, 2022
  • Chen et al., LongLoRA: Efficient Fine-tuning of Long-Context Large Language Models, 2023
  • Peng et al., YaRN: Efficient Context Window Extension of Large Language Models, 2023
  • Rasley et al., DeepSpeed: System Optimizations Enable Training Deep Learning Models with Over 100 Billion Parameters, 2020(含 CPU Offload)

下一章预告

第 9 章回答「模型到底行不行」:从准确率/F1 到 MMLU/GSM8K/HumanEval,从 MT-Bench/Arena 的偏好评估到 RAG/Agent 评估,以及污染、红队与可信评估的雷区。