最近一直有人在后台问我,大模型训练动不动就要几十张甚至上千张显卡,显存到底是怎么省下来的?那么多卡又是怎么协同工作的?我之前在系列前几篇里聊了Transformer架构和反向传播的底层逻辑,这一篇终于要进入最硬核的实战环节了:混合精度训练和分布式训练。
严格来说,这两块内容分开讲都是一门课,但实际训练大模型时它们又是深度绑定的。做混合精度是为了省显存、加快计算,做分布式是为了把省下来的显存和算力池化到更大规模。如果你打算从“能跑通小模型”迈向“能训练大模型”,这两个技术点是绕不开的坎。这篇文章我尽量用大白话,把原理和实操串起来讲明白,希望能帮你少踩几个坑。
1. 为什么要动精度和分布式的念头:一个7B模型的显存账本
先算一笔账,你就知道问题出在哪了。
一个70亿(7B)参数的大模型,按FP32精度来存,每个参数占4字节,光模型权重就是 7 × 10^9 × 4 = 28GB。可训练过程中不能只存权重吧,还得存优化器状态(AdamW的话,每个参数要存一阶动量 m 和二阶动量 v,又是两个FP32),以及反向传播时要用的梯度(又一个FP32)。再加上前向和反向过程的中间激活值,显存需求轻松超过100GB。
一张目前主流的专业显卡,比如NVIDIA A100 80GB或者H100 80GB,单独一张卡连一套完整的7B训练流程都跑不起来,更别提动辄几百B的大模型了。这就是为什么行业里常说“训练大模型不是算力问题,而是显存问题”。显存摆不平,算力再高也是空的。
解决办法有两条路:一是把每个字节“抠”着用,这就是混合精度训练要干的事;二是把数据、模型、优化器状态分散到多张卡上,这就是分布式训练要干的事。两条路不是二选一,而是同时用,这也是目前大模型训练框架的默认配置。
2. 混合精度训练:为什么不是“纯低精度”?
混合精度训练的核心直觉很简单:模型里的某些东西可以用更少的位数来存,只要不让精度崩掉就行。但这里有个关键点,名字叫“混合”,不是“全换”。很多人一开始以为把权重从FP32换成FP16就完事了,这是最容易翻车的地方。
2.1 FP16、BF16和FP32之间的差异
先过一下基础知识。FP32是单精度浮点数,1位符号位、8位指数位、23位尾数位,它能表示的数值范围大约是 ±3.4×10^38,精度能达到小数点后约7位有效数字。FP16是半精度,1位符号位、5位指数位、10位尾数位,范围只有 ±65504,有效数字约3位。
问题就出在这个“3位有效数字”上。大模型训练时,梯度值经常非常小,比如1e-7这个量级。在FP16下,这个数根本表示不了,直接就会变成0,这就是所谓的“下溢”(underflow)。梯度要是变成0了,参数就再也不更新了,模型直接废掉。
为了解决这个问题,NVIDIA和各大框架搞出了BF16(BFloat16),1位符号位、8位指数位、7位尾数位。它的指数位和FP32一样多,所以能表示的数值范围和FP32几乎一致,但尾数少了,精度低一些。这意味着BF16不会像FP16那样出现严重下溢,小梯度仍然能表示,只是不够精确。对大模型训练来说,梯度“丢精度”的影响远小于“直接清零”,因此BF16几乎成了大模型训练的默认选择。
2.2 混合精度训练的完整闭环
光是换数据类型还不够,混合精度训练真正难的是整个训练流程里每种数据用多少位。标准流程是这样:
- 主权重(master weights)保存在FP32,这是模型参数的“权威版本”,每次更新都对着它来;
- 前向传播把FP32权重临时转成FP16/BF16,用低精度计算,目的是加速矩阵乘法和省显存;
- 反向传播得到的梯度,也是低精度的,但存到梯度缓冲区时会转回FP32;
- 优化器状态(Adam里的m和v)始终是FP32,更新时用FP32梯度去更新FP32主权重。
你可能想问,为什么不一开始就把权重存成FP16?我试过这种“激进”做法,结果训练没几步loss就开始失控。原因在于每次更新都是一个很小的修正量(比如1e-6量级),如果主权重是FP16,更新的步长小于FP16的精度分辨率时,这次更新就白做了。FP16的尾数只有10位,一个数越大,能表示的最小区间越大,更新量很容易被“吃掉”。FP32主权重相当于给整个训练过程一个可靠的参考系。
2.3 损失缩放(Loss Scaling)的实际操作
虽然BF16缓解了下溢问题,但FP16仍然怎么绕都绕不开。这就是损失缩放登场的原因。
在混合精度训练的早期版本里(尤其是FP16时代),训练loss经过反向传播得到梯度后,很多梯度的绝对值已经小于FP16能表示的最小正常数(约6×10^-8),基本等同于0。为了不让它们消失,做法是在前向传播时先把loss放大,比如乘以1024,反向传播时梯度整体也被放大了,等梯度存回FP32再加回来。
实操中一般推荐动态损失缩放(Dynamic Loss Scaling),框架会自动监测梯度的溢出情况。如果一段时间内没有溢出,就逐步把缩放因子调大,比如从1024调到2048;如果检测到有inf或nan,就回退并调小。PyTorch的torch.cuda.amp.GradScaler就是这个逻辑,autocast负责把算子自动切换到低精度,GradScaler负责缩放梯度。
我用过一个比较稳妥的组合:
- 混合精度策略:
bfloat16(如果是Ampere及以上架构) - 损失缩放:BF16其实不需要动态缩放,因为它的指数范围够大,但FP16必须配
如果你在V100上只能支持FP16,别偷懒,一定要上动态缩放。实测中不缩放的FP16训练,十次有八次会因为梯度下溢而loss不动。
3. 分布式训练:把一张卡装不下的模型“拆开”
分布式训练的本质,是把模型、数据、梯度拆到多张卡上并行计算,再用通信把结果汇聚起来。这里最大的误解在于:分布式训练不只是“多个GPU跑同一个脚本”,而是有不同层级的并行策略,选择哪种策略直接决定了训练效率。
3.1 数据并行:最直观的并行方式
数据并行(Data Parallelism)是最直觉的思路。每张卡都复制一份完整模型,喂进去的数据是不同的batch。前向计算各自算,反向计算出梯度后,要经过一次全reduce操作把不同卡上的梯度加起来求平均,然后每张卡再用这个平均梯度更新自己的模型副本。
数据平行的通信开销是梯度同步,每轮每张卡都要收发全部模型参数对应的梯度。模型越大,通信量越大。当模型大到单卡装不下时,纯数据并行就失效了。另外还有个细节:数据并行通常要求每张卡上的batch size一致,总batch size = 单卡 batch size × 卡数。如果总batch size过大,收敛效果会有变化,一般会配合学习率调整。
3.2 模型并行与流水线并行:把模型切开
当模型超出单卡显存,数据并行就不够了,这时要把模型本身切开。这里有三类常说的并行:张量并行(如Megatron-LM的做法)、流水线并行(如GPipe、PipeDream)、以及深度学习的另一种分布式范式:ZeRO。
张量并行是把一个Layer里的权重矩阵按行或列切成多块,每张卡负责计算一部分,计算过程中每步都要做all-reduce来同步中间结果。通信频率极高,适合在服务器内部用NVLink组网的高带宽环境。流水线并行则是按Transformer层来切,第1-8层放卡A,第9-16层放卡B,数据像流水线一样依次通过各卡。它的通信频率低很多,但存在流水线气泡问题,各卡会有空闲等待时间。
你会发现这些策略本质都在做相同的一件事:用通信换显存。选择什么并行策略,取决于你的硬件拓扑和模型大小。
3.3 ZeRO:把冗余状态也算清楚
ZeRO(Zero Redundancy Optimizer)是微软DeepSpeed提出的方案,它洞察到一个数据并行中的关键问题:数据并行时,每张卡都存着一份完整的模型权重、梯度和优化器状态,这些是多卡间完全冗余的。
ZeRO把优化器状态、梯度、模型参数分段切开分到不同的卡上。三阶段是:
- ZeRO-1:切分优化器状态,显存节省约4倍;
- ZeRO-2:进一步切分梯度,显存节省约8倍;
- ZeRO-3:连模型参数也切分,训练时按需做all-gather取回。
ZeRO-1和ZeRO-2几乎不影响通信效率,因为它们本来同步梯度时就要通信。但ZeRO-3每个layer前向/反向都要实时收集该层权重,通信量比数据并行高了不少。所以ZeRO-3通常会配合梯度检查点(activation checkpointing)和更精细的通信调度来用。
4. 实战配置与踩坑记录:从跑通到跑快的心路历程
原理听再多,不上手永远是纸上谈兵。这一节我记录一次实际调参过程,用的是一个70亿参数模型,在8张A100(80GB)上的训练配置。机器环境是CUDA 12.1、PyTorch 2.1、DeepSpeed 0.12。
4.1 第一步:显存估算与模型策略选择
70亿参数按FP32主权重算,权重28GB、梯度28GB、Adam优化器状态56GB,合计112GB。这就是为什么单张80GB卡完全装不下的原因。如果开了BF16,训练时每张卡上的进程仍会保留一份FP32主权重和优化器状态,所以112GB这个数值并不是按卡数均分的“总额”,而是实实在在每张卡都要维护的一份。
用上ZeRO-1后,优化器状态被切成8份,每张卡只存一份 56GB/8 = 7GB,显存压力大幅降低。如果只是做ZeRO-1,每张卡仍然需要存完整的28GB权重和28GB梯度,加上中间激活,感觉还是比较紧。稳妥起见我直接选了ZeRO-2并开了activation checkpointing。激活值这一块非常占显存,尤其是序列长度长的时候,开启activation checkpointing后以约30%的计算开销换来大量显存回退,这买卖划算。
最终显存占用大约是:FP32权重28GB + 梯度(切分后 28/8=3.5GB)+ 优化器状态(56/8=7GB) + 激活约10GB。总计约50GB,在80GB卡上余量充足。这个余量给了batch size和sequence length调整空间,训练起来踏实得多。
4.2 第二步:DeepSpeed配置解析
DeepSpeed的配置文件是训练流程里最关键的一份文件,直接决定显存怎么分、通信怎么优化。我当时用的核心配置长这样:
{ "train_batch_size": 64, "train_micro_batch_size_per_gpu": 8, "gradient_accumulation_steps": 1, "fp16": { "enabled": false }, "bf16": { "enabled": true }, "zero_optimization": { "stage": 2, "offload_optimizer": { "device": "none" }, "contiguous_gradients": true, "overlap_comm": true }, "activation_checkpointing": { "partition_activations": true, "cpu_checkpointing": false }, "communication_data_type": "fp16" }几个关键点值得解释一下。
train_batch_size是全局的总batch size,它等于gradient_accumulation_steps × train_micro_batch_size_per_gpu × 卡数。用64 = 1 × 8 × 8,正好对上。如果你改了这个数值但没同步改另外两个,训练可能会直接报错或者出现batch对不上的问题。
bf16.enabled和fp16.enabled只能二选一。在高阶卡上我强烈建议开BF16。我当时在A100上对比过,FP16需要额外维护GradScaler,而且因为动态损失缩放的存在,偶尔会看到loss出现小幅跳变;BF16则不需要,训练过程平滑了很多。
offload_optimizer默认是none,即优化器状态保持在显存里。如果你显存实在吃紧,可以把optimizer状态offload到CPU内存,速度会慢一些,但能把单卡显存占用压到极低。我当时80GB够用,就没开offload。
contiguous_gradients和overlap_comm是DeepSpeed特别有用的两个开关。前者把梯度整理成连续内存块,后者让梯度计算与通信重叠。开启后整体吞吐能提升10%-15%,属于白捡的优化。
4.3 第三步:启动训练的命令与日志解读
DeepSpeed启动命令通常长这样:
deepspeed --num_gpus=8 train.py \ --deepspeed ds_config.json \ --model_name_or_path path/to/model \ --per_device_train_batch_size 8 \ --gradient_accumulation_steps 1 \ --learning_rate 1e-5 \ --bf16跑起来以后,我盯了几个关键指标:
cpu_mem和gpu_mem:确认显存占用与你预估相符。如果显存占用持续增长,先查是不是缓存泄漏;loss:在混合精度场景下,loss出现nan或inf,九成是梯度溢出或学习率过大;throughput(每秒处理的样本数):如果吞吐偏低,优先检查通信瓶颈,其次是数据加载。
我第一次把8卡跑通的时候,发现有个卡的利用率只有60%多。排查了半天,发现是数据加载用了默认的dataloader,num_workers太小,每轮到卡取数据都要等。把num_workers调到16并开启prefetch_factor之后,GPU利用率拉到了90%以上。这个坑很基础但很多人都会踩。
4.4 第四步:通信瓶颈的识别与解决
分布式训练最大的隐形杀手是通信。有时候你看到8张卡都在跑,但吞吐就是上不去,此时八成是卡在梯度同步上。
判断方法很简单:在DeepSpeed配置里把overlap_comm关闭,如果训练时间反而变短了,说明你的通信和计算其实是串行的,并且通信时间占比太高。这时可以优先做三件事:
- 检查是否用的NCCL通信库,并确认网络接口选择正确;
- 调大
train_micro_batch_size_per_gpu,增大单卡计算粒度,减少通信频率; - 在ZeRO-2下把梯度切分打开,用多张卡的通信冗余换取更小的单次通信量。
我调参时把micro batch从4调到8,吞吐直接提升了近30%,因为通信被分摊到了更大的计算量上。这也是为什么很多框架教程反复强调:分布式训练优先调大单卡batch size,再去想其他花活。
5. 常见问题与排查技巧实录
训练大模型的过程,本质就是不断遇见新bug、不断排查的过程。以下是我踩过的一些坑,整理成表格方便你对照。
| 故障现象 | 可能原因 | 解决方案 |
|---|---|---|
| loss为nan,训练直接崩 | 梯度溢出;学习率过大 | 开BF16或动态损失缩放;降低学习率重启 |
| loss长期不变化 | FP16梯度下溢 | 换BF16;开FP16的GradScaler |
| 显存OOM | batch size过大;激活值过多 | 减小micro batch;开启activation checkpointing;ZeRO升级到Stage 2/3 |
| 多卡吞吐远低于预期 | 数据加载瓶颈;通信瓶颈 | 调大num_workers;检查NCCL;增大单卡batch size |
| 不同卡显存占用严重不均 | 流水线切分不平衡 | 改用ZeRO或调整层切分策略 |
| BF16的loss比FP32略高 | 正常现象 | 无需处理,收敛趋势正常即可 |
还有一个很有意思的坑,是我在第一次尝试ZeRO-3时碰到的。模型权重被切分后,每张卡按需去别的卡上取参数,训练速度比ZeRO-2慢了不少。当时我以为配置有问题,后来才发现ZeRO-3必须配合gradient_checkpointing和精心调好的communication_data_type,否则通信量会大到拖垮训练。实测下来,70B以上的模型ZeRO-3才真正划算,7B-13B这个量级ZeRO-2性价比通常更好。
另外提醒一下,千万不要在混合精度训练中用手动乘一个缩放系数来代替框架的GradScaler/autocast。这种看似“控制力更强”的做法,实际很容易遗漏某些算子导致精度不一致,而且排查起来极难定位。框架提供的工具是经过大规模验证的,直接信任它。
6. 训练完成后的模型保存与继续训练技巧
模型训练到一定阶段要保存checkpoint,这里面也有讲究。混合精度训练下,pytorch的save默认保存的是FP32主权重,这没问题。但如果你用的是DeepSpeed,checkpoint的保存和加载最好也走DeepSpeed的接口,否则权重切分状态会和优化器状态对不上,暖启动(从checkpoint继续训练)时会出现各种怪问题。
我习惯每隔500步保存一个checkpoint,如果一个checkpoint保存失败,整个训练就白跑了。保存频率不是越高越好,因为保存checkpoint会打断训练流程、产生额外IO开销。500步对我来说是一个在恢复时间和性能损失之间比较平衡的值。
暖启动时的学习率也要注意。如果是从头训练,前面一两千步一般用warmup把学习率缓慢升高,避免早期梯度震荡;如果是加载checkpoint继续训练,可以把学习率降为原先的50%-70%。我一般会在训练中断后给一个更小的学习率,让loss先从“恢复期”平稳过渡回正常下降轨道。
还有一个容易被忽略的点:多卡训练时,checkpoint只保存主卡的数据。加载时如果用单卡直接加载,需要先把分布式环境初始化好,否则torch.load时会因为缺少分布式上下文而报错。这些细节往往不会写在框架文档里,但实际跑训练时几乎都会遇到。