前言
什么是混合精度训练?
混合精度训练(Mixed Precision Training)是一种在模型训练中同时使用多种数值精度(如FP32 + FP16/BF16)的技术。核心思路:前向传播和反向传播用低精度(16位)计算来加速、省显存,而权重更新和优化器状态用高精度(32位)保留来保证数值稳定性。
打个比方:日常算账用计算器(低精度,快),但银行结算用高精度算盘(FP32,准)。混合精度就是「计算时用计算器,存钱时用算盘」------既快又不会算错钱。
它能将训练显存占用减少约50%,同时在现代GPU(A100/H100)上获得2-3倍加速,是目前大模型训练的标准配置。
一次混合精度训练中,各种精度各司其职:
- FP32 主权重:优化器始终保留一份 FP32 的权重副本(约 4 字节/参数),用于参数更新,保证累计的小幅更新不被舍入吞掉
- FP16/BF16 计算副本:前向传播和反向传播时,把权重临时转成 16 位(2 字节/参数)参与矩阵乘法,这是加速和省显存的主要来源
- FP16/BF16 梯度:反向传播算出的梯度以 16 位存储(FP16 需配合损失缩放,BF16 不需要)
- FP32 优化器状态:AdamW 的动量 m 和二阶矩 v 各占 4 字节/参数,保持 FP32 以保证数值稳定
- FP8(前沿) :H100 等新硬件支持用 8 位(1 字节/参数)做部分矩阵乘法,进一步提速,但主要用在超大规模训练场景
什么是分布式训练?
分布式训练(Distributed Training)是把一个模型的训练任务拆分到多张GPU上并行执行的技术。随着模型规模增长(GPT-3有1750亿参数,单张GPU根本放不下),单卡训练已经不现实,必须靠多卡协作。
主要有三种拆分思路:
- 数据并行(DDP) :每张卡存完整模型,各自处理不同数据,最后汇总梯度------最简单,但模型必须能放进单卡
- 模型并行(TP/PP) :把模型本身拆开(按矩阵切=张量并行,按层切=流水线并行),每张卡只负责一部分模型
- 混合分片(ZeRO/FSDP) :不切模型结构,而是把「参数+梯度+优化器状态」分散存储在各卡上,计算时按需拉取------兼顾显存效率和实现简洁性
核心概念
1.1 为什么需要混合精度和分布式训练?
以GPT-3(1750亿参数)为例,计算训练需要的显存:
| 组成部分 | 计算 | 占用 |
|---|---|---|
| 模型参数 | 175B × 4字节(FP32) | 700GB |
| 梯度 | 175B × 4字节 | 700GB |
| 优化器状态(AdamW):m(动量,梯度移动平均)+ v(二阶矩,梯度平方移动平均),各占4字节/参数,共约2倍参数量 | 175B × 8字节(m和v) | 1400GB |
| 激活值 | 约200GB | 200GB |
| 总计 | 约3000GB |
单张A100 80G GPU只能装80GB------需要38张GPU才能放下GPT-3的训练状态!
两个解决方向:
- 混合精度训练:减少每个参数的存储位数
- 分布式训练:把计算和存储分摊到多个GPU
1.2 混合精度训练
核心思想:不是所有计算都需要FP32(32位浮点)的高精度。大部分计算用FP16(16位)就够了,关键部分保留FP32。
浮点数格式对比
| 精度 | 位宽 | 指数位 | 尾数位 | 显存占用 | 数值范围 |
|---|---|---|---|---|---|
| FP32 | 32位 | 8 | 23 | 1x | ±3.4e38 |
| FP16 | 16位 | 5 | 10 | 0.5x | ±6.5e4 |
| BF16 | 16位 | 8 | 7 | 0.5x | ±3.4e38 |
| FP8 | 8位 | 4/5 | 3/2 | 0.25x | ±4.3e1 |
FP16 vs BF16:
BF16(Brain Float 16)是Google设计的格式:
- 指数位和FP32一样(8位),所以数值范围相同
- 尾数位少(7位 vs FP32的23位),所以精度较低
- 在NVIDIA A100/H100上,BF16和FP16性能相同
- 大模型训练推荐用BF16------因为FP16的数值范围太小(±65504),梯度容易溢出
类比:
- FP32像高清照片------清晰但占空间大
- FP16像JPEG压缩------省空间但可能丢失细节
- BF16像降低分辨率但保持色彩范围------省空间且不溢出
混合精度训练流程
markdown
1. 前向传播:权重(FP32) → 转为FP16 → 计算 → 输出转回FP32
2. 反向传播:梯度用FP16计算
3. 损失缩放(Loss Scaling):将loss乘以缩放因子(如2^16),
防止小梯度在FP16中变为0(underflow)
4. 优化器更新:梯度转回FP32,优化器状态保持FP32
5. 权重更新:FP32精度更新,再转为FP16用于下一步
上面流程里「转为FP16」具体是什么操作?
- 什么是「转」:每个浮点数在内存里就是一段二进制位。FP32有32位(1符号+8指数+23尾数),FP16有16位(1符号+5指数+10尾数)。所谓"权重从FP32转为FP16",就是对每个数重新编码------尾数位从23砍到10(四舍五入),指数位照抄。硬件上一条cast指令完成,整个模型转一遍是毫秒级的事
- 为什么前向/反向用FP16:大模型95%的计算量是矩阵乘法,A100/H100的Tensor Core跑FP16的吞吐量是FP32的好几倍。矩阵乘法里单个数的舍入误差会被求和平均掉,16位精度扛得住
- 为什么权重更新必须用FP32:优化器每步更新量极小(学习率×梯度),相对权重可能是10^-7量级。如果权重本身是FP16(只有10位尾数,约3位有效数字),0.5 + 0.0000001还是0.5------更新被舍入吞掉了,权重永远不动,模型学不到东西。FP32有23位尾数(约7位有效数字),能正确记录微调
- 为什么需要"FP32主权重" :它是"账本",FP16副本是"干活的临时工"。每个step开始时把账本抄一份出去做矩阵乘法,干完把结果记回账本
- 矩阵累加器细节:矩阵乘法内部的累加器其实是FP32,算完再加回FP16存储------防止几万个数相加时FP16的精度不够,误差滚雪球
为什么需要损失缩放?
FP16能表示的最小正数约为 ``。很多梯度的值比这还小,在FP16中会变成0------这就是underflow。
解决方案:把loss乘以一个大的缩放因子(如65536),梯度也相应放大65536倍,不会变成0。更新前再除以65536恢复。数值上等价,只是把数搬进了FP16可表示的范围内。
BF16不需要损失缩放------因为BF16的指数位和FP32一样(8位),数值范围也是±3.4e38,小梯度也能正常表示,不会underflow。这也是当前大模型训练普遍用BF16而非FP16的核心原因之一。
py
# 损失缩放示例
scaler = torch.cuda.amp.GradScaler()
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
loss = model(batch) # 前向传播用BF16
scaler.scale(loss).backward() # 反向传播 + 缩放
scaler.unscale_(optimizer) # 恢复原始梯度
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer) # 参数更新
scaler.update() # 动态调整缩放因子
1.3 分布式训练策略
数据并行(DDP - Distributed Data Parallel)
最简单的方法:每个GPU存一份完整模型,处理不同的数据批次。
yaml
GPU 0: 模型副本 + 数据批次0 → 梯度0 ─┐
GPU 1: 模型副本 + 数据批次1 → 梯度1 ─┤ All-Reduce(梯度平均)
GPU 2: 模型副本 + 数据批次2 → 梯度2 ─┤
GPU 3: 模型副本 + 数据批次3 → 梯度3 ─┘
│
平均梯度 → 各GPU更新参数
优点 :简单,通信开销小(只同步梯度)
缺点:每个GPU都要存完整模型------百亿参数模型单个GPU放不下
ZeRO(Zero Redundancy Optimizer)
DeepSpeek的核心创新。ZeRO分三个阶段逐步优化显存:
| ZeRO Stage | 优化内容 | 显存节省 | 通信开销 |
|---|---|---|---|
| Stage 1 | 优化器状态分片:DDP中每张卡存一份完整的优化器状态(AdamW的m和v,大小=2倍参数),所有卡的内容完全一样。Stage 1把m和v均匀分片到各卡,每卡只存1/N。更新时各卡只更新自己负责的那部分,通信模式跟DDP的All-Reduce一样。【补充】什么是m和v?它们是AdamW优化器给每个参数额外维护的两个统计量:m(动量) 是最近一段时间梯度的移动平均,用来抹平单个batch梯度的噪声,保留稳定的大方向;v(二阶矩) 是梯度平方的移动平均,衡量每个参数梯度的剧烈程度。两者合起来让AdamW能给每个参数自动调节学习率------梯度平稳的参数步子迈大点,梯度剧烈的参数步子收敛点。正因为m和v跟参数一一对应且都是FP32存储,175B模型的优化器状态高达1400GB,远超参数本身的700GB,成为DDP中最浪费显存的部分。 | 4x | 和DDP相同 |
| Stage 2 | 优化器状态 + 梯度分片:Stage 1基础上,再把梯度也分片。DDP中每张卡反向传播后,梯度经All-Reduce同步,最终所有卡都变成同一个平均值------同步后的梯度完全是重复存储。Stage 2用Reduce-Scatter替代All-Reduce,各卡只保留自己负责那份梯度,省掉每卡存完整梯度的开销 | 8x | 略增 |
| Stage 3 | 优化器状态 + 梯度 + 模型参数分片:最激进的一阶段,连模型参数本身也分片。每卡只存1/N的参数,计算时需要哪部分就从对应卡临时拉过来(All-Gather),用完归还。这意味着每张卡不用背完整模型了------百亿参数模型也能训。代价是每次前向/反向都要频繁通信拉参数,网络带宽成为瓶颈 | ∞(理论上) | 大幅增加 |
ZeRO Stage 3 原理:
ini
传统DDP(每个GPU存所有东西):
GPU 0: [参数0][参数1][参数2][优化器0][优化器1][优化器2][梯度0][梯度1][梯度2]
GPU 1: [参数0][参数1][参数2][优化器0][优化器1][优化器2][梯度0][梯度1][梯度2]
GPU 2: [参数0][参数1][参数2][优化器0][优化器1][优化器2][梯度0][梯度1][梯度2]
ZeRO Stage 3(每个GPU只存1/3):
GPU 0: [参数0][优化器0][梯度0]
GPU 1: [参数1][优化器1][梯度1]
GPU 2: [参数2][优化器2][梯度2]
计算时:需要哪部分参数就从对应GPU Gather过来
代价:需要频繁的All-Gather和Reduce-Scatter通信,需要高速网络(NVLink/InfiniBand)。
FSDP(Fully Sharded Data Parallel)
PyTorch官方实现的ZeRO Stage 3,做了更多工程优化:
- 自动分片策略
- 通信和计算重叠(compute-communication overlap)
- 支持CPU offload
1.4 显存组成详解
| 组成部分 | 占用比例 | 说明 |
|---|---|---|
| 模型参数 | ~20% | FP32: 4字节/参数 |
| 梯度 | ~20% | 同参数大小 |
| 优化器状态 | ~40% | AdamW需要2倍参数(m和v) |
| 激活值 | ~15% | 前向传播中间结果 |
| 其他 | ~5% | 临时缓冲区等 |
关键洞察:优化器状态是最大的显存消耗者!ZeRO Stage 1就是通过分片优化器状态来大幅节省显存。
技术细节
2.1 FSDP训练代码
py
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import MixedPrecision
def setup_distributed():
"""初始化分布式训练环境"""
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
return local_rank
def setup_fsdp_model(model, local_rank):
"""用FSDP包装模型"""
# 混合精度配置
mixed_precision_policy = MixedPrecision(
param_dtype=torch.bfloat16, # 参数用BF16
reduce_dtype=torch.bfloat16, # 梯度归约用BF16
buffer_dtype=torch.bfloat16, # 缓冲区用BF16
)
# FSDP包装
model = FSDP(
model,
mixed_precision=mixed_precision_policy,
use_orig_params=True,
device_id=local_rank,
)
return model
def train_step_fsdp(model, batch, optimizer, scheduler):
"""FSDP训练步骤"""
optimizer.zero_grad()
# 前向传播(自动用BF16)
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
loss = model(batch)
# 反向传播
loss.backward()
# 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# 参数更新
optimizer.step()
scheduler.step()
return loss.item()
2.2 不同并行策略对比
| 策略 | 切分方式 | 通信量 | 显存效率 | 实现难度 |
|---|---|---|---|---|
| DDP | 不切分(数据并行) | 小 | 低 | 简单 |
| 张量并行(TP) | 切分矩阵 | 大 | 高 | 中等 |
| 流水线并行(PP) | 切分层 | 中 | 高 | 中等 |
| ZeRO-3/FSDP | 切分参数+梯度+优化器 | 大 | 最高 | 中等 |
大模型训练通常组合使用:FSDP + TP + PP
常见误区
误区1:FP16和BF16差不多
事实:FP16数值范围小(±65504),容易溢出。BF16数值范围和FP32相同,大模型训练推荐用BF16。
误区2:GPU越多训练越快
事实:GPU之间的通信开销会随着数量增加而增大。当通信开销超过计算收益时,增加GPU不再加速训练。
误区3:ZeRO Stage 3总是最好的
事实:ZeRO-3的通信开销最大。对于能放进单机的模型,DDP或ZeRO-1可能更快。
最后
感谢你能看到这里,本文梳理了当下主流的【混合精度训练与分布式训练】流程,希望对你有用
更多 Agent、前端、Node、性能相关的技术文章和实践总结,可以查看我的代码花园: