大模型底层学习(四)- 混合精度训练与分布式训练

18 阅读12分钟

前言

什么是混合精度训练?

混合精度训练(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
激活值约200GB200GB
总计约3000GB

单张A100 80G GPU只能装80GB——需要38张GPU才能放下GPT-3的训练状态!

两个解决方向

  1. 混合精度训练:减少每个参数的存储位数
  2. 分布式训练:把计算和存储分摊到多个GPU

1.2 混合精度训练

核心思想:不是所有计算都需要FP32(32位浮点)的高精度。大部分计算用FP16(16位)就够了,关键部分保留FP32。

浮点数格式对比

精度位宽指数位尾数位显存占用数值范围
FP3232位8231x±3.4e38
FP1616位5100.5x±6.5e4
BF1616位870.5x±3.4e38
FP88位4/53/20.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像降低分辨率但保持色彩范围——省空间且不溢出

混合精度训练流程

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的核心原因之一。

# 损失缩放示例
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存一份完整模型,处理不同的数据批次。

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 原理

传统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训练代码

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、性能相关的技术文章和实践总结,可以查看我的代码花园:

📦 github.com/AdolescentJ…

参考