第 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 | ~ | 高 | 基准,主权重用 |
| FP16 | 1 | 5 | 10 | ~ | 中 | 范围小是死穴 |
| BF16 | 1 | 8 | 7 | ~ | 低(同范围同 FP32) | 范围大是救星 |
| FP8 (E4M3/E5M2) | 1 | 4/5 | 3/2 | 更窄/更宽 | 很低 | H100 时代新增 |
FP16 的问题:尾数多但指数少——能表示很精细的数,但范围只有 6 万量级。训练里梯度经常小到 以下(下溢为 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 | 优化器状态 | 8× | 基本不变 |
| 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:低秩增量
核心假设:微调引起的权重变化 是低秩的。于是不直接学 ,而是分解成两个瘦矩阵:
前向时 ( 初始化为 0,保证起点等于原模型)。可训参数从 降到 ——以 7B 模型、、只对注意力投影加 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 章多租户服务的标配);
- 防遗忘:原权重不动,通用能力天然保留;
- 训练快:优化器状态小一个量级,步速更高。
超参经验: 起步,任务复杂(领域差异大)升到 64-128;lora_alpha 设为 ;target_modules 至少覆盖全部注意力投影,复杂任务加 FFN(gate/up/down_proj)。
8.5.3 QLoRA:4bit 基座 + LoRA
QLoRA 在 LoRA 上再压一层:基座权重量化到 4bit(NF4)存储,LoRA 增量用 BF16 训练,反向时反量化计算。三个关键技术:
- NF4(4-bit NormalFloat):按正态分布分位数设计的量化格点,对呈高斯分布的权重更友好;
- 双重量化:把量化的缩放常数本身也量化,再省 0.4 bit/参数;
- 分页优化器:显存峰值时把优化器状态临时换页到 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 可合并( 合并后与原模型同构)是它统治地位的工程原因——部署侧无任何额外成本。
8.6 长上下文训练
第 3 章讲过上下文窗口与 RoPE。把窗口从 4K 扩到 128K+,训练侧有三重障碍:
- 显存:激活与注意力矩阵随 平方增长——FlashAttention(第 14 章详述)把显存从 降到 ,是长上下文的前提;
- 数据:长文档语料稀缺(网页平均长度远小于 128K),要做长文档拼接与「文档注意力掩码」防跨文档泄漏;
- 位置编码外推:训练时见过的相对距离有限,推理遇到更长序列会崩。
主流扩展方案(都作用于 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 章展开原理。
本章要点回顾
- 高效训练两大目标:省显存(分片/重算/减状态)、省时间(低精度/融合/重叠)。
- BF16 用尾数换指数,范围同 FP32,是稳定性的来源;FP16 要 Loss Scaling;FP8 要细粒度缩放,生产化仍在推进。
- 混合精度分角色:前向反向低精度、优化器与主权重 FP32。
- 梯度检查点 +30% 时间换 √L 级激活显存;ZeRO-3/FSDP 分片冗余状态,通信 +50% 换 N 分之一显存。
- LoRA 用低秩增量把可训参数降到 0.1%,可合并进权重零部署开销;QLoRA 用 4bit 基座把 7B 微调塞进 24GB 卡。
- 长上下文 = FlashAttention + 位置插值(PI/NTK/YaRN)+ 序列并行 + 长文档数据。
- torch.compile 与通信重叠是「不改算法」的最后两级加速。
习题
- 7B 模型全参 BF16 训练,8 卡 FSDP(ZeRO-3)后每卡显存的理论下限是多少(忽略激活与通信缓冲)?加上激活后为什么实际要留 20% 余量?
- 为什么 LoRA 的 矩阵初始化为零?如果 、 都随机初始化会发生什么?
- QLoRA 为什么把基座量化到 NF4 而不是 INT4?从权重分布角度解释。
- 你的 70B LoRA 微调吞吐只有理论 MFU 的 15%。列出至少 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 评估,以及污染、红队与可信评估的雷区。